diff --git a/aibridge/bridge.go b/aibridge/bridge.go index 835c125b6c..e351f58f49 100644 --- a/aibridge/bridge.go +++ b/aibridge/bridge.go @@ -253,12 +253,16 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC return } + cred := interceptor.Credential() traceAttrs := interceptor.TraceAttributes(r) span.SetAttributes(traceAttrs...) ctx = tracing.WithInterceptionAttributesInContext(ctx, traceAttrs) - // Attach the interception ID to the context so every log line - // emitted with this context can be correlated to the interception. - ctx = slog.With(ctx, slog.F("interception_id", interceptor.ID())) + // Attach the interception ID and credential kind to the context so every + // log line emitted with it can be correlated to the interception. + ctx = slog.With(ctx, + slog.F("interception_id", interceptor.ID()), + slog.F("credential_kind", string(cred.Kind())), + ) r = r.WithContext(ctx) // Record usage in the background to not block request flow. @@ -270,7 +274,6 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC asyncRecorder.WithClient(string(client)) interceptor.Setup(logger, asyncRecorder, mcpProxy) - cred := interceptor.Credential() if err := rec.RecordInterception(ctx, &recorder.InterceptionRecord{ ID: interceptor.ID().String(), InitiatorID: actor.ID, @@ -282,8 +285,8 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC Client: string(client), ClientSessionID: sessionID, CorrelatingToolCallID: interceptor.CorrelatingToolCallID(), - CredentialKind: string(cred.Kind), - CredentialHint: cred.Hint, + CredentialKind: string(cred.Kind()), + CredentialHint: cred.Hint(), }); err != nil { span.SetStatus(codes.Error, fmt.Sprintf("failed to record interception: %v", err)) logger.Warn(ctx, "failed to record interception", slog.Error(err)) @@ -297,19 +300,12 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC slog.F("provider", p.Name()), slog.F("user_agent", r.UserAgent()), slog.F("streaming", interceptor.Streaming()), - slog.F("credential_kind", string(cred.Kind)), ) - // Log BYOK credentials. Centralized credentials are set by - // the key failover loop. - credLogFields := []slog.Field{} - if cred.Kind == intercept.CredentialKindBYOK { - credLogFields = append(credLogFields, - slog.F("credential_hint", cred.Hint), - slog.F("credential_length", cred.Length), - ) - } - log.Debug(ctx, "interception started", credLogFields...) + log.Debug(ctx, "interception started", + slog.F("credential_hint", cred.Hint()), + slog.F("credential_length", cred.Length()), + ) if m != nil { m.InterceptionsInflight.WithLabelValues(p.Name(), interceptor.Model(), route).Add(1) defer func() { @@ -321,26 +317,25 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC execErr := cbs.Execute(route, interceptor.Model(), w, func(rw http.ResponseWriter) error { return interceptor.ProcessRequest(rw, r) }) - // For centralized, the hint now reflects the last attempted - // key from the failover loop. - credHint := interceptor.Credential().Hint - credLen := interceptor.Credential().Length + // For a centralized pool, the hint now reflects the last key the + // failover loop attempted. + credCtx := intercept.WithCredentialInfo(ctx, cred) if execErr != nil { if m != nil { m.InterceptionCount.WithLabelValues(p.Name(), interceptor.Model(), metrics.InterceptionCountStatusFailed, route, r.Method, actor.ID, string(client)).Add(1) } span.SetStatus(codes.Error, fmt.Sprintf("interception failed: %v", execErr)) - log.Warn(ctx, "interception failed", slog.Error(execErr), slog.F("credential_hint", credHint), slog.F("credential_length", credLen)) + log.Warn(credCtx, "interception failed", slog.Error(execErr)) } else { if m != nil { m.InterceptionCount.WithLabelValues(p.Name(), interceptor.Model(), metrics.InterceptionCountStatusCompleted, route, r.Method, actor.ID, string(client)).Add(1) } - log.Debug(ctx, "interception ended", slog.F("credential_hint", credHint), slog.F("credential_length", credLen)) + log.Debug(credCtx, "interception ended") } _ = asyncRecorder.RecordInterceptionEnded(ctx, &recorder.InterceptionRecordEnded{ ID: interceptor.ID().String(), - CredentialHint: credHint, + CredentialHint: cred.Hint(), }) // Ensure all recording have completed before completing request. diff --git a/aibridge/config/config.go b/aibridge/config/config.go index 5805741f60..9fdeb32d57 100644 --- a/aibridge/config/config.go +++ b/aibridge/config/config.go @@ -13,31 +13,16 @@ const ( ) // Anthropic carries configuration for an Anthropic provider. -// -// Authentication is mutually exclusive across these three fields, -// set per interception in the provider's CreateInterceptor: -// - KeyPool: centralized requests with automatic key failover. -// - Key: BYOK with X-Api-Key (single attempt, no failover). -// - BYOKBearerToken: BYOK with Authorization Bearer (single -// attempt, no failover). -// -// TODO(ssncferreira): consolidate the three authentication -// fields into a single abstraction per -// https://github.com/coder/aibridge/issues/266. type Anthropic struct { // Name is the provider instance name. If empty, defaults to "anthropic". - Name string - BaseURL string - Key string + Name string + BaseURL string + // KeyPool holds the centralized keys, with automatic key failover. BYOK + // credentials are resolved per request from the incoming headers. KeyPool *keypool.Pool APIDumpDir string CircuitBreaker *CircuitBreaker SendActorHeaders bool - ExtraHeaders map[string]string - // BYOKBearerToken is set in BYOK mode when the user authenticates - // with a access token. When set, the access token is used for upstream - // LLM requests instead of the API key. - BYOKBearerToken string } type AWSBedrock struct { @@ -51,25 +36,16 @@ type AWSBedrock struct { } // OpenAI carries configuration for an OpenAI provider. -// -// Authentication is mutually exclusive across these two fields, -// set per interception in the provider's CreateInterceptor: -// - KeyPool: centralized requests with automatic key failover. -// - Key: BYOK with Authorization Bearer (single attempt, no -// failover). -// -// TODO(ssncferreira): consolidate the authentication fields per -// https://github.com/coder/aibridge/issues/266. type OpenAI struct { // Name is the provider instance name. If empty, defaults to "openai". - Name string - BaseURL string - Key string + Name string + BaseURL string + // KeyPool holds the centralized keys, with automatic key failover. BYOK + // credentials are resolved per request from the incoming headers. KeyPool *keypool.Pool APIDumpDir string CircuitBreaker *CircuitBreaker SendActorHeaders bool - ExtraHeaders map[string]string } type Copilot struct { diff --git a/aibridge/intercept/chatcompletions/base.go b/aibridge/intercept/chatcompletions/base.go index 4a7efcb1d5..e2c4a82ad4 100644 --- a/aibridge/intercept/chatcompletions/base.go +++ b/aibridge/intercept/chatcompletions/base.go @@ -17,7 +17,6 @@ import ( "go.opentelemetry.io/otel/trace" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/apidump" @@ -29,56 +28,47 @@ import ( ) type interceptionBase struct { - id uuid.UUID - providerName string - req *ChatCompletionNewParamsWrapper - cfg config.OpenAI + id uuid.UUID + req *ChatCompletionNewParamsWrapper + + cfg intercept.Config + cred intercept.Credential // clientHeaders are the original HTTP headers from the client request. - clientHeaders http.Header - authHeaderName string + clientHeaders http.Header logger slog.Logger tracer trace.Tracer - recorder recorder.Recorder - mcpProxy mcp.ServerProxier - credential intercept.CredentialInfo + recorder recorder.Recorder + mcpProxy mcp.ServerProxier } -// newCompletionsService builds the SDK service used for upstream -// calls. BYOK auth is set here. Centralized auth is set -// per-attempt by the failover loop. -func (i *interceptionBase) newCompletionsService() openai.ChatCompletionService { - // TODO(ssncferreira): validate auth is configured per - // https://github.com/coder/aibridge/issues/266. - +// newCompletionsService builds the SDK service used for upstream calls. +func (i *interceptionBase) newCompletionsService(ctx context.Context) openai.ChatCompletionService { var opts []option.RequestOption - // BYOK auth. - if i.cfg.KeyPool == nil { - opts = append(opts, option.WithAPIKey(i.cfg.Key)) + // Only BYOK sets its credential here. Centralized keys are injected + // per-attempt in the failover loop. + if byok, ok := intercept.AsBYOK(i.cred); ok { + i.logger.Debug(ctx, "using byok auth", + slog.F("auth_header", byok.Header), slog.F("key_hint", byok.Hint()), + ) + opts = append(opts, option.WithAPIKey(byok.Secret)) } opts = append(opts, option.WithBaseURL(i.cfg.BaseURL)) - // Add extra headers if configured. - // Some providers require additional headers that are not added by the SDK. - // TODO(ssncferreira): remove as part of https://github.com/coder/aibridge/issues/192 - for key, value := range i.cfg.ExtraHeaders { - opts = append(opts, option.WithHeader(key, value)) - } - // Forward client headers to upstream. This middleware runs after the SDK // has built the request, and replaces the outgoing headers with the sanitized // client headers plus provider auth. if i.clientHeaders != nil { opts = append(opts, option.WithMiddleware(func(req *http.Request, next option.MiddlewareNext) (*http.Response, error) { - req.Header = intercept.BuildUpstreamHeaders(req.Header, i.clientHeaders, i.authHeaderName) + req.Header = intercept.BuildUpstreamHeaders(req.Header, i.clientHeaders, i.cred.AuthHeader()) return next(req) })) } // Add API dump middleware if configured - if mw := apidump.NewBridgeMiddleware(i.cfg.APIDumpDir, i.providerName, i.Model(), i.id, i.logger, quartz.NewReal()); mw != nil { + if mw := apidump.NewBridgeMiddleware(i.cfg.APIDumpDir, i.cfg.ProviderName, i.Model(), i.id, i.logger, quartz.NewReal()); mw != nil { opts = append(opts, option.WithMiddleware(mw)) } @@ -89,8 +79,8 @@ func (i *interceptionBase) ID() uuid.UUID { return i.id } -func (i *interceptionBase) Credential() intercept.CredentialInfo { - return i.credential +func (i *interceptionBase) Credential() intercept.Credential { + return i.cred } func (i *interceptionBase) Setup(logger slog.Logger, rec recorder.Recorder, mcpProxy mcp.ServerProxier) { @@ -117,7 +107,7 @@ func (i *interceptionBase) baseTraceAttributes(r *http.Request, streaming bool) attribute.String(tracing.RequestPath, r.URL.Path), attribute.String(tracing.InterceptionID, i.id.String()), attribute.String(tracing.InitiatorID, aibcontext.ActorIDFromContext(r.Context())), - attribute.String(tracing.Provider, i.providerName), + attribute.String(tracing.Provider, i.cfg.ProviderName), attribute.String(tracing.Model, i.Model()), attribute.Bool(tracing.Streaming, streaming), } @@ -219,14 +209,15 @@ func (i *interceptionBase) writeUpstreamError(w http.ResponseWriter, oaiErr *int // code. Returns true if the status was a key-specific failover // trigger so callers can retry with the next key. func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key, err error) bool { - if i.cfg.KeyPool == nil { + cp, ok := intercept.AsCentralizedPool(i.cred) + if !ok { return false } var apiErr *openai.Error if !errors.As(err, &apiErr) { return false } - return i.cfg.KeyPool.MarkKeyOnStatus( + return cp.Pool.MarkKeyOnStatus( ctx, key, apiErr.Response, i.logger, ) } diff --git a/aibridge/intercept/chatcompletions/base_internal_test.go b/aibridge/intercept/chatcompletions/base_internal_test.go index c56d39db10..af31a59bdc 100644 --- a/aibridge/intercept/chatcompletions/base_internal_test.go +++ b/aibridge/intercept/chatcompletions/base_internal_test.go @@ -141,7 +141,7 @@ func TestMarkKeyOnError(t *testing.T) { key, keyPoolErr := pool.Walker().Next() require.Nil(t, keyPoolErr) - base := &interceptionBase{cfg: config.OpenAI{KeyPool: pool}, logger: slog.Make()} + base := &interceptionBase{cred: &intercept.CentralizedPool{Pool: pool}, logger: slog.Make()} got := base.markKeyOnError(context.Background(), key, tc.err) assert.Equal(t, tc.expectedReturn, got) diff --git a/aibridge/intercept/chatcompletions/blocking.go b/aibridge/intercept/chatcompletions/blocking.go index 29cc8bc704..b3d829377b 100644 --- a/aibridge/intercept/chatcompletions/blocking.go +++ b/aibridge/intercept/chatcompletions/blocking.go @@ -16,7 +16,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -33,22 +32,18 @@ type BlockingInterception struct { func NewBlockingInterceptor( id uuid.UUID, req *ChatCompletionNewParamsWrapper, - providerName string, - cfg config.OpenAI, + cfg intercept.Config, + cred intercept.Credential, clientHeaders http.Header, - authHeaderName string, tracer trace.Tracer, - cred intercept.CredentialInfo, ) *BlockingInterception { return &BlockingInterception{interceptionBase: interceptionBase{ - id: id, - providerName: providerName, - req: req, - cfg: cfg, - clientHeaders: clientHeaders, - authHeaderName: authHeaderName, - tracer: tracer, - credential: cred, + id: id, + req: req, + cfg: cfg, + cred: cred, + clientHeaders: clientHeaders, + tracer: tracer, }} } @@ -72,7 +67,7 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req ctx, span := i.tracer.Start(r.Context(), "Intercept.ProcessRequest", trace.WithAttributes(tracing.InterceptionAttributesFromContext(r.Context())...)) defer tracing.EndSpanErr(span, &outErr) - svc := i.newCompletionsService() + svc := i.newCompletionsService(ctx) logger := i.logger.With(slog.F("model", i.req.Model)) var ( @@ -91,7 +86,11 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req // Sum the key attempts across all iterations and record once when the // interception completes. var totalKeyAttempts int - defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }() + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + defer func() { + cp.Pool.RecordAttempts(totalKeyAttempts) + }() + } for { // TODO add outer loop span (https://github.com/coder/aibridge/issues/67) @@ -274,16 +273,16 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req return nil } -// newChatCompletion routes between BYOK (single attempt) and centralized -// failover, returning the upstream completion, the number of key attempts -// made for this call, and any error. +// newChatCompletion routes by credential type, returning the upstream +// completion, the number of key attempts made for this call, and any error. A +// centralized key pool fails over across keys, while BYOK authenticates with a +// single, fixed credential baked into svc, so it makes one attempt. func (i *BlockingInterception) newChatCompletion(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, int, error) { - // BYOK: single attempt, no failover. - if i.cfg.KeyPool == nil { - completion, err := i.newChatCompletionWithKey(ctx, svc, opts) - return completion, 0, err + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + return i.newChatCompletionWithKeyFailover(ctx, svc, cp, opts) } - return i.newChatCompletionWithKeyFailover(ctx, svc, opts) + completion, err := i.newChatCompletionWithKey(intercept.WithCredentialInfo(ctx, i.cred), svc, opts) + return completion, 0, err } // newChatCompletionWithKey performs a single upstream call. @@ -307,18 +306,16 @@ func (i *BlockingInterception) newChatCompletionWithKey(ctx context.Context, svc // on 429 and permanent on 401/403. Errors that aren't key-specific don't // trigger failover and are returned to the caller. It returns the upstream // completion, the number of key attempts made for this call, and any error. -func (i *BlockingInterception) newChatCompletionWithKeyFailover(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, int, error) { - walker := i.cfg.KeyPool.Walker() +func (i *BlockingInterception) newChatCompletionWithKeyFailover(ctx context.Context, svc openai.ChatCompletionService, cp *intercept.CentralizedPool, opts []option.RequestOption) (*openai.ChatCompletion, int, error) { + walker := cp.Pool.Walker() for { - key, keyPoolErr := walker.Next() + key, keyPoolErr := cp.NextKey(walker) if keyPoolErr != nil { return nil, walker.Attempts(), keyPoolErr } - // Record the key in use so the hint reflects the last attempted key. - i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) - i.logger.Debug(ctx, "using centralized api key", - slog.F("credential_hint", i.Credential().Hint), slog.F("credential_length", i.Credential().Length)) + ctx = intercept.WithCredentialInfo(ctx, i.cred) + i.logger.Debug(ctx, "using centralized api key") requestOpts := append([]option.RequestOption{}, opts...) requestOpts = append(requestOpts, option.WithAPIKey(key.Value()), diff --git a/aibridge/intercept/chatcompletions/google_openai_compat_internal_test.go b/aibridge/intercept/chatcompletions/google_openai_compat_internal_test.go index 826dba07b6..a8acd5397a 100644 --- a/aibridge/intercept/chatcompletions/google_openai_compat_internal_test.go +++ b/aibridge/intercept/chatcompletions/google_openai_compat_internal_test.go @@ -7,7 +7,7 @@ import ( "github.com/openai/openai-go/v3/option" "github.com/stretchr/testify/require" - "github.com/coder/coder/v2/aibridge/config" + "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/internal/googleopenai" ) @@ -46,7 +46,7 @@ func TestGoogleOpenAICompatThoughtSignaturePatchSurvivesParamRoundTrip(t *testin body, err := (&interceptionBase{ req: &req, - cfg: config.OpenAI{BaseURL: "https://generativelanguage.googleapis.com/v1beta/openai/"}, + cfg: intercept.Config{BaseURL: "https://generativelanguage.googleapis.com/v1beta/openai/"}, }).chatCompletionRequestBody() require.NoError(t, err) require.Equal(t, googleopenai.DummyThoughtSignature, googleThoughtSignatureFromBody(t, body, 1, 0)) @@ -70,7 +70,7 @@ func TestGoogleOpenAICompatChatCompletionRequestOptions(t *testing.T) { opts := make([]option.RequestOption, 1) updated, overrideBody, err := (&interceptionBase{ req: &req, - cfg: config.OpenAI{BaseURL: "https://generativelanguage.googleapis.com/v1beta/openai/"}, + cfg: intercept.Config{BaseURL: "https://generativelanguage.googleapis.com/v1beta/openai/"}, }).chatCompletionRequestOptions(opts) require.NoError(t, err) require.True(t, overrideBody) diff --git a/aibridge/intercept/chatcompletions/streaming.go b/aibridge/intercept/chatcompletions/streaming.go index 694b09893f..bf76bdcfb4 100644 --- a/aibridge/intercept/chatcompletions/streaming.go +++ b/aibridge/intercept/chatcompletions/streaming.go @@ -20,7 +20,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -38,22 +37,18 @@ type StreamingInterception struct { func NewStreamingInterceptor( id uuid.UUID, req *ChatCompletionNewParamsWrapper, - providerName string, - cfg config.OpenAI, + cfg intercept.Config, + cred intercept.Credential, clientHeaders http.Header, - authHeaderName string, tracer trace.Tracer, - cred intercept.CredentialInfo, ) *StreamingInterception { return &StreamingInterception{interceptionBase: interceptionBase{ - id: id, - providerName: providerName, - req: req, - cfg: cfg, - clientHeaders: clientHeaders, - authHeaderName: authHeaderName, - tracer: tracer, - credential: cred, + id: id, + req: req, + cfg: cfg, + cred: cred, + clientHeaders: clientHeaders, + tracer: tracer, }} } @@ -99,7 +94,7 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re defer cancel() r = r.WithContext(ctx) // Rewire context for SSE cancellation. - svc := i.newCompletionsService() + svc := i.newCompletionsService(ctx) logger := i.logger.With(slog.F("model", i.req.Model)) streamCtx, streamCancel := context.WithCancelCause(ctx) @@ -131,30 +126,29 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re // Sum the key attempts across all iterations and record once when the // interception completes. var totalKeyAttempts int - defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }() + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + defer func() { + cp.Pool.RecordAttempts(totalKeyAttempts) + }() + } for { // TODO add outer loop span (https://github.com/coder/aibridge/issues/67) - // Per-iteration walker. An iteration is either an agentic - // continuation (sending a tool result back in a new - // stream) or a failover retry (previous key marked, try - // the next one). - var walker *keypool.Walker - if i.cfg.KeyPool != nil { - walker = i.cfg.KeyPool.Walker() - } - + // Per-iteration: a pool credential advances its failover walker. An + // iteration is either an agentic continuation or a failover retry after + // the previous key was marked. BYOK has no pool and runs as a single + // attempt. var opts []option.RequestOption - var currentKey *keypool.Key - if walker != nil { - key, keyPoolErr := walker.Next() + var currentPoolKey *keypool.Key + if cp, isPool := intercept.AsCentralizedPool(i.cred); isPool { + walker := cp.Pool.Walker() + key, keyPoolErr := cp.NextKey(walker) if keyPoolErr != nil { + // Pool exhausted in this iteration. Relay the error to the + // client: as an SSE event if events have already been sent, + // or by direct write otherwise. respErr := intercept.ResponseErrorFromKeyPool(keyPoolErr) - // Pool exhausted in this iteration. Relay the - // error to the client: as an SSE event if events - // have already been sent, or by direct write - // otherwise. interceptionErr = respErr if events.IsStreaming() { payload, mErr := i.marshalErr(respErr) @@ -168,22 +162,18 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re } break } - currentKey = key - // Record the key in use so the hint reflects the last attempted key. - i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) - logger.Debug(ctx, "using centralized api key", - slog.F("credential_hint", i.Credential().Hint), slog.F("credential_length", i.Credential().Length)) + logger.Debug(intercept.WithCredentialInfo(ctx, i.cred), "using centralized api key") + currentPoolKey = key opts = append(opts, option.WithAPIKey(key.Value()), - // Disable SDK retries because the failover - // loop handles retries via key rotation. + // Disable SDK retries because the failover loop handles + // retries via key rotation. option.WithMaxRetries(0), ) + totalKeyAttempts += walker.Attempts() } - totalKeyAttempts += walker.Attempts() - // TODO(ssncferreira): inject actor headers directly in the client-header // middleware instead of using SDK options. if actor := aibcontext.ActorFromContext(r.Context()); actor != nil && i.cfg.SendActorHeaders { @@ -319,7 +309,7 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re // Pre-stream failure of this iteration. For // centralized requests, mark the key and retry with // the next one. - if currentKey != nil && i.markKeyOnError(ctx, currentKey, stream.Err()) { + if currentPoolKey != nil && i.markKeyOnError(ctx, currentPoolKey, stream.Err()) { continue } // Non-key error: relay it. Use mapStreamError so that diff --git a/aibridge/intercept/chatcompletions/streaming_internal_test.go b/aibridge/intercept/chatcompletions/streaming_internal_test.go index 2feea1a709..1f6e419555 100644 --- a/aibridge/intercept/chatcompletions/streaming_internal_test.go +++ b/aibridge/intercept/chatcompletions/streaming_internal_test.go @@ -13,7 +13,6 @@ import ( "cdr.dev/slog/v3" "cdr.dev/slog/v3/sloggers/slogtest" - "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/internal/testutil" ) @@ -67,10 +66,10 @@ func TestStreamingInterception_RelaysUpstreamErrorToClient(t *testing.T) { t.Cleanup(mockServer.Close) // Create interceptor with mock server URL - cfg := config.OpenAI{ + cfg := intercept.Config{ BaseURL: mockServer.URL, - Key: "test-key", } + cred := intercept.BYOK{Secret: "test-key", Header: intercept.AuthHeaderAuthorization} req := &ChatCompletionNewParamsWrapper{ ChatCompletionNewParams: openai.ChatCompletionNewParams{ @@ -87,7 +86,7 @@ func TestStreamingInterception_RelaysUpstreamErrorToClient(t *testing.T) { httpReq := httptest.NewRequest(http.MethodPost, "/chat/completions", nil) tracer := otel.Tracer("test") - interceptor := NewStreamingInterceptor(uuid.New(), req, config.ProviderOpenAI, cfg, httpReq.Header, "Authorization", tracer, intercept.CredentialInfo{}) + interceptor := NewStreamingInterceptor(uuid.New(), req, cfg, cred, httpReq.Header, tracer) logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: false}).Leveled(slog.LevelDebug) interceptor.Setup(logger, &testutil.MockRecorder{}, nil) diff --git a/aibridge/intercept/config.go b/aibridge/intercept/config.go new file mode 100644 index 0000000000..d2d42d957f --- /dev/null +++ b/aibridge/intercept/config.go @@ -0,0 +1,19 @@ +package intercept + +// Config is the per-request configuration an interceptor needs to process +// an interception, independent of which provider produced it. Providers +// resolve it in CreateInterceptor and hand it to the API-format +// interceptor. +type Config struct { + // ProviderName is the provider instance name, used for recording, + // logging, and API dumps. + ProviderName string + // BaseURL is the upstream provider's API base URL. + BaseURL string + // APIDumpDir is the directory for dumping API requests and responses, + // or empty when API dumping is disabled. + APIDumpDir string + // SendActorHeaders reports whether actor identity headers should be + // forwarded to the upstream provider. + SendActorHeaders bool +} diff --git a/aibridge/intercept/credential.go b/aibridge/intercept/credential.go index 3343245e38..008c43463d 100644 --- a/aibridge/intercept/credential.go +++ b/aibridge/intercept/credential.go @@ -1,31 +1,140 @@ package intercept -import "github.com/coder/coder/v2/aibridge/utils" +import ( + "context" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/aibridge/keypool" + "github.com/coder/coder/v2/aibridge/utils" +) // CredentialKind identifies how a request was authenticated. // Keep in sync with the credential_kind enum in coderd's database. type CredentialKind string -// Credential kind constants for interception recording. const ( CredentialKindCentralized CredentialKind = "centralized" CredentialKindBYOK CredentialKind = "byok" ) -// CredentialInfo holds credential metadata for an interception. -type CredentialInfo struct { - Kind CredentialKind - Hint string - Length int +// Auth header names shared by providers (which set them on resolved +// credentials) and interceptors (which present credentials under them). +const ( + AuthHeaderXAPIKey = "X-Api-Key" //nolint:gosec // G101 false positive: HTTP header name, not a credential. + AuthHeaderAuthorization = "Authorization" +) + +// Hint placeholders for credentials with no static key value to mask: a pool +// before failover selects a key, and a key resolved dynamically at request time. +const ( + hintFailoverKey = "" + hintBedrockChainKey = "" +) + +// Credential is the per-request upstream authentication for an interception: +// - BYOK: a user-supplied secret. +// - Bedrock: AWS Bedrock credentials, used to sign requests. +// - CentralizedPool: a provider-managed key pool with failover. +type Credential interface { + Kind() CredentialKind + // AuthHeader is the header carrying this request's credential, or empty when + // the credential is not carried in a header. + AuthHeader() string + // Hint is a masked, identifiable fragment of the credential. + Hint() string + // Length is the length of the credential value. + Length() int } -// NewCredentialInfo creates a CredentialInfo from a raw credential. -// The credential is automatically masked before storage so that the -// original secret is never retained. -func NewCredentialInfo(kind CredentialKind, credential string) CredentialInfo { - return CredentialInfo{ - Kind: kind, - Hint: utils.MaskSecret(credential), - Length: len(credential), +// BYOK authenticates with a single user-supplied secret. +type BYOK struct { + Secret string + Header string +} + +func (BYOK) Kind() CredentialKind { return CredentialKindBYOK } +func (b BYOK) AuthHeader() string { return b.Header } +func (b BYOK) Hint() string { return utils.MaskSecret(b.Secret) } +func (b BYOK) Length() int { return len(b.Secret) } + +// Bedrock authenticates with AWS Bedrock: requests are signed (so there is no +// auth header) using either static credentials (when an access key is set) or +// the AWS default credential chain. There is no key pool or failover. +type Bedrock struct { + AccessKey string +} + +func (Bedrock) Kind() CredentialKind { return CredentialKindCentralized } +func (Bedrock) AuthHeader() string { return "" } +func (b Bedrock) Length() int { return len(b.AccessKey) } + +func (b Bedrock) Hint() string { + if b.AccessKey == "" { + return hintBedrockChainKey } + return utils.MaskSecret(b.AccessKey) +} + +// CentralizedPool authenticates with a provider-managed key pool and fails over +// across keys. +type CentralizedPool struct { + Pool *keypool.Pool + Header string + // currentKey is the key most recently handed out by NextKey, nil until the first call. + currentKey *keypool.Key +} + +func (*CentralizedPool) Kind() CredentialKind { return CredentialKindCentralized } +func (c *CentralizedPool) AuthHeader() string { return c.Header } + +func (c *CentralizedPool) Hint() string { + if c.currentKey != nil { + return c.currentKey.Hint() + } + return hintFailoverKey +} + +func (c *CentralizedPool) Length() int { + if c.currentKey != nil { + return c.currentKey.Length() + } + return 0 +} + +// NextKey advances the failover walker and records the selected key as the one +// in use. +func (c *CentralizedPool) NextKey(w *keypool.Walker) (*keypool.Key, *keypool.Error) { + key, err := w.Next() + if err != nil { + return nil, err + } + c.currentKey = key + return key, nil +} + +var ( + _ Credential = BYOK{} + _ Credential = Bedrock{} + _ Credential = &CentralizedPool{} +) + +// AsBYOK reports whether c is a BYOK credential and returns it if so. +func AsBYOK(c Credential) (BYOK, bool) { + b, ok := c.(BYOK) + return b, ok +} + +// AsCentralizedPool reports whether c is a key-pool credential that fails over, +// and returns it if so. +func AsCentralizedPool(c Credential) (*CentralizedPool, bool) { + pool, ok := c.(*CentralizedPool) + return pool, ok +} + +// WithCredentialInfo returns a context carrying the credential hint and length. +func WithCredentialInfo(ctx context.Context, cred Credential) context.Context { + return slog.With(ctx, + slog.F("credential_hint", cred.Hint()), + slog.F("credential_length", cred.Length()), + ) } diff --git a/aibridge/intercept/credential_test.go b/aibridge/intercept/credential_test.go new file mode 100644 index 0000000000..0f78d8dee0 --- /dev/null +++ b/aibridge/intercept/credential_test.go @@ -0,0 +1,138 @@ +package intercept_test + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/aibridge/config" + "github.com/coder/coder/v2/aibridge/intercept" + "github.com/coder/coder/v2/aibridge/keypool" + "github.com/coder/quartz" +) + +// TestCredential covers the public surface of the Credential interface and its +// three implementations (BYOK, Bedrock, CentralizedPool), plus the +// AsBYOK/AsCentralizedPool helpers interceptors use to route. Only +// CentralizedPool fails over, so AsCentralizedPool must be true only for it. +func TestCredential(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + newCred func(t *testing.T) intercept.Credential + expectKind intercept.CredentialKind + expectAuthHeader string + expectHint string + expectLength int + expectAsBYOK bool + expectAsCentralizedPool bool + }{ + { + name: "byok_authorization", + newCred: func(*testing.T) intercept.Credential { + return intercept.BYOK{Secret: "user-bearer-token", Header: intercept.AuthHeaderAuthorization} + }, + expectKind: intercept.CredentialKindBYOK, + expectAuthHeader: intercept.AuthHeaderAuthorization, + expectHint: "us...en", + expectLength: len("user-bearer-token"), + expectAsBYOK: true, + }, + { + name: "byok_xapikey", + newCred: func(*testing.T) intercept.Credential { + return intercept.BYOK{Secret: "user-api-key", Header: intercept.AuthHeaderXAPIKey} + }, + expectKind: intercept.CredentialKindBYOK, + expectAuthHeader: intercept.AuthHeaderXAPIKey, + expectHint: "us...ey", + expectLength: len("user-api-key"), + expectAsBYOK: true, + }, + { + // Bedrock with static AWS credentials: the access key ID is + // masked. AWS signs the request, so there is no auth header. + name: "centralized_bedrock_static", + newCred: func(*testing.T) intercept.Credential { + return intercept.Bedrock{AccessKey: "AKIAIOSFODNN7EXAMPLE"} + }, + expectKind: intercept.CredentialKindCentralized, + expectAuthHeader: "", + expectHint: "AKIA...MPLE", + expectLength: len("AKIAIOSFODNN7EXAMPLE"), + }, + { + // Bedrock with dynamic credentials (AWS default credential chain): + // no static key to mask, so the hint is a descriptive placeholder. + name: "centralized_bedrock_dynamic", + newCred: func(*testing.T) intercept.Credential { + return intercept.Bedrock{AccessKey: ""} + }, + expectKind: intercept.CredentialKindCentralized, + expectAuthHeader: "", + expectHint: "", + expectLength: 0, + }, + { + // Pool before failover selects a key: the hint is a placeholder + // until NextKey hands one out. + name: "centralized_pool_before_key", + newCred: func(t *testing.T) intercept.Credential { + pool, err := keypool.New(config.ProviderAnthropic, []string{"k0-pool-key"}, quartz.NewMock(t), nil) + require.NoError(t, err) + return &intercept.CentralizedPool{Pool: pool, Header: intercept.AuthHeaderXAPIKey} + }, + expectKind: intercept.CredentialKindCentralized, + expectAuthHeader: intercept.AuthHeaderXAPIKey, + expectHint: "", + expectLength: 0, + expectAsCentralizedPool: true, + }, + { + // Pool after NextKey: Hint/Length reflect the selected key. + name: "centralized_pool_after_next_key", + newCred: func(t *testing.T) intercept.Credential { + pool, err := keypool.New(config.ProviderAnthropic, []string{"k0-pool-key"}, quartz.NewMock(t), nil) + require.NoError(t, err) + cp := &intercept.CentralizedPool{Pool: pool, Header: intercept.AuthHeaderXAPIKey} + _, keyErr := cp.NextKey(cp.Pool.Walker()) + require.Nil(t, keyErr) + return cp + }, + expectKind: intercept.CredentialKindCentralized, + expectAuthHeader: intercept.AuthHeaderXAPIKey, + expectHint: "k0...ey", + expectLength: len("k0-pool-key"), + expectAsCentralizedPool: true, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + cred := tc.newCred(t) + + assert.Equal(t, tc.expectKind, cred.Kind(), "Kind") + assert.Equal(t, tc.expectAuthHeader, cred.AuthHeader(), "AuthHeader") + assert.Equal(t, tc.expectHint, cred.Hint(), "Hint") + assert.Equal(t, tc.expectLength, cred.Length(), "Length") + + credBYOK, credBYOKOK := intercept.AsBYOK(cred) + assert.Equal(t, tc.expectAsBYOK, credBYOKOK, "AsBYOK ok") + if tc.expectAsBYOK { + assert.Equal(t, cred, credBYOK, "AsBYOK returns the credential") + } + + credPool, credPoolOK := intercept.AsCentralizedPool(cred) + assert.Equal(t, tc.expectAsCentralizedPool, credPoolOK, "AsCentralizedPool ok") + if tc.expectAsCentralizedPool { + assert.Same(t, cred, credPool, "AsCentralizedPool returns the same pointer") + } else { + assert.Nil(t, credPool, "AsCentralizedPool returns nil when not a pool") + } + }) + } +} diff --git a/aibridge/intercept/interceptor.go b/aibridge/intercept/interceptor.go index 33cbc51dff..23c6890a31 100644 --- a/aibridge/intercept/interceptor.go +++ b/aibridge/intercept/interceptor.go @@ -26,8 +26,10 @@ type Interceptor interface { Streaming() bool // TraceAttributes returns tracing attributes for this [Interceptor] TraceAttributes(*http.Request) []attribute.KeyValue - // Credential returns the credential metadata for this interception. - Credential() CredentialInfo + // Credential returns the credential resolved for this interception. Its + // Hint/Length reflect the key in use (the last failover key for a pool + // credential, otherwise the static credential), for logs and records. + Credential() Credential // CorrelatingToolCallID returns the ID of a tool call result submitted // in the request, if present. This is used to correlate the current // interception back to the previous interception that issued those tool diff --git a/aibridge/intercept/keyfailover_test.go b/aibridge/intercept/keyfailover_test.go index 997ca82705..a20b14d343 100644 --- a/aibridge/intercept/keyfailover_test.go +++ b/aibridge/intercept/keyfailover_test.go @@ -88,13 +88,15 @@ var interceptorCases = []interceptorCase{ agenticStreamErrorEvent: "event: error", streamDoneEvent: "event: message_stop", newInterceptor: func(t *testing.T, streaming bool, upstreamURL string, reqBody []byte, pool *keypool.Pool, byokKey string) intercept.Interceptor { - cfg := config.Anthropic{BaseURL: upstreamURL + "/"} - cred := intercept.NewCredentialInfo(intercept.CredentialKindCentralized, "") + var cred intercept.Credential if pool != nil { - cfg.KeyPool = pool - } else if byokKey != "" { - cfg.Key = byokKey - cred = intercept.NewCredentialInfo(intercept.CredentialKindBYOK, byokKey) + cred = &intercept.CentralizedPool{Pool: pool, Header: "X-Api-Key"} + } else { + cred = intercept.BYOK{Secret: byokKey, Header: "X-Api-Key"} + } + cfg := intercept.Config{ + ProviderName: config.ProviderAnthropic, + BaseURL: upstreamURL + "/", } payload, err := messages.NewRequestPayload(reqBody) @@ -102,9 +104,9 @@ var interceptorCases = []interceptorCase{ id, tracer := uuid.New(), otel.Tracer("keyfailover") if streaming { - return messages.NewStreamingInterceptor(id, payload, config.ProviderAnthropic, cfg, nil, http.Header{}, "X-Api-Key", tracer, cred) + return messages.NewStreamingInterceptor(id, payload, cfg, cred, nil, http.Header{}, tracer) } - return messages.NewBlockingInterceptor(id, payload, config.ProviderAnthropic, cfg, nil, http.Header{}, "X-Api-Key", tracer, cred) + return messages.NewBlockingInterceptor(id, payload, cfg, cred, nil, http.Header{}, tracer) }, }, { @@ -121,13 +123,15 @@ var interceptorCases = []interceptorCase{ agenticStreamErrorEvent: `data: {"error"`, streamDoneEvent: "data: [DONE]", newInterceptor: func(t *testing.T, streaming bool, upstreamURL string, reqBody []byte, pool *keypool.Pool, byokKey string) intercept.Interceptor { - cfg := config.OpenAI{BaseURL: upstreamURL + "/"} - cred := intercept.NewCredentialInfo(intercept.CredentialKindCentralized, "") + var cred intercept.Credential if pool != nil { - cfg.KeyPool = pool - } else if byokKey != "" { - cfg.Key = byokKey - cred = intercept.NewCredentialInfo(intercept.CredentialKindBYOK, byokKey) + cred = &intercept.CentralizedPool{Pool: pool, Header: "Authorization"} + } else { + cred = intercept.BYOK{Secret: byokKey, Header: "Authorization"} + } + cfg := intercept.Config{ + ProviderName: config.ProviderOpenAI, + BaseURL: upstreamURL + "/", } var req chatcompletions.ChatCompletionNewParamsWrapper @@ -135,9 +139,9 @@ var interceptorCases = []interceptorCase{ id, tracer := uuid.New(), otel.Tracer("keyfailover") if streaming { - return chatcompletions.NewStreamingInterceptor(id, &req, config.ProviderOpenAI, cfg, http.Header{}, "Authorization", tracer, cred) + return chatcompletions.NewStreamingInterceptor(id, &req, cfg, cred, http.Header{}, tracer) } - return chatcompletions.NewBlockingInterceptor(id, &req, config.ProviderOpenAI, cfg, http.Header{}, "Authorization", tracer, cred) + return chatcompletions.NewBlockingInterceptor(id, &req, cfg, cred, http.Header{}, tracer) }, }, { @@ -159,13 +163,15 @@ var interceptorCases = []interceptorCase{ }, streamDoneEvent: "event: response.completed", newInterceptor: func(t *testing.T, streaming bool, upstreamURL string, reqBody []byte, pool *keypool.Pool, byokKey string) intercept.Interceptor { - cfg := config.OpenAI{BaseURL: upstreamURL + "/"} - cred := intercept.NewCredentialInfo(intercept.CredentialKindCentralized, "") + var cred intercept.Credential if pool != nil { - cfg.KeyPool = pool - } else if byokKey != "" { - cfg.Key = byokKey - cred = intercept.NewCredentialInfo(intercept.CredentialKindBYOK, byokKey) + cred = &intercept.CentralizedPool{Pool: pool, Header: "Authorization"} + } else { + cred = intercept.BYOK{Secret: byokKey, Header: "Authorization"} + } + cfg := intercept.Config{ + ProviderName: config.ProviderOpenAI, + BaseURL: upstreamURL + "/", } payload, err := responses.NewRequestPayload(reqBody) @@ -173,9 +179,9 @@ var interceptorCases = []interceptorCase{ id, tracer := uuid.New(), otel.Tracer("keyfailover") if streaming { - return responses.NewStreamingInterceptor(id, payload, config.ProviderOpenAI, cfg, http.Header{}, "Authorization", tracer, cred) + return responses.NewStreamingInterceptor(id, payload, cfg, cred, http.Header{}, tracer) } - return responses.NewBlockingInterceptor(id, payload, config.ProviderOpenAI, cfg, http.Header{}, "Authorization", tracer, cred) + return responses.NewBlockingInterceptor(id, payload, cfg, cred, http.Header{}, tracer) }, }, } @@ -371,7 +377,7 @@ func TestInterception_KeyFailover(t *testing.T) { if len(tc.expectedSeenKeys) > 0 { assert.Equal(t, utils.MaskSecret(tc.expectedSeenKeys[len(tc.expectedSeenKeys)-1]), - interceptor.Credential().Hint, "credential hint") + interceptor.Credential().Hint(), "credential hint") } if tc.expectedBodyContains != "" { assert.Contains(t, w.Body.String(), tc.expectedBodyContains, "response body") @@ -546,7 +552,7 @@ func TestInterception_AgenticLoopFailover(t *testing.T) { if len(tc.expectedSeenKeys) > 0 { assert.Equal(t, utils.MaskSecret(tc.expectedSeenKeys[len(tc.expectedSeenKeys)-1]), - interceptor.Credential().Hint, "credential hint") + interceptor.Credential().Hint(), "credential hint") } if tc.expectedBodyContains != "" { assert.Contains(t, w.Body.String(), tc.expectedBodyContains, "response body") diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index b167df4293..b7b414dc8d 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -66,31 +66,30 @@ var bedrockSupportedBetaFlags = map[string]bool{ } type interceptionBase struct { - id uuid.UUID - providerName string - reqPayload RequestPayload + id uuid.UUID + reqPayload RequestPayload - cfg aibconfig.Anthropic + cfg intercept.Config + cred intercept.Credential bedrockCfg *aibconfig.AWSBedrock // clientHeaders are the original HTTP headers from the client request. - clientHeaders http.Header - authHeaderName string + clientHeaders http.Header - tracer trace.Tracer logger slog.Logger + tracer trace.Tracer - recorder recorder.Recorder - mcpProxy mcp.ServerProxier - credential intercept.CredentialInfo + recorder recorder.Recorder + mcpProxy mcp.ServerProxier } func (i *interceptionBase) ID() uuid.UUID { return i.id } -func (i *interceptionBase) Credential() intercept.CredentialInfo { - return i.credential +// Credential returns the credential resolved for this interception. +func (i *interceptionBase) Credential() intercept.Credential { + return i.cred } func (i *interceptionBase) Setup(logger slog.Logger, rec recorder.Recorder, mcpProxy mcp.ServerProxier) { @@ -124,7 +123,7 @@ func (i *interceptionBase) baseTraceAttributes(r *http.Request, streaming bool) attribute.String(tracing.RequestPath, r.URL.Path), attribute.String(tracing.InterceptionID, i.id.String()), attribute.String(tracing.InitiatorID, aibcontext.ActorIDFromContext(r.Context())), - attribute.String(tracing.Provider, i.providerName), + attribute.String(tracing.Provider, i.cfg.ProviderName), attribute.String(tracing.Model, i.Model()), attribute.Bool(tracing.Streaming, streaming), attribute.Bool(tracing.IsBedrock, i.bedrockCfg != nil), @@ -205,50 +204,37 @@ func (i *interceptionBase) isSmallFastModel() bool { return strings.Contains(i.reqPayload.model(), "haiku") } -// newMessagesService builds the SDK service used for upstream -// calls. BYOK auth is set here. Centralized auth is set -// per-attempt by the failover loop. +// newMessagesService builds the SDK service used for upstream calls. func (i *interceptionBase) newMessagesService(ctx context.Context, opts ...option.RequestOption) (anthropic.MessageService, error) { - // TODO(ssncferreira): validate auth is configured per - // https://github.com/coder/aibridge/issues/266. - - // BYOK auth. - if i.cfg.KeyPool == nil { - if i.cfg.BYOKBearerToken != "" { - // BYOK Bearer: Authorization header. - i.logger.Debug(ctx, "using byok access token auth", - slog.F("bearer_hint", utils.MaskSecret(i.cfg.BYOKBearerToken)), - ) - opts = append(opts, option.WithAuthToken(i.cfg.BYOKBearerToken)) - } else { - // BYOK X-Api-Key. - i.logger.Debug(ctx, "using api key auth", - slog.F("api_key_hint", utils.MaskSecret(i.cfg.Key)), - ) - opts = append(opts, option.WithAPIKey(i.cfg.Key)) + // Only BYOK sets its credential here. Centralized keys are injected + // per-attempt in the failover loop. + if byok, ok := intercept.AsBYOK(i.cred); ok { + i.logger.Debug(ctx, "using byok auth", + slog.F("auth_header", byok.Header), slog.F("key_hint", byok.Hint()), + ) + switch byok.Header { + case intercept.AuthHeaderAuthorization: + opts = append(opts, option.WithAuthToken(byok.Secret)) + case intercept.AuthHeaderXAPIKey: + opts = append(opts, option.WithAPIKey(byok.Secret)) + default: + return anthropic.MessageService{}, xerrors.Errorf("unexpected byok auth header: %q", byok.Header) } } opts = append(opts, option.WithBaseURL(i.cfg.BaseURL)) - // Add extra headers if configured. - // Some providers require additional headers that are not added by the SDK. - // TODO(ssncferreira): remove as part of https://github.com/coder/aibridge/issues/192 - for key, value := range i.cfg.ExtraHeaders { - opts = append(opts, option.WithHeader(key, value)) - } - // Forward client headers to upstream. This middleware runs after the SDK // has built the request, and replaces the outgoing headers with the sanitized // client headers plus provider auth. if i.clientHeaders != nil { opts = append(opts, option.WithMiddleware(func(req *http.Request, next option.MiddlewareNext) (*http.Response, error) { - req.Header = intercept.BuildUpstreamHeaders(req.Header, i.clientHeaders, i.authHeaderName) + req.Header = intercept.BuildUpstreamHeaders(req.Header, i.clientHeaders, i.cred.AuthHeader()) return next(req) })) } // Add API dump middleware if configured - if mw := apidump.NewBridgeMiddleware(i.cfg.APIDumpDir, i.providerName, i.Model(), i.id, i.logger, quartz.NewReal()); mw != nil { + if mw := apidump.NewBridgeMiddleware(i.cfg.APIDumpDir, i.cfg.ProviderName, i.Model(), i.id, i.logger, quartz.NewReal()); mw != nil { opts = append(opts, option.WithMiddleware(mw)) } @@ -576,14 +562,15 @@ func accumulateUsage(dest, src any) { // its status code. Returns true if the status was a key-specific // failover trigger so callers can retry with the next key. func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key, err error) bool { - if i.cfg.KeyPool == nil { + cp, ok := intercept.AsCentralizedPool(i.cred) + if !ok { return false } var apiErr *anthropic.Error if !errors.As(err, &apiErr) { return false } - return i.cfg.KeyPool.MarkKeyOnStatus( + return cp.Pool.MarkKeyOnStatus( ctx, key, apiErr.Response, i.logger, ) } diff --git a/aibridge/intercept/messages/base_internal_test.go b/aibridge/intercept/messages/base_internal_test.go index f6323ec795..41133836fa 100644 --- a/aibridge/intercept/messages/base_internal_test.go +++ b/aibridge/intercept/messages/base_internal_test.go @@ -16,6 +16,7 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/aibridge/config" + "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/keypool" "github.com/coder/coder/v2/aibridge/mcp" @@ -1127,7 +1128,7 @@ func TestMarkKeyOnError(t *testing.T) { key, keyPoolErr := pool.Walker().Next() require.Nil(t, keyPoolErr) - base := &interceptionBase{cfg: config.Anthropic{KeyPool: pool}, logger: slog.Make()} + base := &interceptionBase{cred: &intercept.CentralizedPool{Pool: pool}, logger: slog.Make()} got := base.markKeyOnError(context.Background(), key, tc.err) assert.Equal(t, tc.expectedReturn, got) diff --git a/aibridge/intercept/messages/blocking.go b/aibridge/intercept/messages/blocking.go index 4370676ce7..8ee118b5da 100644 --- a/aibridge/intercept/messages/blocking.go +++ b/aibridge/intercept/messages/blocking.go @@ -17,7 +17,7 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" + aibconfig "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -34,24 +34,20 @@ type BlockingInterception struct { func NewBlockingInterceptor( id uuid.UUID, reqPayload RequestPayload, - providerName string, - cfg config.Anthropic, - bedrockCfg *config.AWSBedrock, + cfg intercept.Config, + cred intercept.Credential, + bedrockCfg *aibconfig.AWSBedrock, clientHeaders http.Header, - authHeaderName string, tracer trace.Tracer, - cred intercept.CredentialInfo, ) *BlockingInterception { return &BlockingInterception{interceptionBase: interceptionBase{ - id: id, - providerName: providerName, - reqPayload: reqPayload, - cfg: cfg, - bedrockCfg: bedrockCfg, - clientHeaders: clientHeaders, - authHeaderName: authHeaderName, - tracer: tracer, - credential: cred, + id: id, + reqPayload: reqPayload, + cfg: cfg, + cred: cred, + bedrockCfg: bedrockCfg, + clientHeaders: clientHeaders, + tracer: tracer, }} } @@ -108,7 +104,11 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req // Sum the key attempts across all iterations and record once when the // interception completes. var totalKeyAttempts int - defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }() + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + defer func() { + cp.Pool.RecordAttempts(totalKeyAttempts) + }() + } for { // TODO add outer loop span (https://github.com/coder/aibridge/issues/67) @@ -349,16 +349,16 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req return nil } -// newMessage routes between BYOK (single attempt) and centralized -// failover, returning the upstream message, the number of key attempts -// made for this call, and any error. +// newMessage routes by credential type, returning the upstream message, the +// number of key attempts made for this call, and any error. A centralized key +// pool fails over across keys, while BYOK and Bedrock authenticate with a +// single, fixed credential baked into svc, so they make one attempt. func (i *BlockingInterception) newMessage(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, int, error) { - // BYOK: single attempt, no failover. - if i.cfg.KeyPool == nil { - msg, err := i.newMessageWithKey(ctx, svc) - return msg, 0, err + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + return i.newMessageWithKeyFailover(ctx, svc, cp) } - return i.newMessageWithKeyFailover(ctx, svc) + msg, err := i.newMessageWithKey(intercept.WithCredentialInfo(ctx, i.cred), svc) + return msg, 0, err } // newMessageWithKey performs a single upstream call. @@ -375,18 +375,16 @@ func (i *BlockingInterception) newMessageWithKey(ctx context.Context, svc anthro // 429 and permanent on 401/403. Errors that aren't key-specific don't trigger // failover and are returned to the caller. It returns the upstream message, // the number of key attempts made for this call, and any error. -func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, int, error) { - walker := i.cfg.KeyPool.Walker() +func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, svc anthropic.MessageService, cp *intercept.CentralizedPool) (*anthropic.Message, int, error) { + walker := cp.Pool.Walker() for { - key, keyPoolErr := walker.Next() + key, keyPoolErr := cp.NextKey(walker) if keyPoolErr != nil { return nil, walker.Attempts(), keyPoolErr } - // Record the key in use so the hint reflects the last attempted key. - i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) - i.logger.Debug(ctx, "using centralized api key", - slog.F("credential_hint", i.Credential().Hint), slog.F("credential_length", i.Credential().Length)) + ctx = intercept.WithCredentialInfo(ctx, i.cred) + i.logger.Debug(ctx, "using centralized api key") msg, err := i.newMessageWithKey(ctx, svc, option.WithAPIKey(key.Value()), // Disable SDK retries because the failover loop diff --git a/aibridge/intercept/messages/streaming.go b/aibridge/intercept/messages/streaming.go index 1a383889e3..22d0dff720 100644 --- a/aibridge/intercept/messages/streaming.go +++ b/aibridge/intercept/messages/streaming.go @@ -21,7 +21,7 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" + aibconfig "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -39,24 +39,20 @@ type StreamingInterception struct { func NewStreamingInterceptor( id uuid.UUID, reqPayload RequestPayload, - providerName string, - cfg config.Anthropic, - bedrockCfg *config.AWSBedrock, + cfg intercept.Config, + cred intercept.Credential, + bedrockCfg *aibconfig.AWSBedrock, clientHeaders http.Header, - authHeaderName string, tracer trace.Tracer, - cred intercept.CredentialInfo, ) *StreamingInterception { return &StreamingInterception{interceptionBase: interceptionBase{ - id: id, - providerName: providerName, - reqPayload: reqPayload, - cfg: cfg, - bedrockCfg: bedrockCfg, - clientHeaders: clientHeaders, - authHeaderName: authHeaderName, - tracer: tracer, - credential: cred, + id: id, + reqPayload: reqPayload, + cfg: cfg, + cred: cred, + bedrockCfg: bedrockCfg, + clientHeaders: clientHeaders, + tracer: tracer, }} } @@ -102,6 +98,7 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re // Allow us to interrupt watch via cancel. ctx, cancel := context.WithCancel(ctx) defer cancel() + r = r.WithContext(ctx) // Rewire context for SSE cancellation. logger := i.logger.With(slog.F("model", i.Model())) @@ -156,7 +153,11 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re // Sum the key attempts across all iterations and record once when the // interception completes. var totalKeyAttempts int - defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }() + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + defer func() { + cp.Pool.RecordAttempts(totalKeyAttempts) + }() + } isFirst := true newStream: @@ -167,24 +168,19 @@ newStream: break } - // Per-iteration walker. An iteration is either an agentic - // continuation (sending a tool result back in a new - // stream) or a failover retry (previous key marked, try - // the next one). - var walker *keypool.Walker - if i.cfg.KeyPool != nil { - walker = i.cfg.KeyPool.Walker() - } - + // Per-iteration: a pool credential advances its failover walker. An + // iteration is either an agentic continuation or a failover retry after + // the previous key was marked. BYOK and Bedrock have no pool and run as + // a single attempt. var streamOpts []option.RequestOption - var currentKey *keypool.Key - if walker != nil { - key, keyPoolErr := walker.Next() + var currentPoolKey *keypool.Key + if cp, isPool := intercept.AsCentralizedPool(i.cred); isPool { + walker := cp.Pool.Walker() + key, keyPoolErr := cp.NextKey(walker) if keyPoolErr != nil { - // Pool exhausted in this iteration. Relay the - // error to the client: as an SSE event if events - // have already been sent, or by direct write - // otherwise. + // Pool exhausted in this iteration. Relay the error to the + // client: as an SSE event if events have already been sent, + // or by direct write otherwise. respErr := ResponseErrorFromKeyPool(keyPoolErr) interceptionErr = respErr if events.IsStreaming() { @@ -199,22 +195,18 @@ newStream: } break } - currentKey = key - // Record the key in use so the hint reflects the last attempted key. - i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) - logger.Debug(ctx, "using centralized api key", - slog.F("credential_hint", i.Credential().Hint), slog.F("credential_length", i.Credential().Length)) + logger.Debug(intercept.WithCredentialInfo(ctx, i.cred), "using centralized api key") + currentPoolKey = key streamOpts = append(streamOpts, option.WithAPIKey(key.Value()), - // Disable SDK retries because the failover - // loop handles retries via key rotation. + // Disable SDK retries because the failover loop handles + // retries via key rotation. option.WithMaxRetries(0), ) + totalKeyAttempts += walker.Attempts() } - totalKeyAttempts += walker.Attempts() - stream := i.newStream(streamCtx, svc, streamOpts...) var message anthropic.Message @@ -568,7 +560,7 @@ newStream: // Pre-stream failure of this iteration. For // centralized requests, mark the key and retry with // the next one. - if currentKey != nil && i.markKeyOnError(ctx, currentKey, stream.Err()) { + if currentPoolKey != nil && i.markKeyOnError(ctx, currentPoolKey, stream.Err()) { continue newStream } // Non-key error: relay it. Use mapStreamError so that diff --git a/aibridge/intercept/responses/base.go b/aibridge/intercept/responses/base.go index 6b4521739d..0154457795 100644 --- a/aibridge/intercept/responses/base.go +++ b/aibridge/intercept/responses/base.go @@ -26,7 +26,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/apidump" @@ -42,55 +41,47 @@ const ( ) type responsesInterceptionBase struct { - id uuid.UUID - providerName string + id uuid.UUID + reqPayload RequestPayload + + cfg intercept.Config + cred intercept.Credential + // clientHeaders are the original HTTP headers from the client request. - clientHeaders http.Header - authHeaderName string - reqPayload RequestPayload + clientHeaders http.Header + + logger slog.Logger + tracer trace.Tracer - cfg config.OpenAI recorder recorder.Recorder mcpProxy mcp.ServerProxier - - logger slog.Logger - tracer trace.Tracer - credential intercept.CredentialInfo } -// newResponsesService builds the SDK service used for upstream -// calls. BYOK auth is set here. Centralized auth is set -// per-attempt by the failover loop. -func (i *responsesInterceptionBase) newResponsesService() responses.ResponseService { - // TODO(ssncferreira): validate auth is configured per - // https://github.com/coder/aibridge/issues/266. - +// newResponsesService builds the SDK service used for upstream calls. +func (i *responsesInterceptionBase) newResponsesService(ctx context.Context) responses.ResponseService { var opts []option.RequestOption - // BYOK auth. - if i.cfg.KeyPool == nil { - opts = append(opts, option.WithAPIKey(i.cfg.Key)) + // Only BYOK sets its credential here. Centralized keys are injected + // per-attempt in the failover loop. + if byok, ok := intercept.AsBYOK(i.cred); ok { + i.logger.Debug(ctx, "using byok auth", + slog.F("auth_header", byok.Header), slog.F("key_hint", byok.Hint()), + ) + opts = append(opts, option.WithAPIKey(byok.Secret)) } opts = append(opts, option.WithBaseURL(i.cfg.BaseURL)) - // Add extra headers if configured. - // Some providers require additional headers that are not added by the SDK. - // TODO(ssncferreira): remove as part of https://github.com/coder/aibridge/issues/192 - for key, value := range i.cfg.ExtraHeaders { - opts = append(opts, option.WithHeader(key, value)) - } - // Forward client headers to upstream. This middleware runs after the SDK // has built the request, and replaces the outgoing headers with the sanitized // client headers plus provider auth. if i.clientHeaders != nil { opts = append(opts, option.WithMiddleware(func(req *http.Request, next option.MiddlewareNext) (*http.Response, error) { - req.Header = intercept.BuildUpstreamHeaders(req.Header, i.clientHeaders, i.authHeaderName) + req.Header = intercept.BuildUpstreamHeaders(req.Header, i.clientHeaders, i.cred.AuthHeader()) return next(req) })) } // Add API dump middleware if configured - if mw := apidump.NewBridgeMiddleware(i.cfg.APIDumpDir, i.providerName, i.Model(), i.id, i.logger, quartz.NewReal()); mw != nil { + if mw := apidump.NewBridgeMiddleware(i.cfg.APIDumpDir, i.cfg.ProviderName, i.Model(), i.id, i.logger, quartz.NewReal()); mw != nil { opts = append(opts, option.WithMiddleware(mw)) } @@ -101,8 +92,8 @@ func (i *responsesInterceptionBase) ID() uuid.UUID { return i.id } -func (i *responsesInterceptionBase) Credential() intercept.CredentialInfo { - return i.credential +func (i *responsesInterceptionBase) Credential() intercept.Credential { + return i.cred } func (i *responsesInterceptionBase) Setup(logger slog.Logger, rec recorder.Recorder, mcpProxy mcp.ServerProxier) { @@ -124,7 +115,7 @@ func (i *responsesInterceptionBase) baseTraceAttributes(r *http.Request, streami attribute.String(tracing.RequestPath, r.URL.Path), attribute.String(tracing.InterceptionID, i.id.String()), attribute.String(tracing.InitiatorID, aibcontext.ActorIDFromContext(r.Context())), - attribute.String(tracing.Provider, i.providerName), + attribute.String(tracing.Provider, i.cfg.ProviderName), attribute.String(tracing.Model, i.Model()), attribute.Bool(tracing.Streaming, streaming), } @@ -174,14 +165,15 @@ func (i *responsesInterceptionBase) writeUpstreamError(w http.ResponseWriter, oa // code. Returns true if the status was a key-specific failover // trigger so callers can retry with the next key. func (i *responsesInterceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key, err error) bool { - if i.cfg.KeyPool == nil { + cp, ok := intercept.AsCentralizedPool(i.cred) + if !ok { return false } var apiErr *openai.Error if !errors.As(err, &apiErr) { return false } - return i.cfg.KeyPool.MarkKeyOnStatus( + return cp.Pool.MarkKeyOnStatus( ctx, key, apiErr.Response, i.logger, ) } diff --git a/aibridge/intercept/responses/base_internal_test.go b/aibridge/intercept/responses/base_internal_test.go index 883db116e9..208ee0afc3 100644 --- a/aibridge/intercept/responses/base_internal_test.go +++ b/aibridge/intercept/responses/base_internal_test.go @@ -445,7 +445,7 @@ func TestMarkKeyOnError(t *testing.T) { key, keyPoolErr := pool.Walker().Next() require.Nil(t, keyPoolErr) - base := &responsesInterceptionBase{cfg: config.OpenAI{KeyPool: pool}, logger: slog.Make()} + base := &responsesInterceptionBase{cred: &intercept.CentralizedPool{Pool: pool}, logger: slog.Make()} got := base.markKeyOnError(context.Background(), key, tc.err) assert.Equal(t, tc.expectedReturn, got) diff --git a/aibridge/intercept/responses/blocking.go b/aibridge/intercept/responses/blocking.go index 2236cd3616..6038e2e1e6 100644 --- a/aibridge/intercept/responses/blocking.go +++ b/aibridge/intercept/responses/blocking.go @@ -14,7 +14,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/keypool" @@ -30,23 +29,19 @@ type BlockingResponsesInterceptor struct { func NewBlockingInterceptor( id uuid.UUID, reqPayload RequestPayload, - providerName string, - cfg config.OpenAI, + cfg intercept.Config, + cred intercept.Credential, clientHeaders http.Header, - authHeaderName string, tracer trace.Tracer, - cred intercept.CredentialInfo, ) *BlockingResponsesInterceptor { return &BlockingResponsesInterceptor{ responsesInterceptionBase: responsesInterceptionBase{ - id: id, - providerName: providerName, - reqPayload: reqPayload, - cfg: cfg, - clientHeaders: clientHeaders, - authHeaderName: authHeaderName, - tracer: tracer, - credential: cred, + id: id, + reqPayload: reqPayload, + cfg: cfg, + cred: cred, + clientHeaders: clientHeaders, + tracer: tracer, }, } } @@ -89,10 +84,14 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r * // Sum the key attempts across all iterations and record once when the // interception completes. var totalKeyAttempts int - defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }() + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + defer func() { + cp.Pool.RecordAttempts(totalKeyAttempts) + }() + } for shouldLoop { - srv := i.newResponsesService() + srv := i.newResponsesService(ctx) respCopy = responseCopier{} opts := i.requestOptions(&respCopy) @@ -153,16 +152,16 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r * return errors.Join(upstreamErr, err) } -// newResponse routes between BYOK (single attempt) and centralized failover, -// returning the upstream response, the number of key attempts made for this -// call, and any error. +// newResponse routes by credential type, returning the upstream response, the +// number of key attempts made for this call, and any error. A centralized key +// pool fails over across keys, while BYOK authenticates with a single, fixed +// credential baked into srv, so it makes one attempt. func (i *BlockingResponsesInterceptor) newResponse(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, int, error) { - // BYOK: single attempt, no failover. - if i.cfg.KeyPool == nil { - response, err := i.newResponseWithKey(ctx, srv, opts) - return response, 0, err + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + return i.newResponseWithKeyFailover(ctx, srv, cp, opts) } - return i.newResponseWithKeyFailover(ctx, srv, opts) + response, err := i.newResponseWithKey(intercept.WithCredentialInfo(ctx, i.cred), srv, opts) + return response, 0, err } // newResponseWithKey performs a single upstream call. @@ -179,18 +178,16 @@ func (i *BlockingResponsesInterceptor) newResponseWithKey(ctx context.Context, s // 429 and permanent on 401/403. Errors that aren't key-specific don't trigger // failover and are returned to the caller. It returns the upstream response, // the number of key attempts made for this call, and any error. -func (i *BlockingResponsesInterceptor) newResponseWithKeyFailover(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, int, error) { - walker := i.cfg.KeyPool.Walker() +func (i *BlockingResponsesInterceptor) newResponseWithKeyFailover(ctx context.Context, srv responses.ResponseService, cp *intercept.CentralizedPool, opts []option.RequestOption) (*responses.Response, int, error) { + walker := cp.Pool.Walker() for { - key, keyPoolErr := walker.Next() + key, keyPoolErr := cp.NextKey(walker) if keyPoolErr != nil { return nil, walker.Attempts(), keyPoolErr } - // Record the key in use so the hint reflects the last attempted key. - i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) - i.logger.Debug(ctx, "using centralized api key", - slog.F("credential_hint", i.Credential().Hint), slog.F("credential_length", i.Credential().Length)) + ctx = intercept.WithCredentialInfo(ctx, i.cred) + i.logger.Debug(ctx, "using centralized api key") requestOpts := append([]option.RequestOption{}, opts...) requestOpts = append(requestOpts, option.WithAPIKey(key.Value()), diff --git a/aibridge/intercept/responses/streaming.go b/aibridge/intercept/responses/streaming.go index 617cd144f1..492783f4de 100644 --- a/aibridge/intercept/responses/streaming.go +++ b/aibridge/intercept/responses/streaming.go @@ -16,7 +16,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/intercept/eventstream" @@ -38,23 +37,19 @@ type StreamingResponsesInterceptor struct { func NewStreamingInterceptor( id uuid.UUID, reqPayload RequestPayload, - providerName string, - cfg config.OpenAI, + cfg intercept.Config, + cred intercept.Credential, clientHeaders http.Header, - authHeaderName string, tracer trace.Tracer, - cred intercept.CredentialInfo, ) *StreamingResponsesInterceptor { return &StreamingResponsesInterceptor{ responsesInterceptionBase: responsesInterceptionBase{ - id: id, - providerName: providerName, - reqPayload: reqPayload, - cfg: cfg, - clientHeaders: clientHeaders, - authHeaderName: authHeaderName, - tracer: tracer, - credential: cred, + id: id, + reqPayload: reqPayload, + cfg: cfg, + cred: cred, + clientHeaders: clientHeaders, + tracer: tracer, }, } } @@ -77,6 +72,7 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r ctx, cancel := context.WithCancel(ctx) defer cancel() + r = r.WithContext(ctx) // Rewire context for SSE cancellation. if err := i.validateRequest(ctx, w); err != nil { @@ -104,23 +100,27 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r i.logger.Warn(ctx, "failed to get user prompt", slog.Error(err)) } shouldLoop := true - srv := i.newResponsesService() + srv := i.newResponsesService(ctx) // Sum the key attempts across all iterations and record once when the // interception completes. var totalKeyAttempts int - defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }() + if cp, ok := intercept.AsCentralizedPool(i.cred); ok { + defer func() { + cp.Pool.RecordAttempts(totalKeyAttempts) + }() + } for shouldLoop { shouldLoop = false - // Per-iteration walker. An iteration is either an agentic - // continuation (sending a tool result back in a new - // stream) or a failover retry (previous key marked, try - // the next one). + // A pool credential advances its failover walker. An iteration is an + // agentic continuation or a failover retry after the previous key was + // marked. BYOK has no pool and runs as a single attempt. var walker *keypool.Walker - if i.cfg.KeyPool != nil { - walker = i.cfg.KeyPool.Walker() + cp, isPool := intercept.AsCentralizedPool(i.cred) + if isPool { + walker = cp.Pool.Walker() } // Failover sub-loop: try keys until a stream starts @@ -137,9 +137,9 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r opts = append(opts, intercept.ActorHeadersAsOpenAIOpts(actor)...) } - var currentKey *keypool.Key - if walker != nil { - key, keyPoolErr := walker.Next() + var currentPoolKey *keypool.Key + if isPool && walker != nil { + key, keyPoolErr := cp.NextKey(walker) if keyPoolErr != nil { // Pool exhausted: write the error directly. In // agentic mode the inner loop buffers events @@ -149,12 +149,9 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r i.writeUpstreamError(w, intercept.ResponseErrorFromKeyPool(keyPoolErr)) return xerrors.Errorf("key pool exhausted: %w", keyPoolErr) } - currentKey = key - // Record the key in use so the hint reflects the last attempted key. - i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) - i.logger.Debug(ctx, "using centralized api key", - slog.F("credential_hint", i.Credential().Hint), slog.F("credential_length", i.Credential().Length)) + i.logger.Debug(intercept.WithCredentialInfo(ctx, i.cred), "using centralized api key") + currentPoolKey = key opts = append(opts, option.WithAPIKey(key.Value()), // Disable SDK retries because the failover @@ -168,7 +165,7 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r // Pre-stream failure of this attempt. For // centralized requests, mark the key and // retry with the next one. - if currentKey != nil && i.markKeyOnError(ctx, currentKey, upstreamErr) { + if currentPoolKey != nil && i.markKeyOnError(ctx, currentPoolKey, upstreamErr) { stream.Close() continue } @@ -181,7 +178,9 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r break } - totalKeyAttempts += walker.Attempts() + if isPool { + totalKeyAttempts += walker.Attempts() + } // func scope to defer steam.Close() err := func() error { diff --git a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go index afa6091e2a..1f5ddd71f3 100644 --- a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go +++ b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go @@ -17,6 +17,7 @@ import ( "github.com/stretchr/testify/require" "github.com/coder/coder/v2/aibridge/config" + "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/metrics" "github.com/coder/coder/v2/aibridge/provider" ) @@ -70,7 +71,7 @@ func TestCircuitBreaker_FullRecoveryCycle(t *testing.T) { createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { return provider.NewAnthropic(config.Anthropic{ BaseURL: baseURL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, }, nil) }, @@ -88,7 +89,7 @@ func TestCircuitBreaker_FullRecoveryCycle(t *testing.T) { createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { return provider.NewOpenAI(config.OpenAI{ BaseURL: baseURL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }) }, @@ -237,7 +238,7 @@ func TestCircuitBreaker_HalfOpenFailure(t *testing.T) { createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { return provider.NewAnthropic(config.Anthropic{ BaseURL: baseURL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, }, nil) }, @@ -254,7 +255,7 @@ func TestCircuitBreaker_HalfOpenFailure(t *testing.T) { createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { return provider.NewOpenAI(config.OpenAI{ BaseURL: baseURL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }) }, @@ -374,7 +375,7 @@ func TestCircuitBreaker_HalfOpenMaxRequests(t *testing.T) { createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { return provider.NewAnthropic(config.Anthropic{ BaseURL: baseURL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, }, nil) }, @@ -392,7 +393,7 @@ func TestCircuitBreaker_HalfOpenMaxRequests(t *testing.T) { createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { return provider.NewOpenAI(config.OpenAI{ BaseURL: baseURL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }) }, @@ -555,7 +556,7 @@ func TestCircuitBreaker_PerModelIsolation(t *testing.T) { bridgeServer := newBridgeTestServer(ctx, t, mockUpstream.URL, withCustomProvider(provider.NewAnthropic(config.Anthropic{ BaseURL: mockUpstream.URL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, }, nil)), withMetrics(m), diff --git a/aibridge/internal/integrationtest/helpers.go b/aibridge/internal/integrationtest/helpers.go index 7b6e80c903..4146ccb354 100644 --- a/aibridge/internal/integrationtest/helpers.go +++ b/aibridge/internal/integrationtest/helpers.go @@ -6,6 +6,7 @@ import ( "cdr.dev/slog/v3" "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/aibridge/config" + "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/recorder" ) @@ -13,7 +14,7 @@ import ( func anthropicCfg(url string, key string) config.Anthropic { return config.Anthropic{ BaseURL: url, - Key: key, + KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, key), } } @@ -39,7 +40,7 @@ func bedrockCfg(url string) *config.AWSBedrock { func openAICfg(url string, key string) config.OpenAI { return config.OpenAI{ BaseURL: url, - Key: key, + KeyPool: testutil.SingleKeyPool(config.ProviderOpenAI, key), } } diff --git a/aibridge/internal/testutil/mockprovider.go b/aibridge/internal/testutil/mockprovider.go index 8f6ce2a22c..a607ae68d0 100644 --- a/aibridge/internal/testutil/mockprovider.go +++ b/aibridge/internal/testutil/mockprovider.go @@ -10,8 +10,23 @@ import ( "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/keypool" + "github.com/coder/quartz" ) +// SingleKeyPool builds a centralized key pool containing a single key, or nil +// when key is empty (no centralized credential). It panics if the pool cannot +// be built, which does not happen for a non-empty key. +func SingleKeyPool(name, key string) *keypool.Pool { + if key == "" { + return nil + } + pool, err := keypool.New(name, []string{key}, quartz.NewReal(), nil) + if err != nil { + panic(err) + } + return pool +} + type MockProvider struct { NameStr string URL string diff --git a/aibridge/keypool/keypool.go b/aibridge/keypool/keypool.go index 746c0920f1..4ca4f76ba4 100644 --- a/aibridge/keypool/keypool.go +++ b/aibridge/keypool/keypool.go @@ -151,6 +151,11 @@ func (k *Key) Hint() string { return utils.MaskSecret(k.value) } +// Length returns the length of the key value, for logs. +func (k *Key) Length() int { + return len(k.value) +} + // State returns the current state of the key, derived from its // permanent flag and cooldown deadline. func (k *Key) State() KeyState { diff --git a/aibridge/provider/anthropic.go b/aibridge/provider/anthropic.go index 0757296f81..2b04bdf8b1 100644 --- a/aibridge/provider/anthropic.go +++ b/aibridge/provider/anthropic.go @@ -19,16 +19,8 @@ import ( "github.com/coder/coder/v2/aibridge/keypool" "github.com/coder/coder/v2/aibridge/tracing" "github.com/coder/coder/v2/aibridge/utils" - "github.com/coder/quartz" ) -// anthropicForwardHeaders lists headers from incoming requests that should be -// forwarded to the Anthropic API. -// TODO(ssncferreira): remove as part of https://github.com/coder/aibridge/issues/192 -var anthropicForwardHeaders = []string{ - "Anthropic-Beta", -} - var _ Provider = &Anthropic{} // Anthropic allows for interactions with the Anthropic API. @@ -58,24 +50,6 @@ func NewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropi if cfg.BaseURL == "" { cfg.BaseURL = "https://api.anthropic.com/" } - // Resolve centralized key configuration into KeyPool. - // Precedence: - // 1. cfg.KeyPool (explicit, highest priority). - // 2. cfg.Key (legacy single key). - // After this block cfg.Key is empty so it can only carry a - // BYOK X-Api-Key set per interception in CreateInterceptor. - // TODO(ssncferreira): simplify auth field resolution per - // https://github.com/coder/aibridge/issues/266. - if cfg.KeyPool == nil && cfg.Key != "" { - // keypool.New only fails on empty or duplicate keys, - // neither possible with a single non-empty key. - pool, err := keypool.New(cfg.Name, []string{cfg.Key}, quartz.NewReal(), nil) - if err != nil { - panic(fmt.Sprintf("anthropic provider: build single-key pool: %s", err)) - } - cfg.KeyPool = pool - } - cfg.Key = "" if cfg.CircuitBreaker != nil { cfg.CircuitBreaker.IsFailure = anthropicIsFailure cfg.CircuitBreaker.OpenErrorResponse = anthropicOpenErrorResponse @@ -135,63 +109,62 @@ func (p *Anthropic) CreateInterceptor(_ http.ResponseWriter, r *http.Request, tr return nil, xerrors.Errorf("unmarshal request body: %w", err) } - cfg := p.cfg - cfg.ExtraHeaders = extractAnthropicHeaders(r) - - // At this point the request contains only LLM provider headers. - // Any Coder-specific authentication has already been stripped. - // - // In centralized mode neither Authorization nor X-Api-Key is - // present, so cfg keeps the KeyPool from provider construction - // and the failover loop walks it. - // - // In BYOK mode the user's LLM credentials survive intact and - // failover is disabled by clearing cfg.KeyPool. If X-Api-Key is - // present the user has a personal API key, populate cfg.Key. - // If Authorization is present the user authenticated directly - // with the provider, populate cfg.BYOKBearerToken. When both - // are present, X-Api-Key takes priority to match claude-code - // behavior. - // - // TODO(ssncferreira): consolidate auth field handling per - // https://github.com/coder/aibridge/issues/266. - credKind := intercept.CredentialKindCentralized - var credSecret string - authHeaderName := p.AuthHeader() - if apiKey := r.Header.Get("X-Api-Key"); apiKey != "" { - cfg.Key = apiKey - cfg.KeyPool = nil - authHeaderName = "X-Api-Key" - credKind = intercept.CredentialKindBYOK - credSecret = apiKey - } else if token := utils.ExtractBearerToken(r.Header.Get("Authorization")); token != "" { - cfg.BYOKBearerToken = token - cfg.KeyPool = nil - authHeaderName = "Authorization" - credKind = intercept.CredentialKindBYOK - credSecret = token + cfg := intercept.Config{ + ProviderName: p.Name(), + BaseURL: p.cfg.BaseURL, + APIDumpDir: p.cfg.APIDumpDir, + SendActorHeaders: p.cfg.SendActorHeaders, + } + cred, err := p.resolveCredential(r) + if err != nil { + span.SetStatus(codes.Error, err.Error()) + return nil, xerrors.Errorf("resolve credential: %w", err) } - // Centralized leaves credSecret empty: the hint is set by the - // failover loop on each key attempt and persisted at - // end-of-interception. - cred := intercept.NewCredentialInfo(credKind, credSecret) var interceptor intercept.Interceptor if reqPayload.Stream() { - interceptor = messages.NewStreamingInterceptor(id, reqPayload, p.Name(), cfg, p.bedrockCfg, r.Header, authHeaderName, tracer, cred) + interceptor = messages.NewStreamingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, r.Header, tracer) } else { - interceptor = messages.NewBlockingInterceptor(id, reqPayload, p.Name(), cfg, p.bedrockCfg, r.Header, authHeaderName, tracer, cred) + interceptor = messages.NewBlockingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, r.Header, tracer) } span.SetAttributes(interceptor.TraceAttributes(r)...) return interceptor, nil } +// resolveCredential determines the upstream credential for a request. At this +// point the request contains only LLM provider headers. Any Coder-specific +// authentication has already been stripped. +// +// - X-Api-Key present: BYOK with a personal API key. +// - Authorization present: BYOK with an access token. +// - Neither present: centralized, using the provider's key pool with +// failover. +// +// When both BYOK headers are present, X-Api-Key takes priority to match +// claude-code behavior. Centralized requests require a key pool, except for +// Bedrock providers, which authenticate via AWS signing rather than a pool. +func (p *Anthropic) resolveCredential(r *http.Request) (intercept.Credential, error) { + if apiKey := r.Header.Get(intercept.AuthHeaderXAPIKey); apiKey != "" { + return intercept.BYOK{Secret: apiKey, Header: intercept.AuthHeaderXAPIKey}, nil + } + if token := utils.ExtractBearerToken(r.Header.Get(intercept.AuthHeaderAuthorization)); token != "" { + return intercept.BYOK{Secret: token, Header: intercept.AuthHeaderAuthorization}, nil + } + if p.cfg.KeyPool != nil { + return &intercept.CentralizedPool{Pool: p.cfg.KeyPool, Header: p.AuthHeader()}, nil + } + if p.bedrockCfg != nil { + return intercept.Bedrock{AccessKey: p.bedrockCfg.AccessKey}, nil + } + return nil, ErrNoCredential +} + func (p *Anthropic) BaseURL() string { return p.cfg.BaseURL } func (*Anthropic) AuthHeader() string { - return "X-Api-Key" + return intercept.AuthHeaderXAPIKey } func (p *Anthropic) KeyPool() *keypool.Pool { @@ -203,10 +176,10 @@ func (p *Anthropic) KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverCon Pool: p.cfg.KeyPool, Logger: logger, IsBYOK: func(r *http.Request) bool { - return r.Header.Get("X-Api-Key") != "" || r.Header.Get("Authorization") != "" + return r.Header.Get(intercept.AuthHeaderXAPIKey) != "" || r.Header.Get(intercept.AuthHeaderAuthorization) != "" }, InjectAuthKey: func(h *http.Header, key string) { - h.Set("X-Api-Key", key) + h.Set(intercept.AuthHeaderXAPIKey, key) }, BuildKeyPoolResponse: func(keyPoolErr *keypool.Error) *http.Response { return messages.ResponseErrorFromKeyPool(keyPoolErr).ToResponse() @@ -221,16 +194,3 @@ func (p *Anthropic) CircuitBreakerConfig() *config.CircuitBreaker { func (p *Anthropic) APIDumpDir() string { return p.cfg.APIDumpDir } - -// extractAnthropicHeaders extracts headers required by the Anthropic API from -// the incoming request. -// TODO(ssncferreira): remove as part of https://github.com/coder/aibridge/issues/192 -func extractAnthropicHeaders(r *http.Request) map[string]string { - headers := make(map[string]string, len(anthropicForwardHeaders)) - for _, h := range anthropicForwardHeaders { - if v := r.Header.Get(h); v != "" { - headers[h] = v - } - } - return headers -} diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 285fa3cd04..679469fdc6 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -63,13 +63,6 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { cfg config.Anthropic expectedKeys []string }{ - { - // Legacy single-key path: NewAnthropic builds a - // pool containing just that key. - name: "key_creates_keypool", - cfg: config.Anthropic{Key: "legacy-key"}, - expectedKeys: []string{"legacy-key"}, - }, { // Caller supplies the pool directly. name: "keypool_passed_directly", @@ -77,15 +70,9 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { expectedKeys: []string{"pool-key-0", "pool-key-1"}, }, { - // Both set: KeyPool wins, Key is ignored. - name: "keypool_takes_precedence_over_key", - cfg: config.Anthropic{Key: "legacy-key", KeyPool: pool}, - expectedKeys: []string{"pool-key-0", "pool-key-1"}, - }, - { - // Neither set: no centralized auth available. BYOK - // auth is set per-request in CreateInterceptor. - name: "neither_set_no_centralized_auth", + // No pool: no centralized auth available. BYOK auth is + // resolved per-request in CreateInterceptor. + name: "no_keypool_no_centralized_auth", cfg: config.Anthropic{}, expectedKeys: nil, }, @@ -119,7 +106,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { func TestAnthropic_CreateInterceptor(t *testing.T) { t.Parallel() - provider := NewAnthropic(config.Anthropic{Key: "test-key"}, nil) + provider := NewAnthropic(config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) t.Run("Messages_NonStreamingRequest_BlockingInterceptor", func(t *testing.T) { t.Parallel() @@ -179,7 +166,7 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { provider := NewAnthropic(config.Anthropic{ BaseURL: mockUpstream.URL, - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), }, nil) // Use a realistic multi-beta value as sent by Claude Code clients. @@ -227,49 +214,95 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { }) } -func TestAnthropic_CreateInterceptor_BYOK(t *testing.T) { +func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { t.Parallel() tests := []struct { - name string - setHeaders map[string]string - wantXApiKey string - wantAuthorization string + name string + pool bool // provider has a centralized "test-key" pool + bedrock bool // Bedrock-backed provider (authenticates via AWS signing) + // bedrockStatic, when bedrock is set, configures static AWS credentials. + // False means dynamic mode (AWS default credential chain). + bedrockStatic bool + setHeaders map[string]string + // wantErr, when set, means CreateInterceptor must fail with it. The + // remaining expectations are then ignored. + wantErr error wantCredentialKind intercept.CredentialKind wantCredentialHint string + // Upstream expectations after ProcessRequest. Not checked for Bedrock, + // which signs via AWS rather than forwarding a key header. + wantXApiKey string + wantAuthorization string }{ { - name: "Messages_BYOK_BearerToken", + name: "byok_bearer_token", + pool: true, setHeaders: map[string]string{"Authorization": "Bearer user-access-token"}, - wantAuthorization: "Bearer user-access-token", wantCredentialKind: intercept.CredentialKindBYOK, wantCredentialHint: "us...en", + wantAuthorization: "Bearer user-access-token", }, { - name: "Messages_BYOK_APIKey", + name: "byok_api_key", + pool: true, setHeaders: map[string]string{"X-Api-Key": "user-api-key"}, - wantXApiKey: "user-api-key", wantCredentialKind: intercept.CredentialKindBYOK, wantCredentialHint: "us...ey", + wantXApiKey: "user-api-key", }, { - name: "Messages_Centralized", + name: "byok_bearer_and_api_key", + pool: true, + setHeaders: map[string]string{"Authorization": "Bearer user-access-token", "X-Api-Key": "user-api-key"}, + // X-Api-Key takes priority over Authorization. + wantCredentialKind: intercept.CredentialKindBYOK, + wantCredentialHint: "us...ey", + wantXApiKey: "user-api-key", + }, + { + name: "byok_without_pool", + pool: false, + setHeaders: map[string]string{"X-Api-Key": "user-api-key"}, + wantCredentialKind: intercept.CredentialKindBYOK, + wantCredentialHint: "us...ey", + wantXApiKey: "user-api-key", + }, + { + name: "centralized", + pool: true, setHeaders: map[string]string{}, - wantXApiKey: "test-key", wantCredentialKind: intercept.CredentialKindCentralized, - // Centralized hint is empty at CreateInterceptor; set - // by the key failover loop during ProcessRequest. - wantCredentialHint: "", + // The pool hasn't handed out a key at CreateInterceptor, so the hint + // is a placeholder until the failover loop selects one. + wantCredentialHint: "", + wantXApiKey: "test-key", }, { - name: "Messages_BYOK_BearerToken_And_APIKey", - setHeaders: map[string]string{ - "Authorization": "Bearer user-access-token", - "X-Api-Key": "user-api-key", - }, - wantXApiKey: "user-api-key", - wantCredentialKind: intercept.CredentialKindBYOK, - wantCredentialHint: "us...ey", + // Bedrock dynamic mode: no static access key, so the hint is the + // AWS-credential-chain placeholder. + name: "bedrock_dynamic", + pool: false, + bedrock: true, + setHeaders: map[string]string{}, + wantCredentialKind: intercept.CredentialKindCentralized, + wantCredentialHint: "", + }, + { + // Bedrock static mode: the hint masks the access key ID. + name: "bedrock_static", + pool: false, + bedrock: true, + bedrockStatic: true, + setHeaders: map[string]string{}, + wantCredentialKind: intercept.CredentialKindCentralized, + wantCredentialHint: "AKIA...MPLE", + }, + { + name: "centralized_without_pool_errors", + pool: false, + setHeaders: map[string]string{}, + wantErr: ErrNoCredential, }, } @@ -278,7 +311,6 @@ func TestAnthropic_CreateInterceptor_BYOK(t *testing.T) { t.Parallel() var receivedHeaders http.Header - mockUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { receivedHeaders = r.Header.Clone() w.Header().Set("Content-Type", "application/json") @@ -287,10 +319,19 @@ func TestAnthropic_CreateInterceptor_BYOK(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := NewAnthropic(config.Anthropic{ - BaseURL: mockUpstream.URL, - Key: "test-key", - }, nil) + acfg := config.Anthropic{BaseURL: mockUpstream.URL} + if tc.pool { + acfg.KeyPool = testutil.SingleKeyPool(config.ProviderAnthropic, "test-key") + } + var bedrock *config.AWSBedrock + if tc.bedrock { + bedrock = &config.AWSBedrock{Region: "us-west-2", Model: "m", SmallFastModel: "s"} + if tc.bedrockStatic { + bedrock.AccessKey = "AKIAIOSFODNN7EXAMPLE" + bedrock.AccessKeySecret = "wJalrXUtnFEMI-secret-value" + } + } + provider := NewAnthropic(acfg, bedrock) body := `{"model": "claude-opus-4-5", "max_tokens": 1024, "messages": [{"role": "user", "content": "hello"}], "stream": false}` req := httptest.NewRequest(http.MethodPost, routeMessages, bytes.NewBufferString(body)) @@ -300,19 +341,27 @@ func TestAnthropic_CreateInterceptor_BYOK(t *testing.T) { w := httptest.NewRecorder() interceptor, err := provider.CreateInterceptor(w, req, testTracer) + if tc.wantErr != nil { + require.ErrorIs(t, err, tc.wantErr) + require.Nil(t, interceptor) + return + } require.NoError(t, err) require.NotNil(t, interceptor) cred := interceptor.Credential() - assert.Equal(t, tc.wantCredentialKind, cred.Kind, "credential kind mismatch") - assert.Equal(t, tc.wantCredentialHint, cred.Hint, "credential hint mismatch") + assert.Equal(t, tc.wantCredentialKind, cred.Kind(), "credential kind mismatch") + assert.Equal(t, tc.wantCredentialHint, cred.Hint(), "credential hint mismatch") - logger := slog.Make() - interceptor.Setup(logger, &testutil.MockRecorder{}, nil) + // Bedrock signs via AWS during ProcessRequest (needs real AWS + // credentials), covered by the integration tests. + if tc.bedrock { + return + } + interceptor.Setup(slog.Make(), &testutil.MockRecorder{}, nil) processReq := httptest.NewRequest(http.MethodPost, routeMessages, nil) - err = interceptor.ProcessRequest(w, processReq) - require.NoError(t, err) + require.NoError(t, interceptor.ProcessRequest(w, processReq)) assert.Equal(t, tc.wantXApiKey, receivedHeaders.Get("X-Api-Key")) assert.Equal(t, tc.wantAuthorization, receivedHeaders.Get("Authorization")) @@ -454,51 +503,6 @@ func TestAnthropic_KeyFailoverConfig(t *testing.T) { }) } -func TestExtractAnthropicHeaders(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - headers map[string]string - expected map[string]string - }{ - { - name: "no headers", - headers: map[string]string{}, - expected: map[string]string{}, - }, - { - name: "single beta", - headers: map[string]string{"Anthropic-Beta": "claude-code-20250219"}, - expected: map[string]string{"Anthropic-Beta": "claude-code-20250219"}, - }, - { - name: "multiple betas in single header", - headers: map[string]string{"Anthropic-Beta": "claude-code-20250219,adaptive-thinking-2026-01-28,context-management-2025-06-27,prompt-caching-scope-2026-01-05,effort-2025-11-24"}, - expected: map[string]string{"Anthropic-Beta": "claude-code-20250219,adaptive-thinking-2026-01-28,context-management-2025-06-27,prompt-caching-scope-2026-01-05,effort-2025-11-24"}, - }, - { - name: "ignores other headers", - headers: map[string]string{"Anthropic-Beta": "claude-code-20250219,context-management-2025-06-27", "X-Api-Key": "secret"}, - expected: map[string]string{"Anthropic-Beta": "claude-code-20250219,context-management-2025-06-27"}, - }, - } - - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - - req := httptest.NewRequest(http.MethodPost, "/", nil) - for header, value := range tc.headers { - req.Header.Set(header, value) - } - - result := extractAnthropicHeaders(req) - assert.Equal(t, tc.expected, result) - }) - } -} - func Test_anthropicIsFailure(t *testing.T) { t.Parallel() diff --git a/aibridge/provider/copilot.go b/aibridge/provider/copilot.go index 7453515b92..c3b4f07ea7 100644 --- a/aibridge/provider/copilot.go +++ b/aibridge/provider/copilot.go @@ -34,16 +34,6 @@ var copilotOpenErrorResponse = func() []byte { return []byte(`{"error":{"message":"circuit breaker is open","type":"server_error","code":"service_unavailable"}}`) } -// Headers that need to be forwarded to Copilot API. -// These were determined through manual testing as there is no reference -// of the headers in the official documentation. -// LiteLLM uses the same headers: -// https://docs.litellm.ai/docs/providers/github_copilot -var copilotForwardHeaders = []string{ - "Editor-Version", - "Copilot-Integration-Id", -} - // Copilot implements the Provider interface for GitHub Copilot. // Unlike other providers, Copilot uses per-user API keys that are passed through // the request headers rather than configured statically. @@ -133,7 +123,7 @@ func (p *Copilot) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trac defer tracing.EndSpanErr(span, &outErr) // Extract the per-user Copilot key from the Authorization header. - key := utils.ExtractBearerToken(r.Header.Get("Authorization")) + key := utils.ExtractBearerToken(r.Header.Get(intercept.AuthHeaderAuthorization)) if key == "" { span.SetStatus(codes.Error, "missing authorization") return nil, xerrors.New("missing Copilot authorization: Authorization header not found or invalid") @@ -141,18 +131,14 @@ func (p *Copilot) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trac id := uuid.New() - // Build config for the interceptor using the per-request key. - // Copilot's API is OpenAI-compatible, so it uses the OpenAI interceptors - // that require a config.OpenAI. - cfg := config.OpenAI{ - BaseURL: p.cfg.BaseURL, - Key: key, - APIDumpDir: p.cfg.APIDumpDir, - CircuitBreaker: p.cfg.CircuitBreaker, - ExtraHeaders: extractCopilotHeaders(r), + // Copilot's API is OpenAI-compatible, so it reuses the OpenAI interceptors. + // It is always BYOK: the per-user key arrives in the Authorization header. + cfg := intercept.Config{ + ProviderName: p.Name(), + BaseURL: p.cfg.BaseURL, + APIDumpDir: p.cfg.APIDumpDir, } - - cred := intercept.NewCredentialInfo(intercept.CredentialKindBYOK, key) + cred := intercept.BYOK{Secret: key, Header: intercept.AuthHeaderAuthorization} var interceptor intercept.Interceptor @@ -165,9 +151,9 @@ func (p *Copilot) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trac } if req.Stream { - interceptor = chatcompletions.NewStreamingInterceptor(id, &req, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = chatcompletions.NewStreamingInterceptor(id, &req, cfg, cred, r.Header, tracer) } else { - interceptor = chatcompletions.NewBlockingInterceptor(id, &req, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = chatcompletions.NewBlockingInterceptor(id, &req, cfg, cred, r.Header, tracer) } case routeCopilotResponses: @@ -181,9 +167,9 @@ func (p *Copilot) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trac } if reqPayload.Stream() { - interceptor = responses.NewStreamingInterceptor(id, reqPayload, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = responses.NewStreamingInterceptor(id, reqPayload, cfg, cred, r.Header, tracer) } else { - interceptor = responses.NewBlockingInterceptor(id, reqPayload, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = responses.NewBlockingInterceptor(id, reqPayload, cfg, cred, r.Header, tracer) } default: @@ -194,15 +180,3 @@ func (p *Copilot) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trac span.SetAttributes(interceptor.TraceAttributes(r)...) return interceptor, nil } - -// extractCopilotHeaders extracts headers required by the Copilot API from the -// incoming request. Copilot requires certain client headers to be forwarded. -func extractCopilotHeaders(r *http.Request) map[string]string { - headers := make(map[string]string, len(copilotForwardHeaders)) - for _, h := range copilotForwardHeaders { - if v := r.Header.Get(h); v != "" { - headers[h] = v - } - } - return headers -} diff --git a/aibridge/provider/copilot_internal_test.go b/aibridge/provider/copilot_internal_test.go index 49cb582b78..c29afec72c 100644 --- a/aibridge/provider/copilot_internal_test.go +++ b/aibridge/provider/copilot_internal_test.go @@ -295,48 +295,3 @@ func TestCopilot_CreateInterceptor(t *testing.T) { require.Nil(t, interceptor) }) } - -func TestExtractCopilotHeaders(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - headers map[string]string - expected map[string]string - }{ - { - name: "all headers present", - headers: map[string]string{"Editor-Version": "vscode/1.85.0", "Copilot-Integration-Id": "some-id"}, - expected: map[string]string{"Editor-Version": "vscode/1.85.0", "Copilot-Integration-Id": "some-id"}, - }, - { - name: "some headers present", - headers: map[string]string{"Editor-Version": "vscode/1.85.0"}, - expected: map[string]string{"Editor-Version": "vscode/1.85.0"}, - }, - { - name: "no headers", - headers: map[string]string{}, - expected: map[string]string{}, - }, - { - name: "ignores other headers", - headers: map[string]string{"Editor-Version": "vscode/1.85.0", "Authorization": "Bearer token"}, - expected: map[string]string{"Editor-Version": "vscode/1.85.0"}, - }, - } - - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - - req := httptest.NewRequest(http.MethodPost, "/", nil) - for header, value := range tc.headers { - req.Header.Set(header, value) - } - - result := extractCopilotHeaders(req) - assert.Equal(t, tc.expected, result) - }) - } -} diff --git a/aibridge/provider/openai.go b/aibridge/provider/openai.go index 13763615e6..7ea4fe922e 100644 --- a/aibridge/provider/openai.go +++ b/aibridge/provider/openai.go @@ -20,7 +20,6 @@ import ( "github.com/coder/coder/v2/aibridge/keypool" "github.com/coder/coder/v2/aibridge/tracing" "github.com/coder/coder/v2/aibridge/utils" - "github.com/coder/quartz" ) const ( @@ -47,25 +46,6 @@ func NewOpenAI(cfg config.OpenAI) *OpenAI { if cfg.BaseURL == "" { cfg.BaseURL = "https://api.openai.com/v1/" } - // Resolve centralized key configuration into KeyPool. - // Precedence: - // 1. cfg.KeyPool (explicit, highest priority). - // 2. cfg.Key (legacy single key). - // After this block cfg.Key is empty so it can only carry a - // BYOK Authorization Bearer set per interception in - // CreateInterceptor. - // TODO(ssncferreira): simplify auth field resolution per - // https://github.com/coder/aibridge/issues/266. - if cfg.KeyPool == nil && cfg.Key != "" { - // keypool.New only fails on empty or duplicate keys, - // neither possible with a single non-empty key. - pool, err := keypool.New(cfg.Name, []string{cfg.Key}, quartz.NewReal(), nil) - if err != nil { - panic(fmt.Sprintf("openai provider: build single-key pool: %s", err)) - } - cfg.KeyPool = pool - } - cfg.Key = "" if cfg.CircuitBreaker != nil { cfg.CircuitBreaker.OpenErrorResponse = openAIOpenErrorResponse } @@ -123,31 +103,17 @@ func (p *OpenAI) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trace var interceptor intercept.Interceptor - cfg := p.cfg - // At this point the request contains only LLM provider headers. - // Any Coder-specific authentication has already been stripped. - // - // In centralized mode Authorization is absent, so cfg keeps the - // KeyPool from provider construction and the failover loop walks - // it. - // - // In BYOK mode the user's credential is in Authorization, - // populate cfg.Key and clear cfg.KeyPool so failover is disabled. - // - // TODO(ssncferreira): consolidate auth field handling per - // https://github.com/coder/aibridge/issues/266. - credKind := intercept.CredentialKindCentralized - var credSecret string - if token := utils.ExtractBearerToken(r.Header.Get("Authorization")); token != "" { - cfg.Key = token - cfg.KeyPool = nil - credKind = intercept.CredentialKindBYOK - credSecret = token + cfg := intercept.Config{ + ProviderName: p.Name(), + BaseURL: p.cfg.BaseURL, + APIDumpDir: p.cfg.APIDumpDir, + SendActorHeaders: p.cfg.SendActorHeaders, + } + cred, err := p.resolveCredential(r) + if err != nil { + span.SetStatus(codes.Error, err.Error()) + return nil, xerrors.Errorf("resolve credential: %w", err) } - // Centralized leaves credSecret empty: the hint is set by the - // failover loop on each key attempt and persisted at - // end-of-interception. - cred := intercept.NewCredentialInfo(credKind, credSecret) path := strings.TrimPrefix(r.URL.Path, p.RoutePrefix()) switch path { @@ -158,9 +124,9 @@ func (p *OpenAI) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trace } if req.Stream { - interceptor = chatcompletions.NewStreamingInterceptor(id, &req, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = chatcompletions.NewStreamingInterceptor(id, &req, cfg, cred, r.Header, tracer) } else { - interceptor = chatcompletions.NewBlockingInterceptor(id, &req, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = chatcompletions.NewBlockingInterceptor(id, &req, cfg, cred, r.Header, tracer) } case routeResponses: @@ -173,9 +139,9 @@ func (p *OpenAI) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trace return nil, xerrors.Errorf("unmarshal request body: %w", err) } if reqPayload.Stream() { - interceptor = responses.NewStreamingInterceptor(id, reqPayload, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = responses.NewStreamingInterceptor(id, reqPayload, cfg, cred, r.Header, tracer) } else { - interceptor = responses.NewBlockingInterceptor(id, reqPayload, p.Name(), cfg, r.Header, p.AuthHeader(), tracer, cred) + interceptor = responses.NewBlockingInterceptor(id, reqPayload, cfg, cred, r.Header, tracer) } default: @@ -186,6 +152,21 @@ func (p *OpenAI) CreateInterceptor(_ http.ResponseWriter, r *http.Request, trace return interceptor, nil } +// resolveCredential determines the upstream credential for a request. At this +// point the request contains only LLM provider headers. Any Coder-specific +// authentication has already been stripped. A BYOK token, if present, arrives +// in the Authorization header. Otherwise the request uses the provider's +// centralized key pool with failover, which must be configured. +func (p *OpenAI) resolveCredential(r *http.Request) (intercept.Credential, error) { + if token := utils.ExtractBearerToken(r.Header.Get(intercept.AuthHeaderAuthorization)); token != "" { + return intercept.BYOK{Secret: token, Header: intercept.AuthHeaderAuthorization}, nil + } + if p.cfg.KeyPool == nil { + return nil, ErrNoCredential + } + return &intercept.CentralizedPool{Pool: p.cfg.KeyPool, Header: p.AuthHeader()}, nil +} + func (p *OpenAI) BaseURL() string { return p.cfg.BaseURL } diff --git a/aibridge/provider/openai_internal_test.go b/aibridge/provider/openai_internal_test.go index 6ce11ca221..8f3225b13e 100644 --- a/aibridge/provider/openai_internal_test.go +++ b/aibridge/provider/openai_internal_test.go @@ -198,15 +198,19 @@ func TestOpenAI_TypeAndName(t *testing.T) { } } -func TestOpenAI_CreateInterceptor(t *testing.T) { +func TestOpenAI_CreateInterceptor_Credential(t *testing.T) { t.Parallel() tests := []struct { - name string - route string - requestBody string - responseBody string - setHeaders map[string]string + name string + route string + requestBody string + responseBody string + pool bool // provider has a centralized "centralized-key" pool + setHeaders map[string]string + // wantErr, when set, means CreateInterceptor must fail with it. The + // remaining expectations are then ignored. + wantErr error wantAuthorization string wantCredentialKind intercept.CredentialKind wantCredentialHint string @@ -216,6 +220,7 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { route: routeChatCompletions, requestBody: `{"model": "gpt-4", "messages": [{"role": "user", "content": "hello"}], "stream": false}`, responseBody: chatCompletionResponse, + pool: true, setHeaders: map[string]string{"Authorization": "Bearer user-token"}, wantAuthorization: "Bearer user-token", wantCredentialKind: intercept.CredentialKindBYOK, @@ -226,18 +231,20 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { route: routeChatCompletions, requestBody: `{"model": "gpt-4", "messages": [{"role": "user", "content": "hello"}], "stream": false}`, responseBody: chatCompletionResponse, + pool: true, setHeaders: map[string]string{}, wantAuthorization: "Bearer centralized-key", wantCredentialKind: intercept.CredentialKindCentralized, - // Centralized hint is empty at CreateInterceptor; set - // by the key failover loop during ProcessRequest. - wantCredentialHint: "", + // The pool hasn't handed out a key at CreateInterceptor, so the + // hint is a placeholder until the failover loop selects one. + wantCredentialHint: "", }, { name: "Responses_BYOK", route: routeResponses, requestBody: `{"model": "gpt-5", "input": "hello", "stream": false}`, responseBody: responsesAPIResponse, + pool: true, setHeaders: map[string]string{"Authorization": "Bearer user-token"}, wantAuthorization: "Bearer user-token", wantCredentialKind: intercept.CredentialKindBYOK, @@ -248,12 +255,13 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { route: routeResponses, requestBody: `{"model": "gpt-5", "input": "hello", "stream": false}`, responseBody: responsesAPIResponse, + pool: true, setHeaders: map[string]string{}, wantAuthorization: "Bearer centralized-key", wantCredentialKind: intercept.CredentialKindCentralized, - // Centralized hint is empty at CreateInterceptor; set - // by the key failover loop during ProcessRequest. - wantCredentialHint: "", + // The pool hasn't handed out a key at CreateInterceptor, so the + // hint is a placeholder until the failover loop selects one. + wantCredentialHint: "", }, // X-Api-Key should not appear in production since clients use Authorization, // but ensure it is stripped if it does arrive. @@ -262,6 +270,7 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { route: routeChatCompletions, requestBody: `{"model": "gpt-4", "messages": [{"role": "user", "content": "hello"}], "stream": false}`, responseBody: chatCompletionResponse, + pool: true, setHeaders: map[string]string{ "Authorization": "Bearer user-token", "X-Api-Key": "some-key", @@ -275,6 +284,7 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { route: routeResponses, requestBody: `{"model": "gpt-5", "input": "hello", "stream": false}`, responseBody: responsesAPIResponse, + pool: true, setHeaders: map[string]string{ "Authorization": "Bearer user-token", "X-Api-Key": "some-key", @@ -283,6 +293,27 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { wantCredentialKind: intercept.CredentialKindBYOK, wantCredentialHint: "us...en", }, + { + // BYOK authenticates even with no centralized pool. + name: "ChatCompletions_BYOK_WithoutPool", + route: routeChatCompletions, + requestBody: `{"model": "gpt-4", "messages": [{"role": "user", "content": "hello"}], "stream": false}`, + responseBody: chatCompletionResponse, + pool: false, + setHeaders: map[string]string{"Authorization": "Bearer user-token"}, + wantAuthorization: "Bearer user-token", + wantCredentialKind: intercept.CredentialKindBYOK, + wantCredentialHint: "us...en", + }, + { + // No centralized keys and no Authorization: cannot authenticate. + name: "ChatCompletions_NoCredential", + route: routeChatCompletions, + requestBody: `{"model": "gpt-4", "messages": [{"role": "user", "content": "hello"}], "stream": false}`, + pool: false, + setHeaders: map[string]string{}, + wantErr: ErrNoCredential, + }, } for _, tc := range tests { @@ -300,10 +331,11 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := NewOpenAI(config.OpenAI{ - BaseURL: mockUpstream.URL, - Key: "centralized-key", - }) + ocfg := config.OpenAI{BaseURL: mockUpstream.URL} + if tc.pool { + ocfg.KeyPool = testutil.SingleKeyPool(config.ProviderOpenAI, "centralized-key") + } + provider := NewOpenAI(ocfg) req := httptest.NewRequest(http.MethodPost, provider.RoutePrefix()+tc.route, bytes.NewBufferString(tc.requestBody)) for k, v := range tc.setHeaders { @@ -312,19 +344,22 @@ func TestOpenAI_CreateInterceptor(t *testing.T) { w := httptest.NewRecorder() interceptor, err := provider.CreateInterceptor(w, req, testTracer) + if tc.wantErr != nil { + require.ErrorIs(t, err, tc.wantErr) + require.Nil(t, interceptor) + return + } require.NoError(t, err) require.NotNil(t, interceptor) cred := interceptor.Credential() - assert.Equal(t, tc.wantCredentialKind, cred.Kind, "credential kind mismatch") - assert.Equal(t, tc.wantCredentialHint, cred.Hint, "credential hint mismatch") + assert.Equal(t, tc.wantCredentialKind, cred.Kind(), "credential kind mismatch") + assert.Equal(t, tc.wantCredentialHint, cred.Hint(), "credential hint mismatch") - logger := slog.Make() - interceptor.Setup(logger, &testutil.MockRecorder{}, nil) + interceptor.Setup(slog.Make(), &testutil.MockRecorder{}, nil) processReq := httptest.NewRequest(http.MethodPost, provider.RoutePrefix()+tc.route, nil) - err = interceptor.ProcessRequest(w, processReq) - require.NoError(t, err) + require.NoError(t, interceptor.ProcessRequest(w, processReq)) assert.Equal(t, tc.wantAuthorization, receivedHeaders.Get("Authorization")) assert.Empty(t, receivedHeaders.Get("X-Api-Key"), "X-Api-Key must not be set upstream") @@ -463,7 +498,7 @@ func TestOpenAI_KeyFailoverConfig(t *testing.T) { func BenchmarkOpenAI_CreateInterceptor_ChatCompletions(b *testing.B) { provider := NewOpenAI(config.OpenAI{ BaseURL: "https://api.openai.com/v1/", - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderOpenAI, "test-key"), }) tracer := noop.NewTracerProvider().Tracer("test") @@ -501,7 +536,7 @@ func BenchmarkOpenAI_CreateInterceptor_ChatCompletions(b *testing.B) { func BenchmarkOpenAI_CreateInterceptor_Responses(b *testing.B) { provider := NewOpenAI(config.OpenAI{ BaseURL: "https://api.openai.com/v1/", - Key: "test-key", + KeyPool: testutil.SingleKeyPool(config.ProviderOpenAI, "test-key"), }) tracer := noop.NewTracerProvider().Tracer("test") diff --git a/aibridge/provider/provider.go b/aibridge/provider/provider.go index 310b6f6fcf..22b475a8d2 100644 --- a/aibridge/provider/provider.go +++ b/aibridge/provider/provider.go @@ -14,6 +14,11 @@ import ( var ErrUnknownRoute = xerrors.New("unknown route") +// ErrNoCredential is returned when a request resolves to centralized +// authentication but the provider has no centralized keys configured (and the +// request is not BYOK), so it cannot be authenticated. +var ErrNoCredential = xerrors.New("no credential: request is not BYOK and the provider has no centralized keys") + // Provider defines routes (bridged and passed through) for given provider. // Bridged routes are processed by dedicated interceptors. // diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index 8b29b2653a..c0592e7390 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -18,14 +18,24 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/aibridge" "github.com/coder/coder/v2/aibridge/intercept" + "github.com/coder/coder/v2/aibridge/keypool" agplaibridge "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/coderd/aibridged" mock "github.com/coder/coder/v2/coderd/aibridged/aibridgedmock" "github.com/coder/coder/v2/coderd/aibridged/proto" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/testutil" + "github.com/coder/quartz" ) +// singleKeyPool builds a centralized key pool containing a single key. +func singleKeyPool(t *testing.T, name, key string) *keypool.Pool { + t.Helper() + pool, err := keypool.New(name, []string{key}, quartz.NewReal(), nil) + require.NoError(t, err) + return pool +} + func newTestServer(t *testing.T) (*aibridged.Server, *mock.MockDRPCClient, *mock.MockPooler) { t.Helper() @@ -646,10 +656,12 @@ func TestServeHTTP_ActorHeaders(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{ BaseURL: upstreamSrv.URL, + KeyPool: singleKeyPool(t, "openai", "test-key"), SendActorHeaders: true, }), aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ BaseURL: upstreamSrv.URL, + KeyPool: singleKeyPool(t, "anthropic", "test-key"), SendActorHeaders: true, }, nil), } @@ -753,8 +765,8 @@ func TestRouting(t *testing.T) { client := mock.NewMockDRPCClient(ctrl) providers := []aibridge.Provider{ - aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL}), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL}, nil), + aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL, KeyPool: singleKeyPool(t, "openai", "test-key")}), + aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key")}, nil), } pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, nil, testTracer) require.NoError(t, err) diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 5d907f0726..791c7323a5 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -21,6 +21,7 @@ import ( "github.com/coder/coder/v2/aibridge" "github.com/coder/coder/v2/aibridge/config" + "github.com/coder/coder/v2/aibridge/keypool" aibtracing "github.com/coder/coder/v2/aibridge/tracing" "github.com/coder/coder/v2/coderd/aibridged" "github.com/coder/coder/v2/coderd/aibridgedserver" @@ -33,8 +34,17 @@ import ( "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" "github.com/coder/coder/v2/testutil" + "github.com/coder/quartz" ) +// singleKeyPool builds a centralized key pool containing a single key. +func singleKeyPool(t *testing.T, name, key string) *keypool.Pool { + t.Helper() + pool, err := keypool.New(name, []string{key}, quartz.NewReal(), nil) + require.NoError(t, err) + return pool +} + var testTracer = otel.Tracer("aibridged_inttest") // TestIntegration is not an exhaustive test against the upstream AI providers' SDKs (see coder/aibridge for those). @@ -183,7 +193,7 @@ func TestIntegration(t *testing.T) { require.NoError(t, err) logger := testutil.Logger(t) - providers := []aibridge.Provider{aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: mockOpenAI.URL, Key: "test-centralized-key"})} + providers := []aibridge.Provider{aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: mockOpenAI.URL, KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-centralized-key")})} pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, nil, tracer) require.NoError(t, err) @@ -383,7 +393,7 @@ func TestIntegrationWithMetrics(t *testing.T) { require.NoError(t, err) logger := testutil.Logger(t) - providers := []aibridge.Provider{aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: mockOpenAI.URL})} + providers := []aibridge.Provider{aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: mockOpenAI.URL, KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-centralized-key")})} // Create pool with metrics. pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, metrics, testTracer) @@ -491,11 +501,12 @@ func TestIntegrationCircuitBreaker(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{ BaseURL: mockOpenAI.URL, + KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }), aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ BaseURL: mockAnthropic.URL, - Key: "test-key", + KeyPool: singleKeyPool(t, config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, }, nil), }