From dbf927d4a4600797b5de95428ff2bbf228ce388b Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Thu, 18 Jun 2026 17:51:49 +0000 Subject: [PATCH 01/21] feat: support cross-account Bedrock AssumeRole in AI Bridge --- aibridge/api.go | 4 +- aibridge/bridge_test.go | 22 +- aibridge/config/config.go | 11 + aibridge/intercept/keyfailover_test.go | 4 +- aibridge/intercept/messages/base.go | 59 ++--- .../intercept/messages/base_internal_test.go | 123 ++------- aibridge/intercept/messages/blocking.go | 3 + aibridge/intercept/messages/streaming.go | 3 + .../integrationtest/apidump_internal_test.go | 4 +- .../integrationtest/bridge_internal_test.go | 12 +- .../circuit_breaker_internal_test.go | 8 +- .../keypool_failover_internal_test.go | 2 +- .../internal/integrationtest/setupbridge.go | 15 +- aibridge/passthrough_internal_test.go | 5 +- aibridge/provider/anthropic.go | 35 ++- aibridge/provider/anthropic_internal_test.go | 25 +- aibridge/provider/bedrock.go | 84 +++++++ aibridge/provider/bedrock_internal_test.go | 235 ++++++++++++++++++ cli/aibridged.go | 11 +- cli/server.go | 12 + cli/server_aibridge_internal_test.go | 29 ++- coderd/ai_providers.go | 13 +- coderd/ai_providers_migrate.go | 6 + coderd/ai_providers_test.go | 105 ++++++++ coderd/aibridged/aibridged_test.go | 15 +- coderd/apidoc/docs.go | 7 + coderd/apidoc/swagger.json | 7 + coderd/database/db2sdk/db2sdk.go | 1 + codersdk/aiproviders_bedrock.go | 17 ++ codersdk/deployment.go | 7 + docs/reference/api/general.md | 2 + docs/reference/api/schemas.md | 28 ++- enterprise/aibridged_integration_test.go | 13 +- site/src/api/typesGenerated.ts | 28 +++ 34 files changed, 755 insertions(+), 200 deletions(-) create mode 100644 aibridge/provider/bedrock.go create mode 100644 aibridge/provider/bedrock_internal_test.go diff --git a/aibridge/api.go b/aibridge/api.go index 34dce84ef8873..b74912d565493 100644 --- a/aibridge/api.go +++ b/aibridge/api.go @@ -45,8 +45,8 @@ func AsActor(ctx context.Context, actorID string, metadata recorder.Metadata) co return aibcontext.AsActor(ctx, actorID, metadata) } -func NewAnthropicProvider(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) provider.Provider { - return provider.NewAnthropic(cfg, bedrockCfg) +func NewAnthropicProvider(ctx context.Context, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) (provider.Provider, error) { + return provider.NewAnthropic(ctx, cfg, bedrockCfg) } func NewOpenAIProvider(cfg config.OpenAI) provider.Provider { diff --git a/aibridge/bridge_test.go b/aibridge/bridge_test.go index 9ac7ea9ec3ddb..a51779c2335c8 100644 --- a/aibridge/bridge_test.go +++ b/aibridge/bridge_test.go @@ -2,6 +2,7 @@ package aibridge_test import ( "bytes" + "context" "fmt" "io" "net/http" @@ -22,6 +23,17 @@ import ( var bridgeTestTracer = otel.Tracer("bridge_test") +// mustNewAnthropicProvider builds an Anthropic provider for tests, panicking if +// credential resolution fails. Keeps call sites terse after NewAnthropicProvider +// gained a context and error return. +func mustNewAnthropicProvider(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) aibridge.Provider { + p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) + if err != nil { + panic("build anthropic provider: " + err.Error()) + } + return p +} + func TestValidateProviders(t *testing.T) { t.Parallel() @@ -36,7 +48,7 @@ func TestValidateProviders(t *testing.T) { name: "all_supported_providers", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: "https://api.openai.com/v1/"}), - aibridge.NewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), + mustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: "https://api.individual.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-business", BaseURL: "https://api.business.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-enterprise", BaseURL: "https://api.enterprise.githubcopilot.com"}), @@ -46,7 +58,7 @@ func TestValidateProviders(t *testing.T) { name: "default_names_and_base_urls", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{}), - aibridge.NewAnthropicProvider(config.Anthropic{}, nil), + mustNewAnthropicProvider(config.Anthropic{}, nil), aibridge.NewCopilotProvider(config.Copilot{}), }, }, @@ -150,7 +162,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "anthropic_no_base_path", requestPath: "/anthropic/v1/models", provider: func(baseURL string) provider.Provider { - return aibridge.NewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + return mustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/models", }, @@ -159,7 +171,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { baseURLPath: "/v1", requestPath: "/anthropic/v1/models", provider: func(baseURL string) provider.Provider { - return aibridge.NewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + return mustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/v1/models", }, @@ -217,7 +229,7 @@ func TestRequestBodySizeLimit(t *testing.T) { return aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: baseURL}) } newAnthropic := func(baseURL string) provider.Provider { - return aibridge.NewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) + return mustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) } newCopilot := func(baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: baseURL}) diff --git a/aibridge/config/config.go b/aibridge/config/config.go index 5805741f603f9..377c94175125d 100644 --- a/aibridge/config/config.go +++ b/aibridge/config/config.go @@ -48,6 +48,17 @@ type AWSBedrock struct { // (https://bedrock-runtime.{region}.amazonaws.com). // This is useful for routing requests through a proxy or for testing. BaseURL string + // RoleARN, when set, is assumed via STS before calling Bedrock. The base + // identity (static keys or the AWS SDK default credential chain, e.g. + // IRSA / Pod Identity / instance profile) signs the AssumeRole call, and + // the resulting temporary credentials sign Bedrock requests. This enables + // cross-account access. + RoleARN string + // ExternalID is sent on the AssumeRole call for confused-deputy protection + // when assuming a role in an account this gateway does not own. Optional. + ExternalID string + // SessionName is the STS role session name. Auto-generated when empty. + SessionName string } // OpenAI carries configuration for an OpenAI provider. diff --git a/aibridge/intercept/keyfailover_test.go b/aibridge/intercept/keyfailover_test.go index 997ca82705a9e..293de70934271 100644 --- a/aibridge/intercept/keyfailover_test.go +++ b/aibridge/intercept/keyfailover_test.go @@ -102,9 +102,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, config.ProviderAnthropic, cfg, nil, nil, http.Header{}, "X-Api-Key", tracer, cred) } - return messages.NewBlockingInterceptor(id, payload, config.ProviderAnthropic, cfg, nil, http.Header{}, "X-Api-Key", tracer, cred) + return messages.NewBlockingInterceptor(id, payload, config.ProviderAnthropic, cfg, nil, nil, http.Header{}, "X-Api-Key", tracer, cred) }, }, { diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index b167df42937cd..741162f0362b8 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -16,8 +16,7 @@ import ( "github.com/anthropics/anthropic-sdk-go/option" "github.com/anthropics/anthropic-sdk-go/shared" "github.com/anthropics/anthropic-sdk-go/shared/constant" - "github.com/aws/aws-sdk-go-v2/config" - "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" "go.opentelemetry.io/otel/attribute" "go.opentelemetry.io/otel/trace" @@ -72,6 +71,9 @@ type interceptionBase struct { cfg aibconfig.Anthropic bedrockCfg *aibconfig.AWSBedrock + // bedrockCreds is the cached AWS credentials provider resolved once per + // provider (including any assumed role). nil when not Bedrock-backed. + bedrockCreds aws.CredentialsProvider // clientHeaders are the original HTTP headers from the client request. clientHeaders http.Header @@ -253,8 +255,6 @@ func (i *interceptionBase) newMessagesService(ctx context.Context, opts ...optio } if i.bedrockCfg != nil { - ctx, cancel := context.WithTimeout(ctx, time.Second*30) - defer cancel() bedrockOpts, err := i.withAWSBedrockOptions(ctx, i.bedrockCfg) if err != nil { return anthropic.MessageService{}, err @@ -276,11 +276,14 @@ func (i *interceptionBase) withBody() option.RequestOption { // withAWSBedrockOptions returns request options for authenticating with AWS Bedrock. // -// When both AccessKey and AccessKeySecret are set in the aibridge config, they are -// used directly as static credentials. Otherwise, the AWS SDK default credential chain -// resolves credentials (environment variables, shared config/credentials files, IAM -// roles, IRSA, SSO, IMDS, etc.). -func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconfig.AWSBedrock) ([]option.RequestOption, error) { +// Credentials are resolved once per provider and supplied via i.bedrockCreds: +// either static credentials, the AWS SDK default credential chain (environment +// variables, shared config/credentials files, IAM roles, IRSA, SSO, IMDS, etc.), +// or a role assumed via STS when a target role ARN is configured. i.bedrockCreds +// is a shared, rotating credentials cache, so the per-request Retrieve below is +// served from that cache (a cache hit) and does not re-resolve or re-assume on +// every request; it only fails fast if credentials cannot be resolved. +func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconfig.AWSBedrock) ([]option.RequestOption, error) { if cfg == nil { return nil, xerrors.New("nil config given") } @@ -293,39 +296,19 @@ func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconf if cfg.SmallFastModel == "" { return nil, xerrors.New("small fast model required") } - - loadOpts := []func(*config.LoadOptions) error{ - config.WithRegion(cfg.Region), - } - - // Use static credentials when explicitly provided, otherwise fall back to the SDK default credential chain. - switch { - // Both set: use static credentials directly. - case cfg.AccessKey != "" && cfg.AccessKeySecret != "": - loadOpts = append(loadOpts, config.WithCredentialsProvider( - credentials.NewStaticCredentialsProvider( - cfg.AccessKey, - cfg.AccessKeySecret, - "", - ), - )) - // Only one set: misconfiguration. - case cfg.AccessKey != "" || cfg.AccessKeySecret != "": - return nil, xerrors.New("both access key and access key secret must be provided together") - // Neither set: SDK default credential chain resolves credentials. - default: + if i.bedrockCreds == nil { + return nil, xerrors.New("bedrock credentials not resolved") } - awsCfg, err := config.LoadDefaultConfig(ctx, loadOpts...) - if err != nil { - return nil, xerrors.Errorf("failed to load AWS Bedrock config: %w", err) + // Fail fast: ensure credentials can be resolved before signing. Served from + // the shared cache, so this does not re-assume the role on each request. + if _, err := i.bedrockCreds.Retrieve(ctx); err != nil { + return nil, xerrors.Errorf("no AWS credentials found: %w", err) } - // Fail fast: ensure credentials can be resolved before making any requests. - // awsCfg already carries the credentials provider, and the Bedrock middleware - // will call Retrieve on it when signing each request. - if _, err := awsCfg.Credentials.Retrieve(ctx); err != nil { - return nil, xerrors.Errorf("no AWS credentials found: %w", err) + awsCfg := aws.Config{ + Region: cfg.Region, + Credentials: i.bedrockCreds, } var out []option.RequestOption diff --git a/aibridge/intercept/messages/base_internal_test.go b/aibridge/intercept/messages/base_internal_test.go index f6323ec795a9d..bb4addb64ff21 100644 --- a/aibridge/intercept/messages/base_internal_test.go +++ b/aibridge/intercept/messages/base_internal_test.go @@ -9,6 +9,7 @@ import ( "github.com/anthropics/anthropic-sdk-go" "github.com/anthropics/anthropic-sdk-go/shared/constant" + "github.com/aws/aws-sdk-go-v2/credentials" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" @@ -139,29 +140,6 @@ func TestAWSBedrockValidation(t *testing.T) { expectError: true, errorMsg: "region or base url required", }, - { - name: "missing access key", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - AccessKeySecret: "test-secret", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - expectError: true, - errorMsg: "both access key and access key secret must be provided together", - }, - { - name: "missing access key secret", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - AccessKey: "test-key", - AccessKeySecret: "", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - expectError: true, - errorMsg: "both access key and access key secret must be provided together", - }, { name: "missing model", cfg: &config.AWSBedrock{ @@ -204,7 +182,12 @@ func TestAWSBedrockValidation(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - base := &interceptionBase{} + // Credentials are resolved once per provider and supplied to the + // interceptor; the per-request path only assembles options around + // them. A static stub stands in for the resolved provider. + base := &interceptionBase{ + bedrockCreds: credentials.NewStaticCredentialsProvider("test-key", "test-secret", ""), + } opts, err := base.withAWSBedrockOptions(context.Background(), tt.cfg) if tt.expectError { @@ -218,87 +201,19 @@ func TestAWSBedrockValidation(t *testing.T) { } } -// TestAWSBedrockCredentialChain tests credential resolution via the AWS SDK default credential chain. -// NOTE: Cannot use t.Parallel() here because subtests use t.Setenv which requires sequential execution. -func TestAWSBedrockCredentialChain(t *testing.T) { - tests := []struct { - name string - cfg *config.AWSBedrock - envVars map[string]string - expectError bool - errorMsg string - }{ - { - name: "temporary credentials via env", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - envVars: map[string]string{ - "AWS_ACCESS_KEY_ID": "test-key", - "AWS_SECRET_ACCESS_KEY": "test-secret", - }, - }, - { - name: "temporary credentials with session token via env", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - envVars: map[string]string{ - "AWS_ACCESS_KEY_ID": "test-key", - "AWS_SECRET_ACCESS_KEY": "test-secret", - "AWS_SESSION_TOKEN": "test-session-token", - }, - }, - { - // When static credentials are not provided and no environment credentials are set, - // the SDK default credential chain fails to resolve credentials. - name: "error when no credential source is configured", - cfg: &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }, - envVars: map[string]string{ - "AWS_ACCESS_KEY_ID": "", - "AWS_SECRET_ACCESS_KEY": "", - "AWS_SESSION_TOKEN": "", - "AWS_PROFILE": "", - "AWS_SHARED_CREDENTIALS_FILE": "/dev/null", - "AWS_CONFIG_FILE": "/dev/null", - "AWS_WEB_IDENTITY_TOKEN_FILE": "", - "AWS_ROLE_ARN": "", - "AWS_ROLE_SESSION_NAME": "", - "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI": "", - "AWS_CONTAINER_CREDENTIALS_FULL_URI": "", - "AWS_CONTAINER_AUTHORIZATION_TOKEN": "", - "AWS_EC2_METADATA_DISABLED": "true", - }, - expectError: true, - errorMsg: "no AWS credentials found", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - for key, val := range tt.envVars { - t.Setenv(key, val) - } - base := &interceptionBase{} - opts, err := base.withAWSBedrockOptions(context.Background(), tt.cfg) +// TestAWSBedrockOptionsRequireResolvedCredentials verifies that option assembly +// fails when the per-provider credentials provider was not resolved. +func TestAWSBedrockOptionsRequireResolvedCredentials(t *testing.T) { + t.Parallel() - if tt.expectError { - require.Error(t, err) - require.Contains(t, err.Error(), tt.errorMsg) - } else { - require.NotEmpty(t, opts) - require.NoError(t, err) - } - }) - } + base := &interceptionBase{} + _, err := base.withAWSBedrockOptions(context.Background(), &config.AWSBedrock{ + Region: "us-east-1", + Model: "test-model", + SmallFastModel: "test-small-model", + }) + require.Error(t, err) + require.Contains(t, err.Error(), "bedrock credentials not resolved") } func TestAccumulateUsage(t *testing.T) { diff --git a/aibridge/intercept/messages/blocking.go b/aibridge/intercept/messages/blocking.go index 4370676ce7e85..a0db85236ee8d 100644 --- a/aibridge/intercept/messages/blocking.go +++ b/aibridge/intercept/messages/blocking.go @@ -9,6 +9,7 @@ import ( "github.com/anthropics/anthropic-sdk-go" "github.com/anthropics/anthropic-sdk-go/option" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" mcplib "github.com/mark3labs/mcp-go/mcp" "github.com/tidwall/sjson" @@ -37,6 +38,7 @@ func NewBlockingInterceptor( providerName string, cfg config.Anthropic, bedrockCfg *config.AWSBedrock, + bedrockCreds aws.CredentialsProvider, clientHeaders http.Header, authHeaderName string, tracer trace.Tracer, @@ -48,6 +50,7 @@ func NewBlockingInterceptor( reqPayload: reqPayload, cfg: cfg, bedrockCfg: bedrockCfg, + bedrockCreds: bedrockCreds, clientHeaders: clientHeaders, authHeaderName: authHeaderName, tracer: tracer, diff --git a/aibridge/intercept/messages/streaming.go b/aibridge/intercept/messages/streaming.go index 1a383889e3416..7b438aebbf853 100644 --- a/aibridge/intercept/messages/streaming.go +++ b/aibridge/intercept/messages/streaming.go @@ -13,6 +13,7 @@ import ( "github.com/anthropics/anthropic-sdk-go/option" "github.com/anthropics/anthropic-sdk-go/packages/ssestream" "github.com/anthropics/anthropic-sdk-go/shared/constant" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" mcplib "github.com/mark3labs/mcp-go/mcp" "github.com/tidwall/sjson" @@ -42,6 +43,7 @@ func NewStreamingInterceptor( providerName string, cfg config.Anthropic, bedrockCfg *config.AWSBedrock, + bedrockCreds aws.CredentialsProvider, clientHeaders http.Header, authHeaderName string, tracer trace.Tracer, @@ -53,6 +55,7 @@ func NewStreamingInterceptor( reqPayload: reqPayload, cfg: cfg, bedrockCfg: bedrockCfg, + bedrockCreds: bedrockCreds, clientHeaders: clientHeaders, authHeaderName: authHeaderName, tracer: tracer, diff --git a/aibridge/internal/integrationtest/apidump_internal_test.go b/aibridge/internal/integrationtest/apidump_internal_test.go index 42811cb362ac0..cff3a10eed01a 100644 --- a/aibridge/internal/integrationtest/apidump_internal_test.go +++ b/aibridge/internal/integrationtest/apidump_internal_test.go @@ -39,7 +39,7 @@ func TestAPIDump(t *testing.T) { name: "anthropic", fixture: fixtures.AntSimple, providerFunc: func(addr, dumpDir string) aibridge.Provider { - return provider.NewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return mustNewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, path: pathAnthropicMessages, expectProviderDir: config.ProviderAnthropic, @@ -219,7 +219,7 @@ func TestAPIDumpPassthrough(t *testing.T) { { name: "anthropic", providerFunc: func(addr string, dumpDir string) aibridge.Provider { - return provider.NewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return mustNewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, requestPath: "/anthropic/v1/models", expectDumpName: "-v1-models-", diff --git a/aibridge/internal/integrationtest/bridge_internal_test.go b/aibridge/internal/integrationtest/bridge_internal_test.go index ef226db2b989b..62d9cd36c9ac0 100644 --- a/aibridge/internal/integrationtest/bridge_internal_test.go +++ b/aibridge/internal/integrationtest/bridge_internal_test.go @@ -317,7 +317,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, "http://unused", - withCustomProvider(provider.NewAnthropic(anthropicCfg("http://unused", apiKey), bedrockCfg)), + withCustomProvider(mustNewAnthropic(anthropicCfg("http://unused", apiKey), bedrockCfg)), ) resp, err := bridgeServer.makeRequest(t, http.MethodPost, pathAnthropicMessages, fixtures.Request(t, fixtures.AntSingleBuiltinTool)) @@ -353,7 +353,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(upstream.URL, apiKey), bedrockCfg)), + withCustomProvider(mustNewAnthropic(anthropicCfg(upstream.URL, apiKey), bedrockCfg)), ) // Make API call to aibridge for Anthropic /v1/messages, which will be routed via AWS Bedrock. @@ -483,7 +483,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(upstream.URL, apiKey), bCfg)), + withCustomProvider(mustNewAnthropic(anthropicCfg(upstream.URL, apiKey), bCfg)), ) reqBody, err := sjson.SetBytes(fix.Request(), "stream", streaming) @@ -637,7 +637,7 @@ func TestAWSBedrockIntegration(t *testing.T) { bCfg.Region = region bridgeServer := newBridgeTestServer(ctx, t, mockEgressProxy.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), + withCustomProvider(mustNewAnthropic(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), ) // Sends a bridge request through a mock egress proxy that @@ -2292,7 +2292,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return provider.NewAnthropic(cfg, nil) + return mustNewAnthropic(cfg, nil) }, fixture: fixtures.AntSimple, streaming: true, @@ -2303,7 +2303,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return provider.NewAnthropic(cfg, nil) + return mustNewAnthropic(cfg, nil) }, fixture: fixtures.AntSimple, streaming: false, diff --git a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go index afa6091e2a949..411401b55b0a5 100644 --- a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go +++ b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go @@ -68,7 +68,7 @@ func TestCircuitBreaker_FullRecoveryCycle(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + return mustNewAnthropic(config.Anthropic{ BaseURL: baseURL, Key: "test-key", CircuitBreaker: cbConfig, @@ -235,7 +235,7 @@ func TestCircuitBreaker_HalfOpenFailure(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + return mustNewAnthropic(config.Anthropic{ BaseURL: baseURL, Key: "test-key", CircuitBreaker: cbConfig, @@ -372,7 +372,7 @@ func TestCircuitBreaker_HalfOpenMaxRequests(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + return mustNewAnthropic(config.Anthropic{ BaseURL: baseURL, Key: "test-key", CircuitBreaker: cbConfig, @@ -553,7 +553,7 @@ func TestCircuitBreaker_PerModelIsolation(t *testing.T) { } ctx := t.Context() bridgeServer := newBridgeTestServer(ctx, t, mockUpstream.URL, - withCustomProvider(provider.NewAnthropic(config.Anthropic{ + withCustomProvider(mustNewAnthropic(config.Anthropic{ BaseURL: mockUpstream.URL, Key: "test-key", CircuitBreaker: cbConfig, diff --git a/aibridge/internal/integrationtest/keypool_failover_internal_test.go b/aibridge/internal/integrationtest/keypool_failover_internal_test.go index f186fafa36afb..00f9bc8728c02 100644 --- a/aibridge/internal/integrationtest/keypool_failover_internal_test.go +++ b/aibridge/internal/integrationtest/keypool_failover_internal_test.go @@ -219,7 +219,7 @@ func TestAnthropic_KeyFailover(t *testing.T) { t.Cleanup(upstream.Close) bridgeServer := newBridgeTestServer(t.Context(), t, upstream.URL, - withCustomProvider(provider.NewAnthropic(config.Anthropic{ + withCustomProvider(mustNewAnthropic(config.Anthropic{ BaseURL: upstream.URL, KeyPool: pool, }, nil)), diff --git a/aibridge/internal/integrationtest/setupbridge.go b/aibridge/internal/integrationtest/setupbridge.go index e63f554a009e0..d0a05abc0a300 100644 --- a/aibridge/internal/integrationtest/setupbridge.go +++ b/aibridge/internal/integrationtest/setupbridge.go @@ -249,15 +249,26 @@ func setupInjectedToolTest( return bridgeServer, mockMCP, resp } +// mustNewAnthropic builds an Anthropic provider for tests, panicking if +// credential resolution fails. NewAnthropic resolves Bedrock credentials at +// construction, so this keeps the many test call sites terse. +func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) aibridge.Provider { + p, err := provider.NewAnthropic(context.Background(), cfg, bedrockCfg) + if err != nil { + panic("build anthropic provider: " + err.Error()) + } + return p +} + // newDefaultProvider creates a Provider with default test configuration. func newDefaultProvider(providerType string, addr string) aibridge.Provider { switch providerType { case config.ProviderAnthropic: - return provider.NewAnthropic(anthropicCfg(addr, apiKey), nil) + return mustNewAnthropic(anthropicCfg(addr, apiKey), nil) case config.ProviderOpenAI: return provider.NewOpenAI(openAICfg(addr, apiKey)) case providerBedrock: - return provider.NewAnthropic(anthropicCfg(addr, apiKey), bedrockCfg(addr)) + return mustNewAnthropic(anthropicCfg(addr, apiKey), bedrockCfg(addr)) default: panic("unknown provider type: " + providerType) } diff --git a/aibridge/passthrough_internal_test.go b/aibridge/passthrough_internal_test.go index c095281bf3927..09c4f552591a6 100644 --- a/aibridge/passthrough_internal_test.go +++ b/aibridge/passthrough_internal_test.go @@ -1,6 +1,7 @@ package aibridge import ( + "context" "crypto/tls" "io" "maps" @@ -335,10 +336,12 @@ func TestPassthrough_KeyFailover(t *testing.T) { r.Header.Set("X-Api-Key", key) }, newProvider: func(baseURL string, pool *keypool.Pool) provider.Provider { - return provider.NewAnthropic(config.Anthropic{ + p, err := provider.NewAnthropic(context.Background(), config.Anthropic{ BaseURL: baseURL, KeyPool: pool, }, nil) + require.NoError(t, err) + return p }, }, { diff --git a/aibridge/provider/anthropic.go b/aibridge/provider/anthropic.go index 0757296f814e7..787fee30403bd 100644 --- a/aibridge/provider/anthropic.go +++ b/aibridge/provider/anthropic.go @@ -1,11 +1,13 @@ package provider import ( + "context" "fmt" "io" "net/http" "strings" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" "go.opentelemetry.io/otel/codes" "go.opentelemetry.io/otel/trace" @@ -35,6 +37,11 @@ var _ Provider = &Anthropic{} type Anthropic struct { cfg config.Anthropic bedrockCfg *config.AWSBedrock + // bedrockCreds is the AWS credentials provider (including any assumed + // role), resolved once at construction and shared across requests so + // per-request retrieval is served from its cache rather than re-resolved. + // nil when this provider is not Bedrock-backed. + bedrockCreds aws.CredentialsProvider } const routeMessages = "/v1/messages" // https://docs.anthropic.com/en/api/messages @@ -51,7 +58,7 @@ var anthropicIsFailure = func(statusCode int) bool { return circuitbreaker.DefaultIsFailure(statusCode) } -func NewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { +func NewAnthropic(ctx context.Context, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) (*Anthropic, error) { if cfg.Name == "" { cfg.Name = config.ProviderAnthropic } @@ -81,10 +88,24 @@ func NewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropi cfg.CircuitBreaker.OpenErrorResponse = anthropicOpenErrorResponse } - return &Anthropic{ - cfg: cfg, - bedrockCfg: bedrockCfg, + // Resolve the AWS credentials provider once. This performs no network + // call (the base identity and any AssumeRole resolve lazily on first + // retrieval); it only wires up the provider chain, so it is cheap to run + // at construction and on every provider reload. + var bedrockCreds aws.CredentialsProvider + if bedrockCfg != nil { + var err error + bedrockCreds, err = buildBedrockCredentials(ctx, *bedrockCfg) + if err != nil { + return nil, xerrors.Errorf("build bedrock credentials: %w", err) + } } + + return &Anthropic{ + cfg: cfg, + bedrockCfg: bedrockCfg, + bedrockCreds: bedrockCreds, + }, nil } func (*Anthropic) Type() string { @@ -176,11 +197,13 @@ func (p *Anthropic) CreateInterceptor(_ http.ResponseWriter, r *http.Request, tr // end-of-interception. cred := intercept.NewCredentialInfo(credKind, credSecret) + // bedrockCreds was resolved once at construction; it is a shared, + // rotating credentials cache. nil for non-Bedrock providers. 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, p.Name(), cfg, p.bedrockCfg, p.bedrockCreds, r.Header, authHeaderName, tracer, cred) } else { - interceptor = messages.NewBlockingInterceptor(id, reqPayload, p.Name(), cfg, p.bedrockCfg, r.Header, authHeaderName, tracer, cred) + interceptor = messages.NewBlockingInterceptor(id, reqPayload, p.Name(), cfg, p.bedrockCfg, p.bedrockCreds, r.Header, authHeaderName, tracer, cred) } span.SetAttributes(interceptor.TraceAttributes(r)...) return interceptor, nil diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 285fa3cd04c72..d342fb7b48f55 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -2,6 +2,7 @@ package provider import ( "bytes" + "context" "net/http" "net/http/httptest" "testing" @@ -18,6 +19,18 @@ import ( "github.com/coder/quartz" ) +// mustNewAnthropic builds an Anthropic provider for tests, panicking if +// credential resolution fails (it cannot for the non-Bedrock configs used +// here). Keeps the call sites terse after NewAnthropic gained a context and +// error return. +func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { + p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) + if err != nil { + panic("build anthropic provider: " + err.Error()) + } + return p +} + func TestAnthropic_TypeAndName(t *testing.T) { t.Parallel() @@ -45,7 +58,7 @@ func TestAnthropic_TypeAndName(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := NewAnthropic(tc.cfg, nil) + p := mustNewAnthropic(tc.cfg, nil) assert.Equal(t, tc.expectType, p.Type()) assert.Equal(t, tc.expectName, p.Name()) }) @@ -94,7 +107,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := NewAnthropic(tc.cfg, nil) + p := mustNewAnthropic(tc.cfg, nil) if tc.expectedKeys == nil { assert.Nil(t, p.cfg.KeyPool, "expected no KeyPool") @@ -119,7 +132,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { func TestAnthropic_CreateInterceptor(t *testing.T) { t.Parallel() - provider := NewAnthropic(config.Anthropic{Key: "test-key"}, nil) + provider := mustNewAnthropic(config.Anthropic{Key: "test-key"}, nil) t.Run("Messages_NonStreamingRequest_BlockingInterceptor", func(t *testing.T) { t.Parallel() @@ -177,7 +190,7 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := NewAnthropic(config.Anthropic{ + provider := mustNewAnthropic(config.Anthropic{ BaseURL: mockUpstream.URL, Key: "test-key", }, nil) @@ -287,7 +300,7 @@ func TestAnthropic_CreateInterceptor_BYOK(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := NewAnthropic(config.Anthropic{ + provider := mustNewAnthropic(config.Anthropic{ BaseURL: mockUpstream.URL, Key: "test-key", }, nil) @@ -326,7 +339,7 @@ func TestAnthropic_KeyFailoverConfig(t *testing.T) { pool, err := keypool.New(config.ProviderAnthropic, []string{"k0", "k1"}, quartz.NewMock(t), nil) require.NoError(t, err) - p := NewAnthropic(config.Anthropic{KeyPool: pool}, nil) + p := mustNewAnthropic(config.Anthropic{KeyPool: pool}, nil) cfg := p.KeyFailoverConfig(slog.Make()) diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go new file mode 100644 index 0000000000000..c77b1b904cd28 --- /dev/null +++ b/aibridge/provider/bedrock.go @@ -0,0 +1,84 @@ +package provider + +import ( + "context" + + "github.com/aws/aws-sdk-go-v2/aws" + awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/sts" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/aibridge/config" +) + +// defaultBedrockSessionName is the STS role session name used when the provider +// does not configure one. A stable value keeps AssumeRole calls identifiable in +// CloudTrail. +const defaultBedrockSessionName = "coder-aibridge" + +// buildBedrockCredentials resolves the base identity (static keys or the AWS SDK +// default credential chain, which covers IRSA, Pod Identity, instance profile, +// shared profile, and environment variables) and, when a target role ARN is +// configured, assumes that role via STS. +// +// The result is wrapped in aws.NewCredentialsCache, which caches and rotates the +// resolved temporary credentials. The provider is resolved once per Bedrock +// provider (at construction) and shared across requests, so per-request +// credential retrieval is served from this cache rather than re-resolving (and +// re-assuming) on every request. No network call is made here: the base +// identity and any AssumeRole are resolved lazily on first retrieval. +func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.CredentialsProvider, error) { + if cfg.Region == "" && cfg.BaseURL == "" { + return nil, xerrors.New("region or base url required") + } + + var loadOpts []func(*awsconfig.LoadOptions) error + if cfg.Region != "" { + loadOpts = append(loadOpts, awsconfig.WithRegion(cfg.Region)) + } + + // Use static credentials when explicitly provided, otherwise fall back to + // the SDK default credential chain. + switch { + // Both set: use static credentials directly. + case cfg.AccessKey != "" && cfg.AccessKeySecret != "": + loadOpts = append(loadOpts, awsconfig.WithCredentialsProvider( + credentials.NewStaticCredentialsProvider( + cfg.AccessKey, + cfg.AccessKeySecret, + "", + ), + )) + // Only one set: misconfiguration. + case cfg.AccessKey != "" || cfg.AccessKeySecret != "": + return nil, xerrors.New("both access key and access key secret must be provided together") + // Neither set: SDK default credential chain resolves the base identity. + default: + } + + base, err := awsconfig.LoadDefaultConfig(ctx, loadOpts...) + if err != nil { + return nil, xerrors.Errorf("failed to load AWS Bedrock config: %w", err) + } + + // The base identity signs requests directly unless a target role is + // configured, in which case it signs the AssumeRole call and the resulting + // temporary credentials sign Bedrock requests. + credsProvider := base.Credentials + if cfg.RoleARN != "" { + sessionName := cfg.SessionName + if sessionName == "" { + sessionName = defaultBedrockSessionName + } + credsProvider = stscreds.NewAssumeRoleProvider(sts.NewFromConfig(base), cfg.RoleARN, func(o *stscreds.AssumeRoleOptions) { + o.RoleSessionName = sessionName + if cfg.ExternalID != "" { + o.ExternalID = aws.String(cfg.ExternalID) + } + }) + } + + return aws.NewCredentialsCache(credsProvider), nil +} diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go new file mode 100644 index 0000000000000..4cba363e8093a --- /dev/null +++ b/aibridge/provider/bedrock_internal_test.go @@ -0,0 +1,235 @@ +package provider + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/aibridge/config" +) + +// TestBuildBedrockCredentialsValidation covers the input validation that does +// not require resolving credentials. +func TestBuildBedrockCredentialsValidation(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + cfg config.AWSBedrock + errorMsg string + }{ + { + name: "missing region and base url", + cfg: config.AWSBedrock{}, + errorMsg: "region or base url required", + }, + { + name: "missing access key", + cfg: config.AWSBedrock{ + Region: "us-east-1", + AccessKeySecret: "test-secret", + }, + errorMsg: "both access key and access key secret must be provided together", + }, + { + name: "missing access key secret", + cfg: config.AWSBedrock{ + Region: "us-east-1", + AccessKey: "test-key", + }, + errorMsg: "both access key and access key secret must be provided together", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + _, err := buildBedrockCredentials(context.Background(), tt.cfg) + require.Error(t, err) + require.Contains(t, err.Error(), tt.errorMsg) + }) + } +} + +// TestBuildBedrockCredentialsStatic resolves static credentials with no network +// access. +func TestBuildBedrockCredentialsStatic(t *testing.T) { + t.Parallel() + + creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + AccessKey: "test-key", + AccessKeySecret: "test-secret", + }) + require.NoError(t, err) + + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "test-key", got.AccessKeyID) + require.Equal(t, "test-secret", got.SecretAccessKey) +} + +// TestBuildBedrockCredentialsDefaultChain covers resolution via the AWS SDK +// default credential chain (here, environment variables) and the failure mode +// when no credential source is configured. +// NOTE: no t.Parallel() because the subtests use t.Setenv. +func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { + tests := []struct { + name string + envVars map[string]string + expectError bool + }{ + { + name: "credentials via env", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "test-key", + "AWS_SECRET_ACCESS_KEY": "test-secret", + }, + }, + { + name: "credentials with session token via env", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "test-key", + "AWS_SECRET_ACCESS_KEY": "test-secret", + "AWS_SESSION_TOKEN": "test-session-token", + }, + }, + { + name: "error when no credential source is configured", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "", + "AWS_SECRET_ACCESS_KEY": "", + "AWS_SESSION_TOKEN": "", + "AWS_PROFILE": "", + "AWS_SHARED_CREDENTIALS_FILE": "/dev/null", + "AWS_CONFIG_FILE": "/dev/null", + "AWS_WEB_IDENTITY_TOKEN_FILE": "", + "AWS_ROLE_ARN": "", + "AWS_ROLE_SESSION_NAME": "", + "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI": "", + "AWS_CONTAINER_CREDENTIALS_FULL_URI": "", + "AWS_CONTAINER_AUTHORIZATION_TOKEN": "", + "AWS_EC2_METADATA_DISABLED": "true", + }, + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for key, val := range tt.envVars { + t.Setenv(key, val) + } + + // buildBedrockCredentials wires up the provider chain without a + // network call, so it succeeds regardless of credential + // availability; resolution failures surface on Retrieve. + creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + }) + require.NoError(t, err) + require.NotNil(t, creds) + + _, err = creds.Retrieve(context.Background()) + if tt.expectError { + require.Error(t, err) + return + } + require.NoError(t, err) + }) + } +} + +// TestBuildBedrockCredentialsAssumeRole drives the STS AssumeRole path against a +// mock endpoint, asserting that the configured role ARN, external ID, and +// session name are sent and that the returned temporary credentials are used. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { + var gotRoleARN, gotExternalID, gotSessionName string + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + gotRoleARN = r.Form.Get("RoleArn") + gotExternalID = r.Form.Get("ExternalId") + gotSessionName = r.Form.Get("RoleSessionName") + + w.Header().Set("Content-Type", "text/xml") + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2999-01-01T00:00:00Z + + + arn:aws:sts::123456789012:assumed-role/target/coder + AROAEXAMPLE:coder + + +`)) + })) + defer sts.Close() + + // Point the STS client at the mock and provide static base credentials so + // the base identity resolves without network access. + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + ExternalID: "shared-secret", + SessionName: "coder", + }) + require.NoError(t, err) + + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "ASIAASSUMED", got.AccessKeyID) + require.Equal(t, "assumed-secret", got.SecretAccessKey) + require.Equal(t, "assumed-token", got.SessionToken) + + require.Equal(t, "arn:aws:iam::123456789012:role/target", gotRoleARN) + require.Equal(t, "shared-secret", gotExternalID) + require.Equal(t, "coder", gotSessionName) +} + +// TestBuildBedrockCredentialsAssumeRoleDefaultSessionName verifies a session +// name is sent even when the provider does not configure one. +func TestBuildBedrockCredentialsAssumeRoleDefaultSessionName(t *testing.T) { + var gotSessionName string + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + gotSessionName = r.Form.Get("RoleSessionName") + w.Header().Set("Content-Type", "text/xml") + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2999-01-01T00:00:00Z + + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + _, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, defaultBedrockSessionName, gotSessionName) +} diff --git a/cli/aibridged.go b/cli/aibridged.go index e838c1db5ab0a..f3f21c805de2e 100644 --- a/cli/aibridged.go +++ b/cli/aibridged.go @@ -175,7 +175,7 @@ func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridg if row.Enabled { enabledCount++ } - prov, err := buildAIProviderFromRow(row, keysByProvider[row.ID], cfg, metrics) + prov, err := buildAIProviderFromRow(ctx, row, keysByProvider[row.ID], cfg, metrics) if err != nil { outcome.Status = aibridged.ProviderStatusError outcome.Err = err @@ -210,6 +210,7 @@ func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridg // Disabled: true; settings decode, key loading, and credential checks // are skipped because the provider will never call upstream. func buildAIProviderFromRow( + ctx context.Context, row database.AIProvider, keys []database.AIProviderKey, cfg codersdk.AIBridgeConfig, @@ -284,14 +285,14 @@ func buildAIProviderFromRow( return nil, xerrors.Errorf("anthropic key pool: %w", err) } } - return aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + return aibridge.NewAnthropicProvider(ctx, aibridge.AnthropicConfig{ Name: row.Name, BaseURL: row.BaseUrl, KeyPool: pool, APIDumpDir: dumpDir, CircuitBreaker: cbCfg, SendActorHeaders: sendActorHeaders, - }, bedrock), nil + }, bedrock) case database.AIProviderTypeCopilot: // Copilot is always BYOK; the per-user token is supplied on each @@ -341,6 +342,7 @@ func bedrockConfigFromRow(row database.AIProvider, settings codersdk.AIProviderS } accessKey := ptr.NilToEmpty(bedrockSettings.AccessKey) accessKeySecret := ptr.NilToEmpty(bedrockSettings.AccessKeySecret) + externalID := ptr.NilToEmpty(bedrockSettings.ExternalID) return &aibridge.AWSBedrockConfig{ BaseURL: row.BaseUrl, Region: bedrockSettings.Region, @@ -348,6 +350,9 @@ func bedrockConfigFromRow(row database.AIProvider, settings codersdk.AIProviderS AccessKeySecret: accessKeySecret, Model: bedrockSettings.Model, SmallFastModel: bedrockSettings.SmallFastModel, + RoleARN: bedrockSettings.RoleARN, + ExternalID: externalID, + SessionName: bedrockSettings.SessionName, } } diff --git a/cli/server.go b/cli/server.go index b521f7b4b63f5..7c0f6bba1a03d 100644 --- a/cli/server.go +++ b/cli/server.go @@ -3132,6 +3132,12 @@ func ReadAIProvidersFromEnv(logger slog.Logger, environ []string) ([]codersdk.AI p.BedrockRegion, bedrockKey, bedrockSecret, p.BedrockModel, p.BedrockSmallFastModel, ) + settings.RoleARN = p.BedrockRoleARN + settings.SessionName = p.BedrockSessionName + if p.BedrockExternalID != "" { + externalID := p.BedrockExternalID + settings.ExternalID = &externalID + } isBedrock := codersdk.IsBedrockConfigured(p.BedrockBaseURL, settings) // BEDROCK_* fields are accepted on anthropic (mutually exclusive @@ -3294,6 +3300,12 @@ func readAIProvidersForPrefix(logger slog.Logger, environ []string, prefix strin provider.BedrockModel = v.Value case "BEDROCK_SMALL_FAST_MODEL": provider.BedrockSmallFastModel = v.Value + case "BEDROCK_ROLE_ARN": + provider.BedrockRoleARN = v.Value + case "BEDROCK_EXTERNAL_ID": + provider.BedrockExternalID = v.Value + case "BEDROCK_SESSION_NAME": + provider.BedrockSessionName = v.Value default: logger.Warn(context.Background(), "ignoring unknown AI provider field (check for typos)", slog.F("env", fullName), diff --git a/cli/server_aibridge_internal_test.go b/cli/server_aibridge_internal_test.go index 21711b0289e57..e6be4287382ac 100644 --- a/cli/server_aibridge_internal_test.go +++ b/cli/server_aibridge_internal_test.go @@ -131,6 +131,31 @@ func TestReadAIProvidersFromEnv(t *testing.T) { }, }, }, + { + name: "BedrockAssumeRoleFields", + env: []string{ + "CODER_AIBRIDGE_PROVIDER_0_TYPE=bedrock", + "CODER_AIBRIDGE_PROVIDER_0_NAME=bedrock-unit-a", + "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_REGION=us-east-1", + "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_MODEL=anthropic.claude-3-sonnet", + "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_SMALL_FAST_MODEL=anthropic.claude-3-haiku", + "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_ROLE_ARN=arn:aws:iam::123456789012:role/target", + "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_EXTERNAL_ID=shared-secret", + "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_SESSION_NAME=coder", + }, + expected: []codersdk.AIProviderConfig{ + { + Type: string(database.AIProviderTypeBedrock), + Name: "bedrock-unit-a", + BedrockRegion: "us-east-1", + BedrockModel: "anthropic.claude-3-sonnet", + BedrockSmallFastModel: "anthropic.claude-3-haiku", + BedrockRoleARN: "arn:aws:iam::123456789012:role/target", + BedrockExternalID: "shared-secret", + BedrockSessionName: "coder", + }, + }, + }, { name: "OutOfOrderIndices", env: []string{ @@ -741,7 +766,7 @@ func TestBuildAIProviderFromRowSetsAPIDumpDir(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - provider, err := buildAIProviderFromRow(tt.row, nil, codersdk.AIBridgeConfig{ + provider, err := buildAIProviderFromRow(t.Context(), tt.row, nil, codersdk.AIBridgeConfig{ AllowBYOK: serpent.Bool(true), APIDumpDir: serpent.String(dumpDir), }, nil) @@ -755,7 +780,7 @@ func TestBuildAIProviderFromRowSetsAPIDumpDir(t *testing.T) { func TestBuildAIProviderFromRowBedrockWithoutSettings(t *testing.T) { t.Parallel() - _, err := buildAIProviderFromRow(database.AIProvider{ + _, err := buildAIProviderFromRow(t.Context(), database.AIProvider{ Enabled: true, Type: database.AIProviderTypeBedrock, Name: "bedrock-no-settings", diff --git a/coderd/ai_providers.go b/coderd/ai_providers.go index 49e5bb8750403..bc83e88880185 100644 --- a/coderd/ai_providers.go +++ b/coderd/ai_providers.go @@ -754,11 +754,11 @@ func encodeAIProviderSettings(s codersdk.AIProviderSettings) (sql.NullString, er } // mergeAIProviderSettings overlays a patch onto an existing settings -// value. Write-only fields (Bedrock AccessKey and AccessKeySecret) use -// pointers so the patch can distinguish "omitted, keep existing" (nil) -// from "explicitly clear" (pointer to empty string) - e.g. when an -// admin migrates from static AWS credentials to IAM role-based auth -// in a single PATCH. +// value. Write-only fields (Bedrock AccessKey, AccessKeySecret, and +// ExternalID) use pointers so the patch can distinguish "omitted, keep +// existing" (nil) from "explicitly clear" (pointer to empty string) - +// e.g. when an admin migrates from static AWS credentials to IAM +// role-based auth in a single PATCH. func mergeAIProviderSettings(existing, patch codersdk.AIProviderSettings) codersdk.AIProviderSettings { if patch.Bedrock == nil { // Patch carries no type-specific data; treat as a clear. @@ -772,6 +772,9 @@ func mergeAIProviderSettings(existing, patch codersdk.AIProviderSettings) coders if merged.AccessKeySecret == nil { merged.AccessKeySecret = existing.Bedrock.AccessKeySecret } + if merged.ExternalID == nil { + merged.ExternalID = existing.Bedrock.ExternalID + } } return codersdk.AIProviderSettings{Bedrock: &merged} } diff --git a/coderd/ai_providers_migrate.go b/coderd/ai_providers_migrate.go index 7ec6a997b0dfe..2130d2e95a171 100644 --- a/coderd/ai_providers_migrate.go +++ b/coderd/ai_providers_migrate.go @@ -404,6 +404,12 @@ func providersFromEnv(ctx context.Context, cfg codersdk.AIBridgeConfig, logger s p.BedrockModel, p.BedrockSmallFastModel, ) + bedrock.RoleARN = p.BedrockRoleARN + bedrock.SessionName = p.BedrockSessionName + if p.BedrockExternalID != "" { + externalID := p.BedrockExternalID + bedrock.ExternalID = &externalID + } isBedrock = codersdk.IsBedrockConfigured(p.BedrockBaseURL, bedrock) if isBedrock { dp.Bedrock = &bedrock diff --git a/coderd/ai_providers_test.go b/coderd/ai_providers_test.go index b9bfd283f1c9e..b120c78e0ce9d 100644 --- a/coderd/ai_providers_test.go +++ b/coderd/ai_providers_test.go @@ -1501,4 +1501,109 @@ func TestAIProviderSettingsMerge(t *testing.T) { require.NotNil(t, persisted.Bedrock.AccessKeySecret) require.Equal(t, "secret-new", *persisted.Bedrock.AccessKeySecret) }) + + t.Run("MigrateStaticToRoleWithExternalID", func(t *testing.T) { + t.Parallel() + // An admin migrating from static AWS credentials to cross-account + // role assumption clears the keys and sets a role ARN and external + // ID in a single PATCH. The external ID is a write-only secret: it + // must persist but never be echoed back in API responses. + client, db := coderdtest.NewWithDatabase(t, nil) + _ = coderdtest.CreateFirstUser(t, client) + ctx := testutil.Context(t, testutil.WaitLong) + + //nolint:gocritic // Owner role is the audience for this endpoint. + created, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{ + Type: codersdk.AIProviderTypeAnthropic, + Name: "merge-role", + Enabled: true, + BaseURL: "https://bedrock-runtime.us-east-1.amazonaws.com/", + Settings: codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + AccessKey: ptr.Ref("AKIA-old"), //nolint:gosec // test fixture, not a real credential + AccessKeySecret: ptr.Ref("secret-old"), + }, + }, + }) + require.NoError(t, err) + + updated, err := client.UpdateAIProvider(ctx, created.Name, codersdk.UpdateAIProviderRequest{ + Settings: &codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + AccessKey: ptr.Ref(""), + AccessKeySecret: ptr.Ref(""), + RoleARN: "arn:aws:iam::123456789012:role/target", + SessionName: "coder", + ExternalID: ptr.Ref("shared-secret"), + }, + }, + }) + require.NoError(t, err) + + // The API response must redact the write-only external ID while + // still echoing the non-secret role fields. + require.NotNil(t, updated.Settings.Bedrock) + require.Nil(t, updated.Settings.Bedrock.ExternalID) + require.Equal(t, "arn:aws:iam::123456789012:role/target", updated.Settings.Bedrock.RoleARN) + require.Equal(t, "coder", updated.Settings.Bedrock.SessionName) + + //nolint:gocritic // Test reads the row to verify write-only fields. + row, err := db.GetAIProviderByID(dbauthz.AsSystemRestricted(ctx), created.ID) + require.NoError(t, err) + persisted, err := db2sdk.AIProviderSettings(row.Settings) + require.NoError(t, err) + require.NotNil(t, persisted.Bedrock) + require.Equal(t, "arn:aws:iam::123456789012:role/target", persisted.Bedrock.RoleARN) + require.Equal(t, "coder", persisted.Bedrock.SessionName) + require.NotNil(t, persisted.Bedrock.ExternalID) + require.Equal(t, "shared-secret", *persisted.Bedrock.ExternalID) + require.NotNil(t, persisted.Bedrock.AccessKey) + require.Equal(t, "", *persisted.Bedrock.AccessKey) + }) + + t.Run("OmittedExternalIDPreservesExisting", func(t *testing.T) { + t.Parallel() + // A PATCH that omits the external ID must keep the existing value, + // matching the write-only semantics of AccessKeySecret. + client, db := coderdtest.NewWithDatabase(t, nil) + _ = coderdtest.CreateFirstUser(t, client) + ctx := testutil.Context(t, testutil.WaitLong) + + //nolint:gocritic // Owner role is the audience for this endpoint. + created, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{ + Type: codersdk.AIProviderTypeAnthropic, + Name: "merge-role-omit", + Enabled: true, + BaseURL: "https://bedrock-runtime.us-east-1.amazonaws.com/", + Settings: codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + ExternalID: ptr.Ref("shared-secret"), + }, + }, + }) + require.NoError(t, err) + + _, err = client.UpdateAIProvider(ctx, created.Name, codersdk.UpdateAIProviderRequest{ + Settings: &codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }, + }, + }) + require.NoError(t, err) + + //nolint:gocritic // Test reads the row to verify write-only fields. + row, err := db.GetAIProviderByID(dbauthz.AsSystemRestricted(ctx), created.ID) + require.NoError(t, err) + persisted, err := db2sdk.AIProviderSettings(row.Settings) + require.NoError(t, err) + require.NotNil(t, persisted.Bedrock) + require.NotNil(t, persisted.Bedrock.ExternalID) + require.Equal(t, "shared-secret", *persisted.Bedrock.ExternalID) + }) } diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index 8b29b2653aaa1..1dafd8377f508 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -26,6 +26,17 @@ import ( "github.com/coder/coder/v2/testutil" ) +// mustNewAnthropicProvider builds an Anthropic provider for tests, panicking if +// credential resolution fails. Keeps call sites terse after NewAnthropicProvider +// gained a context and error return. +func mustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { + p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) + if err != nil { + panic("build anthropic provider: " + err.Error()) + } + return p +} + func newTestServer(t *testing.T) (*aibridged.Server, *mock.MockDRPCClient, *mock.MockPooler) { t.Helper() @@ -648,7 +659,7 @@ func TestServeHTTP_ActorHeaders(t *testing.T) { BaseURL: upstreamSrv.URL, SendActorHeaders: true, }), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + mustNewAnthropicProvider(aibridge.AnthropicConfig{ BaseURL: upstreamSrv.URL, SendActorHeaders: true, }, nil), @@ -754,7 +765,7 @@ func TestRouting(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL}), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL}, nil), + mustNewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL}, nil), } pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, nil, testTracer) require.NoError(t, err) diff --git a/coderd/apidoc/docs.go b/coderd/apidoc/docs.go index 24eaec094114b..d11dba579304c 100644 --- a/coderd/apidoc/docs.go +++ b/coderd/apidoc/docs.go @@ -15218,6 +15218,13 @@ const docTemplate = `{ "bedrock_region": { "type": "string" }, + "bedrock_role_arn": { + "description": "BedrockRoleARN, when set, is the IAM role assumed via STS before\ncalling Bedrock, enabling cross-account access. BedrockExternalID is\nthe optional external ID sent on the AssumeRole call (write-only).\nBedrockSessionName is the STS session name, auto-generated when empty.", + "type": "string" + }, + "bedrock_session_name": { + "type": "string" + }, "bedrock_small_fast_model": { "type": "string" }, diff --git a/coderd/apidoc/swagger.json b/coderd/apidoc/swagger.json index 26c4aff908ff3..5e91d75dbc364 100644 --- a/coderd/apidoc/swagger.json +++ b/coderd/apidoc/swagger.json @@ -13574,6 +13574,13 @@ "bedrock_region": { "type": "string" }, + "bedrock_role_arn": { + "description": "BedrockRoleARN, when set, is the IAM role assumed via STS before\ncalling Bedrock, enabling cross-account access. BedrockExternalID is\nthe optional external ID sent on the AssumeRole call (write-only).\nBedrockSessionName is the STS session name, auto-generated when empty.", + "type": "string" + }, + "bedrock_session_name": { + "type": "string" + }, "bedrock_small_fast_model": { "type": "string" }, diff --git a/coderd/database/db2sdk/db2sdk.go b/coderd/database/db2sdk/db2sdk.go index b2ccca09d6da8..b82b212a6bb35 100644 --- a/coderd/database/db2sdk/db2sdk.go +++ b/coderd/database/db2sdk/db2sdk.go @@ -112,6 +112,7 @@ func redactAIProviderSettings(s codersdk.AIProviderSettings) codersdk.AIProvider b := *out.Bedrock b.AccessKey = nil b.AccessKeySecret = nil + b.ExternalID = nil out.Bedrock = &b } return out diff --git a/codersdk/aiproviders_bedrock.go b/codersdk/aiproviders_bedrock.go index 88edcb0017ba7..6067f0a4070b6 100644 --- a/codersdk/aiproviders_bedrock.go +++ b/codersdk/aiproviders_bedrock.go @@ -30,6 +30,20 @@ type AIProviderBedrockSettings struct { // AccessKeySecret is the AWS secret access key paired with // AccessKey. Write-only. AccessKeySecret *string `json:"access_key_secret,omitempty"` + // RoleARN, when set, is the IAM role assumed via STS before calling + // Bedrock. The base identity (static keys or the AWS environment, e.g. + // IRSA / instance profile) signs the AssumeRole call, and the resulting + // temporary credentials sign Bedrock requests. Enables cross-account + // access so usage bills to the target account. Non-secret. + RoleARN string `json:"role_arn,omitempty"` + // SessionName is the STS role session name used when assuming RoleARN. + // Auto-generated when empty. Non-secret. + SessionName string `json:"session_name,omitempty"` + // ExternalID is sent on the AssumeRole call for confused-deputy + // protection when assuming a role in an account the deployment does not + // own. Write-only (servers strip it from GET and list responses); a + // pointer for the same PATCH-omit semantics as AccessKeySecret. + ExternalID *string `json:"external_id,omitempty"` } // IsConfigured reports whether any load-bearing Bedrock field is set, @@ -47,6 +61,9 @@ func (b AIProviderBedrockSettings) IsConfigured() bool { if b.Region != "" { return true } + if b.RoleARN != "" { + return true + } if b.AccessKey != nil && *b.AccessKey != "" { return true } diff --git a/codersdk/deployment.go b/codersdk/deployment.go index 392270235beb9..f911e5db45815 100644 --- a/codersdk/deployment.go +++ b/codersdk/deployment.go @@ -4870,6 +4870,13 @@ type AIProviderConfig struct { BedrockAccessKeySecrets []string `json:"-"` BedrockModel string `json:"bedrock_model,omitempty"` BedrockSmallFastModel string `json:"bedrock_small_fast_model,omitempty"` + // BedrockRoleARN, when set, is the IAM role assumed via STS before + // calling Bedrock, enabling cross-account access. BedrockExternalID is + // the optional external ID sent on the AssumeRole call (write-only). + // BedrockSessionName is the STS session name, auto-generated when empty. + BedrockRoleARN string `json:"bedrock_role_arn,omitempty"` + BedrockExternalID string `json:"-"` + BedrockSessionName string `json:"bedrock_session_name,omitempty"` } type AIBridgeProxyConfig struct { diff --git a/docs/reference/api/general.md b/docs/reference/api/general.md index 4ae0f8d73b4dd..28c967e579c30 100644 --- a/docs/reference/api/general.md +++ b/docs/reference/api/general.md @@ -213,6 +213,8 @@ curl -X GET http://coder-server:8080/api/v2/deployment/config \ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", + "bedrock_role_arn": "string", + "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" diff --git a/docs/reference/api/schemas.md b/docs/reference/api/schemas.md index 8998babc87da7..b86b465129cb1 100644 --- a/docs/reference/api/schemas.md +++ b/docs/reference/api/schemas.md @@ -419,6 +419,8 @@ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", + "bedrock_role_arn": "string", + "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -952,6 +954,8 @@ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", + "bedrock_role_arn": "string", + "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -1044,6 +1048,8 @@ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", + "bedrock_role_arn": "string", + "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -1052,14 +1058,16 @@ ### Properties -| Name | Type | Required | Restrictions | Description | -|----------------------------|--------|----------|--------------|-------------------------------------------------------------------------------------------------------------------------------------------------------| -| `base_url` | string | false | | Base URL is the base URL of the upstream provider API. | -| `bedrock_model` | string | false | | | -| `bedrock_region` | string | false | | | -| `bedrock_small_fast_model` | string | false | | | -| `name` | string | false | | Name is the unique instance identifier used for routing. Defaults to Type if not provided. | -| `type` | string | false | | Type is the provider type. Valid values are: "openai", "anthropic", "azure", "bedrock", "google", "openai-compat", "openrouter", "vercel", "copilot". | +| Name | Type | Required | Restrictions | Description | +|----------------------------|--------|----------|--------------|----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| `base_url` | string | false | | Base URL is the base URL of the upstream provider API. | +| `bedrock_model` | string | false | | | +| `bedrock_region` | string | false | | | +| `bedrock_role_arn` | string | false | | Bedrock role arn when set, is the IAM role assumed via STS before calling Bedrock, enabling cross-account access. BedrockExternalID is the optional external ID sent on the AssumeRole call (write-only). BedrockSessionName is the STS session name, auto-generated when empty. | +| `bedrock_session_name` | string | false | | | +| `bedrock_small_fast_model` | string | false | | | +| `name` | string | false | | Name is the unique instance identifier used for routing. Defaults to Type if not provided. | +| `type` | string | false | | Type is the provider type. Valid values are: "openai", "anthropic", "azure", "bedrock", "google", "openai-compat", "openrouter", "vercel", "copilot". | ## codersdk.AIProviderKey @@ -5516,6 +5524,8 @@ CreateWorkspaceRequest provides options for creating a new workspace. Only one o "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", + "bedrock_role_arn": "string", + "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -6118,6 +6128,8 @@ CreateWorkspaceRequest provides options for creating a new workspace. Only one o "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", + "bedrock_role_arn": "string", + "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 5d907f0726492..9cacfae02baf5 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -37,6 +37,17 @@ import ( var testTracer = otel.Tracer("aibridged_inttest") +// mustNewAnthropicProvider builds an Anthropic provider for tests, panicking if +// credential resolution fails. Keeps call sites terse after NewAnthropicProvider +// gained a context and error return. +func mustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { + p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) + if err != nil { + panic("build anthropic provider: " + err.Error()) + } + return p +} + // TestIntegration is not an exhaustive test against the upstream AI providers' SDKs (see coder/aibridge for those). // This test validates that: // - intercepted requests can be authenticated/authorized @@ -493,7 +504,7 @@ func TestIntegrationCircuitBreaker(t *testing.T) { BaseURL: mockOpenAI.URL, CircuitBreaker: cbConfig, }), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + mustNewAnthropicProvider(aibridge.AnthropicConfig{ BaseURL: mockAnthropic.URL, Key: "test-key", CircuitBreaker: cbConfig, diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 9b580dfec6cc7..ce09fbd1470ac 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -309,6 +309,26 @@ export interface AIProviderBedrockSettings { * AccessKey. Write-only. */ readonly access_key_secret?: string; + /** + * RoleARN, when set, is the IAM role assumed via STS before calling + * Bedrock. The base identity (static keys or the AWS environment, e.g. + * IRSA / instance profile) signs the AssumeRole call, and the resulting + * temporary credentials sign Bedrock requests. Enables cross-account + * access so usage bills to the target account. Non-secret. + */ + readonly role_arn?: string; + /** + * SessionName is the STS role session name used when assuming RoleARN. + * Auto-generated when empty. Non-secret. + */ + readonly session_name?: string; + /** + * ExternalID is sent on the AssumeRole call for confused-deputy + * protection when assuming a role in an account the deployment does not + * own. Write-only (servers strip it from GET and list responses); a + * pointer for the same PATCH-omit semantics as AccessKeySecret. + */ + readonly external_id?: string; } // From codersdk/aiproviders_bedrock.go @@ -344,6 +364,14 @@ export interface AIProviderConfig { readonly bedrock_region?: string; readonly bedrock_model?: string; readonly bedrock_small_fast_model?: string; + /** + * BedrockRoleARN, when set, is the IAM role assumed via STS before + * calling Bedrock, enabling cross-account access. BedrockExternalID is + * the optional external ID sent on the AssumeRole call (write-only). + * BedrockSessionName is the STS session name, auto-generated when empty. + */ + readonly bedrock_role_arn?: string; + readonly bedrock_session_name?: string; } // From codersdk/aiproviders.go From 65cb77c38bb61ac50f901f92146fe83049ca3af4 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Sat, 20 Jun 2026 18:07:52 +0000 Subject: [PATCH 02/21] feat: support cross-account Bedrock AssumeRole in AI Bridge --- aibridge/bridge_test.go | 3 - aibridge/config/config.go | 10 +-- aibridge/intercept/credential.go | 2 +- aibridge/intercept/credential_test.go | 2 +- aibridge/intercept/keyfailover_test.go | 4 +- aibridge/intercept/messages/base.go | 59 +++++++------- .../intercept/messages/base_internal_test.go | 55 ++++++------- aibridge/intercept/messages/blocking.go | 8 +- aibridge/intercept/messages/streaming.go | 8 +- .../integrationtest/bridge_internal_test.go | 15 +--- .../internal/integrationtest/setupbridge.go | 3 - aibridge/provider/anthropic.go | 40 ++++------ aibridge/provider/anthropic_internal_test.go | 6 +- aibridge/provider/bedrock.go | 39 ++++------ aibridge/provider/bedrock_internal_test.go | 78 ++++++------------- cli/aibridged.go | 3 - cli/server.go | 12 --- cli/server_aibridge_internal_test.go | 25 ------ coderd/ai_providers.go | 13 ++-- coderd/ai_providers_migrate.go | 6 -- coderd/ai_providers_test.go | 61 +-------------- coderd/aibridged/aibridged_test.go | 3 - coderd/apidoc/docs.go | 7 -- coderd/apidoc/swagger.json | 7 -- coderd/database/db2sdk/db2sdk.go | 1 - codersdk/aiproviders_bedrock.go | 13 +--- codersdk/deployment.go | 7 -- docs/reference/api/general.md | 2 - docs/reference/api/schemas.md | 28 ++----- enterprise/aibridged_integration_test.go | 3 - go.mod | 2 +- site/src/api/typesGenerated.ts | 25 +----- .../ProvidersPage/components/ProviderForm.tsx | 13 ++++ .../components/providerFormApiMap.test.ts | 3 + .../components/providerFormApiMap.ts | 5 ++ 35 files changed, 165 insertions(+), 406 deletions(-) diff --git a/aibridge/bridge_test.go b/aibridge/bridge_test.go index a51779c2335c8..6ee283d3e9601 100644 --- a/aibridge/bridge_test.go +++ b/aibridge/bridge_test.go @@ -23,9 +23,6 @@ import ( var bridgeTestTracer = otel.Tracer("bridge_test") -// mustNewAnthropicProvider builds an Anthropic provider for tests, panicking if -// credential resolution fails. Keeps call sites terse after NewAnthropicProvider -// gained a context and error return. func mustNewAnthropicProvider(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) aibridge.Provider { p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) if err != nil { diff --git a/aibridge/config/config.go b/aibridge/config/config.go index 3696e88c6a2d1..3dc76841aa3e1 100644 --- a/aibridge/config/config.go +++ b/aibridge/config/config.go @@ -35,15 +35,9 @@ type AWSBedrock struct { BaseURL string // RoleARN, when set, is assumed via STS before calling Bedrock. The base // identity (static keys or the AWS SDK default credential chain, e.g. - // IRSA / Pod Identity / instance profile) signs the AssumeRole call, and - // the resulting temporary credentials sign Bedrock requests. This enables - // cross-account access. + // IRSA / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + // call, and the resulting temporary credentials sign Bedrock requests. RoleARN string - // ExternalID is sent on the AssumeRole call for confused-deputy protection - // when assuming a role in an account this gateway does not own. Optional. - ExternalID string - // SessionName is the STS role session name. Auto-generated when empty. - SessionName string } // OpenAI carries configuration for an OpenAI provider. diff --git a/aibridge/intercept/credential.go b/aibridge/intercept/credential.go index 008c43463df01..24016b69504a8 100644 --- a/aibridge/intercept/credential.go +++ b/aibridge/intercept/credential.go @@ -28,7 +28,7 @@ const ( // before failover selects a key, and a key resolved dynamically at request time. const ( hintFailoverKey = "" - hintBedrockChainKey = "" + hintBedrockChainKey = "" ) // Credential is the per-request upstream authentication for an interception: diff --git a/aibridge/intercept/credential_test.go b/aibridge/intercept/credential_test.go index 0f78d8dee04a8..6c147380aba5c 100644 --- a/aibridge/intercept/credential_test.go +++ b/aibridge/intercept/credential_test.go @@ -72,7 +72,7 @@ func TestCredential(t *testing.T) { }, expectKind: intercept.CredentialKindCentralized, expectAuthHeader: "", - expectHint: "", + expectHint: "", expectLength: 0, }, { diff --git a/aibridge/intercept/keyfailover_test.go b/aibridge/intercept/keyfailover_test.go index c786cf57dae42..700e3429e7f77 100644 --- a/aibridge/intercept/keyfailover_test.go +++ b/aibridge/intercept/keyfailover_test.go @@ -95,9 +95,9 @@ var interceptorCases = []interceptorCase{ id, tracer := uuid.New(), otel.Tracer("keyfailover") if streaming { - return messages.NewStreamingInterceptor(id, payload, cfg, cred, nil, nil, http.Header{}, tracer) + return messages.NewStreamingInterceptor(id, payload, cfg, cred, nil, http.Header{}, tracer) } - return messages.NewBlockingInterceptor(id, payload, cfg, cred, nil, nil, http.Header{}, tracer) + return messages.NewBlockingInterceptor(id, payload, cfg, cred, nil, http.Header{}, tracer) }, }, { diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index 8198fc58079e2..dd11ec3e74438 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -64,16 +64,21 @@ var bedrockSupportedBetaFlags = map[string]bool{ "tool-examples-2025-10-29": true, } +// BedrockRuntime carries everything a Bedrock-backed interception needs: the +// static Bedrock config plus the AWS credentials provider. +type BedrockRuntime struct { + Cfg aibconfig.AWSBedrock + Creds aws.CredentialsProvider +} + type interceptionBase struct { id uuid.UUID reqPayload RequestPayload - cfg intercept.Config - cred intercept.Credential - bedrockCfg *aibconfig.AWSBedrock - // bedrockCreds is the cached AWS credentials provider resolved once per - // provider (including any assumed role). nil when not Bedrock-backed. - bedrockCreds aws.CredentialsProvider + cfg intercept.Config + cred intercept.Credential + // bedrock is nil for non-Bedrock providers. + bedrock *BedrockRuntime // clientHeaders are the original HTTP headers from the client request. clientHeaders http.Header @@ -109,10 +114,10 @@ func (i *interceptionBase) Model() string { return "coder-aibridge-unknown" } - if i.bedrockCfg != nil { - model := i.bedrockCfg.Model + if i.bedrock != nil { + model := i.bedrock.Cfg.Model if i.isSmallFastModel() { - model = i.bedrockCfg.SmallFastModel + model = i.bedrock.Cfg.SmallFastModel } return model } @@ -128,7 +133,7 @@ func (i *interceptionBase) baseTraceAttributes(r *http.Request, streaming bool) 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), + attribute.Bool(tracing.IsBedrock, i.bedrock != nil), } } @@ -240,8 +245,10 @@ func (i *interceptionBase) newMessagesService(ctx context.Context, opts ...optio opts = append(opts, option.WithMiddleware(mw)) } - if i.bedrockCfg != nil { - bedrockOpts, err := i.withAWSBedrockOptions(ctx, i.bedrockCfg) + if i.bedrock != nil { + ctx, cancel := context.WithTimeout(ctx, time.Second*30) + defer cancel() + bedrockOpts, err := i.withAWSBedrockOptions(ctx) if err != nil { return anthropic.MessageService{}, err } @@ -262,17 +269,13 @@ func (i *interceptionBase) withBody() option.RequestOption { // withAWSBedrockOptions returns request options for authenticating with AWS Bedrock. // -// Credentials are resolved once per provider and supplied via i.bedrockCreds: -// either static credentials, the AWS SDK default credential chain (environment -// variables, shared config/credentials files, IAM roles, IRSA, SSO, IMDS, etc.), -// or a role assumed via STS when a target role ARN is configured. i.bedrockCreds -// is a shared, rotating credentials cache, so the per-request Retrieve below is -// served from that cache (a cache hit) and does not re-resolve or re-assume on -// every request; it only fails fast if credentials cannot be resolved. -func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconfig.AWSBedrock) ([]option.RequestOption, error) { - if cfg == nil { - return nil, xerrors.New("nil config given") +// Credentials come from i.bedrock.Creds, it is a shared credentials cache, so the per-request +// Retrieve below is served from that cache and does not re-resolve or re-assume on every request. +func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context) ([]option.RequestOption, error) { + if i.bedrock == nil { + return nil, xerrors.New("nil bedrock runtime") } + cfg := i.bedrock.Cfg if cfg.Region == "" && cfg.BaseURL == "" { return nil, xerrors.New("region or base url required") } @@ -282,19 +285,17 @@ func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibco if cfg.SmallFastModel == "" { return nil, xerrors.New("small fast model required") } - if i.bedrockCreds == nil { - return nil, xerrors.New("bedrock credentials not resolved") - } // Fail fast: ensure credentials can be resolved before signing. Served from - // the shared cache, so this does not re-assume the role on each request. - if _, err := i.bedrockCreds.Retrieve(ctx); err != nil { + // the shared cache on most requests (no network); on the cold or refresh + // path this performs the actual STS/IMDS call. + if _, err := i.bedrock.Creds.Retrieve(ctx); err != nil { return nil, xerrors.Errorf("no AWS credentials found: %w", err) } awsCfg := aws.Config{ Region: cfg.Region, - Credentials: i.bedrockCreds, + Credentials: i.bedrock.Creds, } var out []option.RequestOption @@ -319,7 +320,7 @@ func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibco // don't support adaptive thinking natively, or enabled thinking to adaptive for models that only support // adaptive (Opus 4.7+). func (i *interceptionBase) augmentRequestForBedrock() { - if i.bedrockCfg == nil { + if i.bedrock == nil { return } diff --git a/aibridge/intercept/messages/base_internal_test.go b/aibridge/intercept/messages/base_internal_test.go index ad2be1a9b1159..0b2ee4d77be84 100644 --- a/aibridge/intercept/messages/base_internal_test.go +++ b/aibridge/intercept/messages/base_internal_test.go @@ -88,14 +88,14 @@ func TestAWSBedrockValidation(t *testing.T) { tests := []struct { name string - cfg *config.AWSBedrock + cfg config.AWSBedrock expectError bool errorMsg string }{ // Valid cases: static credentials. { name: "static credentials with region", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -105,7 +105,7 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "static credentials with base url", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ BaseURL: "http://bedrock.internal", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -120,7 +120,7 @@ func TestAWSBedrockValidation(t *testing.T) { // // See TestAWSBedrockIntegration which validates this. name: "static credentials with base url & region", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -131,7 +131,7 @@ func TestAWSBedrockValidation(t *testing.T) { // Invalid cases. { name: "missing region & base url", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -143,7 +143,7 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "missing model", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -155,7 +155,7 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "missing small fast model", - cfg: &config.AWSBedrock{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -167,29 +167,23 @@ func TestAWSBedrockValidation(t *testing.T) { }, { name: "all fields empty", - cfg: &config.AWSBedrock{}, + cfg: config.AWSBedrock{}, expectError: true, errorMsg: "region or base url required", }, - { - name: "nil config", - cfg: nil, - expectError: true, - errorMsg: "nil config given", - }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - // Credentials are resolved once per provider and supplied to the - // interceptor; the per-request path only assembles options around - // them. A static stub stands in for the resolved provider. base := &interceptionBase{ - bedrockCreds: credentials.NewStaticCredentialsProvider("test-key", "test-secret", ""), + bedrock: &BedrockRuntime{ + Cfg: tt.cfg, + Creds: credentials.NewStaticCredentialsProvider("test-key", "test-secret", ""), + }, } - opts, err := base.withAWSBedrockOptions(context.Background(), tt.cfg) + opts, err := base.withAWSBedrockOptions(context.Background()) if tt.expectError { require.Error(t, err) @@ -202,19 +196,16 @@ func TestAWSBedrockValidation(t *testing.T) { } } -// TestAWSBedrockOptionsRequireResolvedCredentials verifies that option assembly -// fails when the per-provider credentials provider was not resolved. -func TestAWSBedrockOptionsRequireResolvedCredentials(t *testing.T) { +// TestAWSBedrockOptionsRequireRuntime verifies that option assembly fails when +// the Bedrock runtime was not set. This should never happen in practice, since +// withAWSBedrockOptions is only called when i.bedrock != nil. +func TestAWSBedrockOptionsRequireRuntime(t *testing.T) { t.Parallel() base := &interceptionBase{} - _, err := base.withAWSBedrockOptions(context.Background(), &config.AWSBedrock{ - Region: "us-east-1", - Model: "test-model", - SmallFastModel: "test-small-model", - }) + _, err := base.withAWSBedrockOptions(context.Background()) require.Error(t, err) - require.Contains(t, err.Error(), "bedrock credentials not resolved") + require.Contains(t, err.Error(), "nil bedrock runtime") } func TestAccumulateUsage(t *testing.T) { @@ -793,9 +784,11 @@ func TestAugmentRequestForBedrock_AdaptiveThinking(t *testing.T) { i := &interceptionBase{ reqPayload: mustMessagesPayload(t, tc.requestBody), - bedrockCfg: &config.AWSBedrock{ - Model: tc.bedrockModel, - SmallFastModel: "anthropic.claude-haiku-3-5", + bedrock: &BedrockRuntime{ + Cfg: config.AWSBedrock{ + Model: tc.bedrockModel, + SmallFastModel: "anthropic.claude-haiku-3-5", + }, }, clientHeaders: clientHeaders, logger: slog.Make(), diff --git a/aibridge/intercept/messages/blocking.go b/aibridge/intercept/messages/blocking.go index 0d65e759c1778..8b3d9cc6138cb 100644 --- a/aibridge/intercept/messages/blocking.go +++ b/aibridge/intercept/messages/blocking.go @@ -9,7 +9,6 @@ import ( "github.com/anthropics/anthropic-sdk-go" "github.com/anthropics/anthropic-sdk-go/option" - "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" mcplib "github.com/mark3labs/mcp-go/mcp" "github.com/tidwall/sjson" @@ -18,7 +17,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - 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" @@ -37,8 +35,7 @@ func NewBlockingInterceptor( reqPayload RequestPayload, cfg intercept.Config, cred intercept.Credential, - bedrockCfg *aibconfig.AWSBedrock, - bedrockCreds aws.CredentialsProvider, + bedrock *BedrockRuntime, clientHeaders http.Header, tracer trace.Tracer, ) *BlockingInterception { @@ -47,8 +44,7 @@ func NewBlockingInterceptor( reqPayload: reqPayload, cfg: cfg, cred: cred, - bedrockCfg: bedrockCfg, - bedrockCreds: bedrockCreds, + bedrock: bedrock, clientHeaders: clientHeaders, tracer: tracer, }} diff --git a/aibridge/intercept/messages/streaming.go b/aibridge/intercept/messages/streaming.go index f7c849b225b2b..2369c17654975 100644 --- a/aibridge/intercept/messages/streaming.go +++ b/aibridge/intercept/messages/streaming.go @@ -13,7 +13,6 @@ import ( "github.com/anthropics/anthropic-sdk-go/option" "github.com/anthropics/anthropic-sdk-go/packages/ssestream" "github.com/anthropics/anthropic-sdk-go/shared/constant" - "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" mcplib "github.com/mark3labs/mcp-go/mcp" "github.com/tidwall/sjson" @@ -22,7 +21,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - 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" @@ -42,8 +40,7 @@ func NewStreamingInterceptor( reqPayload RequestPayload, cfg intercept.Config, cred intercept.Credential, - bedrockCfg *aibconfig.AWSBedrock, - bedrockCreds aws.CredentialsProvider, + bedrock *BedrockRuntime, clientHeaders http.Header, tracer trace.Tracer, ) *StreamingInterception { @@ -52,8 +49,7 @@ func NewStreamingInterceptor( reqPayload: reqPayload, cfg: cfg, cred: cred, - bedrockCfg: bedrockCfg, - bedrockCreds: bedrockCreds, + bedrock: bedrock, clientHeaders: clientHeaders, tracer: tracer, }} diff --git a/aibridge/internal/integrationtest/bridge_internal_test.go b/aibridge/internal/integrationtest/bridge_internal_test.go index 62d9cd36c9ac0..2609a942d5f5d 100644 --- a/aibridge/internal/integrationtest/bridge_internal_test.go +++ b/aibridge/internal/integrationtest/bridge_internal_test.go @@ -316,19 +316,8 @@ func TestAWSBedrockIntegration(t *testing.T) { SmallFastModel: "test-haiku", } - bridgeServer := newBridgeTestServer(ctx, t, "http://unused", - withCustomProvider(mustNewAnthropic(anthropicCfg("http://unused", apiKey), bedrockCfg)), - ) - - resp, err := bridgeServer.makeRequest(t, http.MethodPost, pathAnthropicMessages, fixtures.Request(t, fixtures.AntSingleBuiltinTool)) - require.NoError(t, err) - defer resp.Body.Close() - - require.Equal(t, http.StatusInternalServerError, resp.StatusCode) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - require.Contains(t, string(body), "create anthropic client") - require.Contains(t, string(body), "region or base url required") + _, err := provider.NewAnthropic(ctx, anthropicCfg("http://unused", apiKey), bedrockCfg) + require.ErrorContains(t, err, "region or base url required") }) t.Run("/v1/messages", func(t *testing.T) { diff --git a/aibridge/internal/integrationtest/setupbridge.go b/aibridge/internal/integrationtest/setupbridge.go index d0a05abc0a300..f8121589ee36a 100644 --- a/aibridge/internal/integrationtest/setupbridge.go +++ b/aibridge/internal/integrationtest/setupbridge.go @@ -249,9 +249,6 @@ func setupInjectedToolTest( return bridgeServer, mockMCP, resp } -// mustNewAnthropic builds an Anthropic provider for tests, panicking if -// credential resolution fails. NewAnthropic resolves Bedrock credentials at -// construction, so this keeps the many test call sites terse. func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) aibridge.Provider { p, err := provider.NewAnthropic(context.Background(), cfg, bedrockCfg) if err != nil { diff --git a/aibridge/provider/anthropic.go b/aibridge/provider/anthropic.go index 6b6c4304f1af2..d83af7786ab18 100644 --- a/aibridge/provider/anthropic.go +++ b/aibridge/provider/anthropic.go @@ -7,7 +7,6 @@ import ( "net/http" "strings" - "github.com/aws/aws-sdk-go-v2/aws" "github.com/google/uuid" "go.opentelemetry.io/otel/codes" "go.opentelemetry.io/otel/trace" @@ -27,13 +26,9 @@ var _ Provider = &Anthropic{} // Anthropic allows for interactions with the Anthropic API. type Anthropic struct { - cfg config.Anthropic - bedrockCfg *config.AWSBedrock - // bedrockCreds is the AWS credentials provider (including any assumed - // role), resolved once at construction and shared across requests so - // per-request retrieval is served from its cache rather than re-resolved. - // nil when this provider is not Bedrock-backed. - bedrockCreds aws.CredentialsProvider + cfg config.Anthropic + // bedrock is nil for non-Bedrock providers. + bedrock *messages.BedrockRuntime } const routeMessages = "/v1/messages" // https://docs.anthropic.com/en/api/messages @@ -62,23 +57,22 @@ func NewAnthropic(ctx context.Context, cfg config.Anthropic, bedrockCfg *config. cfg.CircuitBreaker.OpenErrorResponse = anthropicOpenErrorResponse } - // Resolve the AWS credentials provider once. This performs no network - // call (the base identity and any AssumeRole resolve lazily on first - // retrieval); it only wires up the provider chain, so it is cheap to run - // at construction and on every provider reload. - var bedrockCreds aws.CredentialsProvider + // Resolve the AWS credentials provider once and bundle it with the config. + // This performs no network call (the base identity and any AssumeRole + // resolve lazily on first retrieval); it only wires up the provider chain, + // so it is cheap to run at construction. + var bedrock *messages.BedrockRuntime if bedrockCfg != nil { - var err error - bedrockCreds, err = buildBedrockCredentials(ctx, *bedrockCfg) + creds, err := buildBedrockCredentials(ctx, *bedrockCfg) if err != nil { return nil, xerrors.Errorf("build bedrock credentials: %w", err) } + bedrock = &messages.BedrockRuntime{Cfg: *bedrockCfg, Creds: creds} } return &Anthropic{ - cfg: cfg, - bedrockCfg: bedrockCfg, - bedrockCreds: bedrockCreds, + cfg: cfg, + bedrock: bedrock, }, nil } @@ -142,13 +136,11 @@ func (p *Anthropic) CreateInterceptor(_ http.ResponseWriter, r *http.Request, tr return nil, xerrors.Errorf("resolve credential: %w", err) } - // bedrockCreds was resolved once at construction; it is a shared, - // rotating credentials cache. nil for non-Bedrock providers. var interceptor intercept.Interceptor if reqPayload.Stream() { - interceptor = messages.NewStreamingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, p.bedrockCreds, r.Header, tracer) + interceptor = messages.NewStreamingInterceptor(id, reqPayload, cfg, cred, p.bedrock, r.Header, tracer) } else { - interceptor = messages.NewBlockingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, p.bedrockCreds, r.Header, tracer) + interceptor = messages.NewBlockingInterceptor(id, reqPayload, cfg, cred, p.bedrock, r.Header, tracer) } span.SetAttributes(interceptor.TraceAttributes(r)...) return interceptor, nil @@ -176,8 +168,8 @@ func (p *Anthropic) resolveCredential(r *http.Request) (intercept.Credential, er 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 + if p.bedrock != nil { + return intercept.Bedrock{AccessKey: p.bedrock.Cfg.AccessKey}, nil } return nil, ErrNoCredential } diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 7ef3ba72e14b4..8427a31671c91 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -19,10 +19,6 @@ import ( "github.com/coder/quartz" ) -// mustNewAnthropic builds an Anthropic provider for tests, panicking if -// credential resolution fails (it cannot for the non-Bedrock configs used -// here). Keeps the call sites terse after NewAnthropic gained a context and -// error return. func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) if err != nil { @@ -299,7 +295,7 @@ func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { bedrock: true, setHeaders: map[string]string{}, wantCredentialKind: intercept.CredentialKindCentralized, - wantCredentialHint: "", + wantCredentialHint: "", }, { // Bedrock static mode: the hint masks the access key ID. diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go index c77b1b904cd28..d7418348d09c6 100644 --- a/aibridge/provider/bedrock.go +++ b/aibridge/provider/bedrock.go @@ -13,22 +13,22 @@ import ( "github.com/coder/coder/v2/aibridge/config" ) -// defaultBedrockSessionName is the STS role session name used when the provider -// does not configure one. A stable value keeps AssumeRole calls identifiable in -// CloudTrail. -const defaultBedrockSessionName = "coder-aibridge" +// bedrockSessionName is the STS role session name attached to AssumeRole calls. +// A stable value keeps them identifiable in CloudTrail. +const bedrockSessionName = "coder-aigateway" -// buildBedrockCredentials resolves the base identity (static keys or the AWS SDK -// default credential chain, which covers IRSA, Pod Identity, instance profile, -// shared profile, and environment variables) and, when a target role ARN is -// configured, assumes that role via STS. +// buildBedrockCredentials resolves the base identity and, when a role ARN +// is configured, assumes that role via STS. The base identity is either +// static keys or the AWS SDK default credential chain, which covers IRSA, +// EKS Pod Identity, EC2 Instance Profile, and more. // -// The result is wrapped in aws.NewCredentialsCache, which caches and rotates the -// resolved temporary credentials. The provider is resolved once per Bedrock -// provider (at construction) and shared across requests, so per-request -// credential retrieval is served from this cache rather than re-resolving (and -// re-assuming) on every request. No network call is made here: the base -// identity and any AssumeRole are resolved lazily on first retrieval. +// The result is wrapped in aws.NewCredentialsCache, which caches and rotates +// the resolved temporary credentials. buildBedrockCredentials should be called +// once when the Bedrock provider is constructed, and the returned Credential +// Provider should be shared across all LLM requests to the Bedrock Provider, +// so per-request credential retrieval is served from this cache rather than +// re-resolving (and re-assuming) on every request. No network call is made here: +// the base identity and any AssumeRole are resolved lazily on first retrieval. func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.CredentialsProvider, error) { if cfg.Region == "" && cfg.BaseURL == "" { return nil, xerrors.New("region or base url required") @@ -63,20 +63,13 @@ func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.Cr return nil, xerrors.Errorf("failed to load AWS Bedrock config: %w", err) } - // The base identity signs requests directly unless a target role is + // The base identity signs requests directly unless a target RoleARN is // configured, in which case it signs the AssumeRole call and the resulting // temporary credentials sign Bedrock requests. credsProvider := base.Credentials if cfg.RoleARN != "" { - sessionName := cfg.SessionName - if sessionName == "" { - sessionName = defaultBedrockSessionName - } credsProvider = stscreds.NewAssumeRoleProvider(sts.NewFromConfig(base), cfg.RoleARN, func(o *stscreds.AssumeRoleOptions) { - o.RoleSessionName = sessionName - if cfg.ExternalID != "" { - o.ExternalID = aws.String(cfg.ExternalID) - } + o.RoleSessionName = bedrockSessionName }) } diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index 4cba363e8093a..5e1c1c5fd091d 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -55,8 +55,7 @@ func TestBuildBedrockCredentialsValidation(t *testing.T) { } } -// TestBuildBedrockCredentialsStatic resolves static credentials with no network -// access. +// TestBuildBedrockCredentialsStatic resolves static credentials. func TestBuildBedrockCredentialsStatic(t *testing.T) { t.Parallel() @@ -74,14 +73,16 @@ func TestBuildBedrockCredentialsStatic(t *testing.T) { } // TestBuildBedrockCredentialsDefaultChain covers resolution via the AWS SDK -// default credential chain (here, environment variables) and the failure mode -// when no credential source is configured. +// default credential chain. // NOTE: no t.Parallel() because the subtests use t.Setenv. func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { tests := []struct { name string envVars map[string]string expectError bool + wantKey string + wantSecret string + wantToken string }{ { name: "credentials via env", @@ -89,6 +90,8 @@ func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { "AWS_ACCESS_KEY_ID": "test-key", "AWS_SECRET_ACCESS_KEY": "test-secret", }, + wantKey: "test-key", + wantSecret: "test-secret", }, { name: "credentials with session token via env", @@ -97,6 +100,9 @@ func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { "AWS_SECRET_ACCESS_KEY": "test-secret", "AWS_SESSION_TOKEN": "test-session-token", }, + wantKey: "test-key", + wantSecret: "test-secret", + wantToken: "test-session-token", }, { name: "error when no credential source is configured", @@ -125,35 +131,37 @@ func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { t.Setenv(key, val) } - // buildBedrockCredentials wires up the provider chain without a - // network call, so it succeeds regardless of credential - // availability; resolution failures surface on Retrieve. + // buildBedrockCredentials only wires up the provider chain; it + // does not resolve credentials, so it succeeds regardless of + // credential availability. Resolution failures surface on Retrieve. creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ Region: "us-east-1", }) require.NoError(t, err) require.NotNil(t, creds) - _, err = creds.Retrieve(context.Background()) + got, err := creds.Retrieve(context.Background()) if tt.expectError { require.Error(t, err) return } require.NoError(t, err) + require.Equal(t, tt.wantKey, got.AccessKeyID) + require.Equal(t, tt.wantSecret, got.SecretAccessKey) + require.Equal(t, tt.wantToken, got.SessionToken) }) } } // TestBuildBedrockCredentialsAssumeRole drives the STS AssumeRole path against a -// mock endpoint, asserting that the configured role ARN, external ID, and -// session name are sent and that the returned temporary credentials are used. +// mock endpoint, asserting that the configured role ARN and the stable session +// name are sent and that the returned temporary credentials are used. // NOTE: no t.Parallel() because it uses t.Setenv. func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { - var gotRoleARN, gotExternalID, gotSessionName string + var gotRoleARN, gotSessionName string sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { require.NoError(t, r.ParseForm()) gotRoleARN = r.Form.Get("RoleArn") - gotExternalID = r.Form.Get("ExternalId") gotSessionName = r.Form.Get("RoleSessionName") w.Header().Set("Content-Type", "text/xml") @@ -175,16 +183,14 @@ func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { defer sts.Close() // Point the STS client at the mock and provide static base credentials so - // the base identity resolves without network access. + // the base identity resolves without additional network calls. t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) t.Setenv("AWS_ACCESS_KEY_ID", "base-key") t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ - Region: "us-east-1", - RoleARN: "arn:aws:iam::123456789012:role/target", - ExternalID: "shared-secret", - SessionName: "coder", + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", }) require.NoError(t, err) @@ -195,41 +201,5 @@ func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { require.Equal(t, "assumed-token", got.SessionToken) require.Equal(t, "arn:aws:iam::123456789012:role/target", gotRoleARN) - require.Equal(t, "shared-secret", gotExternalID) - require.Equal(t, "coder", gotSessionName) -} - -// TestBuildBedrockCredentialsAssumeRoleDefaultSessionName verifies a session -// name is sent even when the provider does not configure one. -func TestBuildBedrockCredentialsAssumeRoleDefaultSessionName(t *testing.T) { - var gotSessionName string - sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - require.NoError(t, r.ParseForm()) - gotSessionName = r.Form.Get("RoleSessionName") - w.Header().Set("Content-Type", "text/xml") - _, _ = w.Write([]byte(` - - - ASIAASSUMED - assumed-secret - assumed-token - 2999-01-01T00:00:00Z - - -`)) - })) - defer sts.Close() - - t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) - t.Setenv("AWS_ACCESS_KEY_ID", "base-key") - t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") - - creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ - Region: "us-east-1", - RoleARN: "arn:aws:iam::123456789012:role/target", - }) - require.NoError(t, err) - _, err = creds.Retrieve(context.Background()) - require.NoError(t, err) - require.Equal(t, defaultBedrockSessionName, gotSessionName) + require.Equal(t, bedrockSessionName, gotSessionName) } diff --git a/cli/aibridged.go b/cli/aibridged.go index f3f21c805de2e..9aa2ea5c27843 100644 --- a/cli/aibridged.go +++ b/cli/aibridged.go @@ -342,7 +342,6 @@ func bedrockConfigFromRow(row database.AIProvider, settings codersdk.AIProviderS } accessKey := ptr.NilToEmpty(bedrockSettings.AccessKey) accessKeySecret := ptr.NilToEmpty(bedrockSettings.AccessKeySecret) - externalID := ptr.NilToEmpty(bedrockSettings.ExternalID) return &aibridge.AWSBedrockConfig{ BaseURL: row.BaseUrl, Region: bedrockSettings.Region, @@ -351,8 +350,6 @@ func bedrockConfigFromRow(row database.AIProvider, settings codersdk.AIProviderS Model: bedrockSettings.Model, SmallFastModel: bedrockSettings.SmallFastModel, RoleARN: bedrockSettings.RoleARN, - ExternalID: externalID, - SessionName: bedrockSettings.SessionName, } } diff --git a/cli/server.go b/cli/server.go index ca14e83bdac1f..cbf4adaa41226 100644 --- a/cli/server.go +++ b/cli/server.go @@ -3131,12 +3131,6 @@ func ReadAIProvidersFromEnv(logger slog.Logger, environ []string) ([]codersdk.AI p.BedrockRegion, bedrockKey, bedrockSecret, p.BedrockModel, p.BedrockSmallFastModel, ) - settings.RoleARN = p.BedrockRoleARN - settings.SessionName = p.BedrockSessionName - if p.BedrockExternalID != "" { - externalID := p.BedrockExternalID - settings.ExternalID = &externalID - } isBedrock := codersdk.IsBedrockConfigured(p.BedrockBaseURL, settings) // BEDROCK_* fields are accepted on anthropic (mutually exclusive @@ -3299,12 +3293,6 @@ func readAIProvidersForPrefix(logger slog.Logger, environ []string, prefix strin provider.BedrockModel = v.Value case "BEDROCK_SMALL_FAST_MODEL": provider.BedrockSmallFastModel = v.Value - case "BEDROCK_ROLE_ARN": - provider.BedrockRoleARN = v.Value - case "BEDROCK_EXTERNAL_ID": - provider.BedrockExternalID = v.Value - case "BEDROCK_SESSION_NAME": - provider.BedrockSessionName = v.Value default: logger.Warn(context.Background(), "ignoring unknown AI provider field (check for typos)", slog.F("env", fullName), diff --git a/cli/server_aibridge_internal_test.go b/cli/server_aibridge_internal_test.go index e6be4287382ac..cb08ec530beef 100644 --- a/cli/server_aibridge_internal_test.go +++ b/cli/server_aibridge_internal_test.go @@ -131,31 +131,6 @@ func TestReadAIProvidersFromEnv(t *testing.T) { }, }, }, - { - name: "BedrockAssumeRoleFields", - env: []string{ - "CODER_AIBRIDGE_PROVIDER_0_TYPE=bedrock", - "CODER_AIBRIDGE_PROVIDER_0_NAME=bedrock-unit-a", - "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_REGION=us-east-1", - "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_MODEL=anthropic.claude-3-sonnet", - "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_SMALL_FAST_MODEL=anthropic.claude-3-haiku", - "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_ROLE_ARN=arn:aws:iam::123456789012:role/target", - "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_EXTERNAL_ID=shared-secret", - "CODER_AIBRIDGE_PROVIDER_0_BEDROCK_SESSION_NAME=coder", - }, - expected: []codersdk.AIProviderConfig{ - { - Type: string(database.AIProviderTypeBedrock), - Name: "bedrock-unit-a", - BedrockRegion: "us-east-1", - BedrockModel: "anthropic.claude-3-sonnet", - BedrockSmallFastModel: "anthropic.claude-3-haiku", - BedrockRoleARN: "arn:aws:iam::123456789012:role/target", - BedrockExternalID: "shared-secret", - BedrockSessionName: "coder", - }, - }, - }, { name: "OutOfOrderIndices", env: []string{ diff --git a/coderd/ai_providers.go b/coderd/ai_providers.go index bc83e88880185..49e5bb8750403 100644 --- a/coderd/ai_providers.go +++ b/coderd/ai_providers.go @@ -754,11 +754,11 @@ func encodeAIProviderSettings(s codersdk.AIProviderSettings) (sql.NullString, er } // mergeAIProviderSettings overlays a patch onto an existing settings -// value. Write-only fields (Bedrock AccessKey, AccessKeySecret, and -// ExternalID) use pointers so the patch can distinguish "omitted, keep -// existing" (nil) from "explicitly clear" (pointer to empty string) - -// e.g. when an admin migrates from static AWS credentials to IAM -// role-based auth in a single PATCH. +// value. Write-only fields (Bedrock AccessKey and AccessKeySecret) use +// pointers so the patch can distinguish "omitted, keep existing" (nil) +// from "explicitly clear" (pointer to empty string) - e.g. when an +// admin migrates from static AWS credentials to IAM role-based auth +// in a single PATCH. func mergeAIProviderSettings(existing, patch codersdk.AIProviderSettings) codersdk.AIProviderSettings { if patch.Bedrock == nil { // Patch carries no type-specific data; treat as a clear. @@ -772,9 +772,6 @@ func mergeAIProviderSettings(existing, patch codersdk.AIProviderSettings) coders if merged.AccessKeySecret == nil { merged.AccessKeySecret = existing.Bedrock.AccessKeySecret } - if merged.ExternalID == nil { - merged.ExternalID = existing.Bedrock.ExternalID - } } return codersdk.AIProviderSettings{Bedrock: &merged} } diff --git a/coderd/ai_providers_migrate.go b/coderd/ai_providers_migrate.go index 2130d2e95a171..7ec6a997b0dfe 100644 --- a/coderd/ai_providers_migrate.go +++ b/coderd/ai_providers_migrate.go @@ -404,12 +404,6 @@ func providersFromEnv(ctx context.Context, cfg codersdk.AIBridgeConfig, logger s p.BedrockModel, p.BedrockSmallFastModel, ) - bedrock.RoleARN = p.BedrockRoleARN - bedrock.SessionName = p.BedrockSessionName - if p.BedrockExternalID != "" { - externalID := p.BedrockExternalID - bedrock.ExternalID = &externalID - } isBedrock = codersdk.IsBedrockConfigured(p.BedrockBaseURL, bedrock) if isBedrock { dp.Bedrock = &bedrock diff --git a/coderd/ai_providers_test.go b/coderd/ai_providers_test.go index b120c78e0ce9d..ae3d4c577e0c3 100644 --- a/coderd/ai_providers_test.go +++ b/coderd/ai_providers_test.go @@ -1502,12 +1502,10 @@ func TestAIProviderSettingsMerge(t *testing.T) { require.Equal(t, "secret-new", *persisted.Bedrock.AccessKeySecret) }) - t.Run("MigrateStaticToRoleWithExternalID", func(t *testing.T) { + t.Run("MigrateStaticToRole", func(t *testing.T) { t.Parallel() - // An admin migrating from static AWS credentials to cross-account - // role assumption clears the keys and sets a role ARN and external - // ID in a single PATCH. The external ID is a write-only secret: it - // must persist but never be echoed back in API responses. + // An admin migrating from static AWS credentials to IAM role assumption + // clears the keys and sets a role ARN in a single PATCH. client, db := coderdtest.NewWithDatabase(t, nil) _ = coderdtest.CreateFirstUser(t, client) ctx := testutil.Context(t, testutil.WaitLong) @@ -1535,19 +1533,13 @@ func TestAIProviderSettingsMerge(t *testing.T) { AccessKey: ptr.Ref(""), AccessKeySecret: ptr.Ref(""), RoleARN: "arn:aws:iam::123456789012:role/target", - SessionName: "coder", - ExternalID: ptr.Ref("shared-secret"), }, }, }) require.NoError(t, err) - // The API response must redact the write-only external ID while - // still echoing the non-secret role fields. require.NotNil(t, updated.Settings.Bedrock) - require.Nil(t, updated.Settings.Bedrock.ExternalID) require.Equal(t, "arn:aws:iam::123456789012:role/target", updated.Settings.Bedrock.RoleARN) - require.Equal(t, "coder", updated.Settings.Bedrock.SessionName) //nolint:gocritic // Test reads the row to verify write-only fields. row, err := db.GetAIProviderByID(dbauthz.AsSystemRestricted(ctx), created.ID) @@ -1556,54 +1548,7 @@ func TestAIProviderSettingsMerge(t *testing.T) { require.NoError(t, err) require.NotNil(t, persisted.Bedrock) require.Equal(t, "arn:aws:iam::123456789012:role/target", persisted.Bedrock.RoleARN) - require.Equal(t, "coder", persisted.Bedrock.SessionName) - require.NotNil(t, persisted.Bedrock.ExternalID) - require.Equal(t, "shared-secret", *persisted.Bedrock.ExternalID) require.NotNil(t, persisted.Bedrock.AccessKey) require.Equal(t, "", *persisted.Bedrock.AccessKey) }) - - t.Run("OmittedExternalIDPreservesExisting", func(t *testing.T) { - t.Parallel() - // A PATCH that omits the external ID must keep the existing value, - // matching the write-only semantics of AccessKeySecret. - client, db := coderdtest.NewWithDatabase(t, nil) - _ = coderdtest.CreateFirstUser(t, client) - ctx := testutil.Context(t, testutil.WaitLong) - - //nolint:gocritic // Owner role is the audience for this endpoint. - created, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{ - Type: codersdk.AIProviderTypeAnthropic, - Name: "merge-role-omit", - Enabled: true, - BaseURL: "https://bedrock-runtime.us-east-1.amazonaws.com/", - Settings: codersdk.AIProviderSettings{ - Bedrock: &codersdk.AIProviderBedrockSettings{ - Region: "us-east-1", - RoleARN: "arn:aws:iam::123456789012:role/target", - ExternalID: ptr.Ref("shared-secret"), - }, - }, - }) - require.NoError(t, err) - - _, err = client.UpdateAIProvider(ctx, created.Name, codersdk.UpdateAIProviderRequest{ - Settings: &codersdk.AIProviderSettings{ - Bedrock: &codersdk.AIProviderBedrockSettings{ - Region: "us-east-1", - RoleARN: "arn:aws:iam::123456789012:role/target", - }, - }, - }) - require.NoError(t, err) - - //nolint:gocritic // Test reads the row to verify write-only fields. - row, err := db.GetAIProviderByID(dbauthz.AsSystemRestricted(ctx), created.ID) - require.NoError(t, err) - persisted, err := db2sdk.AIProviderSettings(row.Settings) - require.NoError(t, err) - require.NotNil(t, persisted.Bedrock) - require.NotNil(t, persisted.Bedrock.ExternalID) - require.Equal(t, "shared-secret", *persisted.Bedrock.ExternalID) - }) } diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index 7f9134c7a6aef..e0477dbe3ed9a 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -28,9 +28,6 @@ import ( "github.com/coder/quartz" ) -// mustNewAnthropicProvider builds an Anthropic provider for tests, panicking if -// credential resolution fails. Keeps call sites terse after NewAnthropicProvider -// gained a context and error return. func mustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) if err != nil { diff --git a/coderd/apidoc/docs.go b/coderd/apidoc/docs.go index 0c3375bb30900..fd369b9097e0c 100644 --- a/coderd/apidoc/docs.go +++ b/coderd/apidoc/docs.go @@ -15253,13 +15253,6 @@ const docTemplate = `{ "bedrock_region": { "type": "string" }, - "bedrock_role_arn": { - "description": "BedrockRoleARN, when set, is the IAM role assumed via STS before\ncalling Bedrock, enabling cross-account access. BedrockExternalID is\nthe optional external ID sent on the AssumeRole call (write-only).\nBedrockSessionName is the STS session name, auto-generated when empty.", - "type": "string" - }, - "bedrock_session_name": { - "type": "string" - }, "bedrock_small_fast_model": { "type": "string" }, diff --git a/coderd/apidoc/swagger.json b/coderd/apidoc/swagger.json index b8183339a6804..32bfbd47549f5 100644 --- a/coderd/apidoc/swagger.json +++ b/coderd/apidoc/swagger.json @@ -13605,13 +13605,6 @@ "bedrock_region": { "type": "string" }, - "bedrock_role_arn": { - "description": "BedrockRoleARN, when set, is the IAM role assumed via STS before\ncalling Bedrock, enabling cross-account access. BedrockExternalID is\nthe optional external ID sent on the AssumeRole call (write-only).\nBedrockSessionName is the STS session name, auto-generated when empty.", - "type": "string" - }, - "bedrock_session_name": { - "type": "string" - }, "bedrock_small_fast_model": { "type": "string" }, diff --git a/coderd/database/db2sdk/db2sdk.go b/coderd/database/db2sdk/db2sdk.go index b82b212a6bb35..b2ccca09d6da8 100644 --- a/coderd/database/db2sdk/db2sdk.go +++ b/coderd/database/db2sdk/db2sdk.go @@ -112,7 +112,6 @@ func redactAIProviderSettings(s codersdk.AIProviderSettings) codersdk.AIProvider b := *out.Bedrock b.AccessKey = nil b.AccessKeySecret = nil - b.ExternalID = nil out.Bedrock = &b } return out diff --git a/codersdk/aiproviders_bedrock.go b/codersdk/aiproviders_bedrock.go index 6067f0a4070b6..360b763333c08 100644 --- a/codersdk/aiproviders_bedrock.go +++ b/codersdk/aiproviders_bedrock.go @@ -32,18 +32,9 @@ type AIProviderBedrockSettings struct { AccessKeySecret *string `json:"access_key_secret,omitempty"` // RoleARN, when set, is the IAM role assumed via STS before calling // Bedrock. The base identity (static keys or the AWS environment, e.g. - // IRSA / instance profile) signs the AssumeRole call, and the resulting - // temporary credentials sign Bedrock requests. Enables cross-account - // access so usage bills to the target account. Non-secret. + // IRSA / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + // call, and the resulting temporary credentials sign Bedrock requests. RoleARN string `json:"role_arn,omitempty"` - // SessionName is the STS role session name used when assuming RoleARN. - // Auto-generated when empty. Non-secret. - SessionName string `json:"session_name,omitempty"` - // ExternalID is sent on the AssumeRole call for confused-deputy - // protection when assuming a role in an account the deployment does not - // own. Write-only (servers strip it from GET and list responses); a - // pointer for the same PATCH-omit semantics as AccessKeySecret. - ExternalID *string `json:"external_id,omitempty"` } // IsConfigured reports whether any load-bearing Bedrock field is set, diff --git a/codersdk/deployment.go b/codersdk/deployment.go index f911e5db45815..392270235beb9 100644 --- a/codersdk/deployment.go +++ b/codersdk/deployment.go @@ -4870,13 +4870,6 @@ type AIProviderConfig struct { BedrockAccessKeySecrets []string `json:"-"` BedrockModel string `json:"bedrock_model,omitempty"` BedrockSmallFastModel string `json:"bedrock_small_fast_model,omitempty"` - // BedrockRoleARN, when set, is the IAM role assumed via STS before - // calling Bedrock, enabling cross-account access. BedrockExternalID is - // the optional external ID sent on the AssumeRole call (write-only). - // BedrockSessionName is the STS session name, auto-generated when empty. - BedrockRoleARN string `json:"bedrock_role_arn,omitempty"` - BedrockExternalID string `json:"-"` - BedrockSessionName string `json:"bedrock_session_name,omitempty"` } type AIBridgeProxyConfig struct { diff --git a/docs/reference/api/general.md b/docs/reference/api/general.md index 28c967e579c30..4ae0f8d73b4dd 100644 --- a/docs/reference/api/general.md +++ b/docs/reference/api/general.md @@ -213,8 +213,6 @@ curl -X GET http://coder-server:8080/api/v2/deployment/config \ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", - "bedrock_role_arn": "string", - "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" diff --git a/docs/reference/api/schemas.md b/docs/reference/api/schemas.md index 68290fe389f11..595d110747d15 100644 --- a/docs/reference/api/schemas.md +++ b/docs/reference/api/schemas.md @@ -419,8 +419,6 @@ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", - "bedrock_role_arn": "string", - "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -954,8 +952,6 @@ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", - "bedrock_role_arn": "string", - "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -1048,8 +1044,6 @@ "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", - "bedrock_role_arn": "string", - "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -1058,16 +1052,14 @@ ### Properties -| Name | Type | Required | Restrictions | Description | -|----------------------------|--------|----------|--------------|----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| -| `base_url` | string | false | | Base URL is the base URL of the upstream provider API. | -| `bedrock_model` | string | false | | | -| `bedrock_region` | string | false | | | -| `bedrock_role_arn` | string | false | | Bedrock role arn when set, is the IAM role assumed via STS before calling Bedrock, enabling cross-account access. BedrockExternalID is the optional external ID sent on the AssumeRole call (write-only). BedrockSessionName is the STS session name, auto-generated when empty. | -| `bedrock_session_name` | string | false | | | -| `bedrock_small_fast_model` | string | false | | | -| `name` | string | false | | Name is the unique instance identifier used for routing. Defaults to Type if not provided. | -| `type` | string | false | | Type is the provider type. Valid values are: "openai", "anthropic", "azure", "bedrock", "google", "openai-compat", "openrouter", "vercel", "copilot". | +| Name | Type | Required | Restrictions | Description | +|----------------------------|--------|----------|--------------|-------------------------------------------------------------------------------------------------------------------------------------------------------| +| `base_url` | string | false | | Base URL is the base URL of the upstream provider API. | +| `bedrock_model` | string | false | | | +| `bedrock_region` | string | false | | | +| `bedrock_small_fast_model` | string | false | | | +| `name` | string | false | | Name is the unique instance identifier used for routing. Defaults to Type if not provided. | +| `type` | string | false | | Type is the provider type. Valid values are: "openai", "anthropic", "azure", "bedrock", "google", "openai-compat", "openrouter", "vercel", "copilot". | ## codersdk.AIProviderKey @@ -5546,8 +5538,6 @@ CreateWorkspaceRequest provides options for creating a new workspace. Only one o "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", - "bedrock_role_arn": "string", - "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" @@ -6150,8 +6140,6 @@ CreateWorkspaceRequest provides options for creating a new workspace. Only one o "base_url": "string", "bedrock_model": "string", "bedrock_region": "string", - "bedrock_role_arn": "string", - "bedrock_session_name": "string", "bedrock_small_fast_model": "string", "name": "string", "type": "string" diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 3201118680bf1..d9e067bfb9472 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -47,9 +47,6 @@ func singleKeyPool(t *testing.T, name, key string) *keypool.Pool { var testTracer = otel.Tracer("aibridged_inttest") -// mustNewAnthropicProvider builds an Anthropic provider for tests, panicking if -// credential resolution fails. Keeps call sites terse after NewAnthropicProvider -// gained a context and error return. func mustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) if err != nil { diff --git a/go.mod b/go.mod index 174500ea5f242..8316c92c03d25 100644 --- a/go.mod +++ b/go.mod @@ -311,7 +311,7 @@ require ( github.com/aws/aws-sdk-go-v2/service/ssm v1.67.4 // indirect github.com/aws/aws-sdk-go-v2/service/sso v1.31.3 // indirect github.com/aws/aws-sdk-go-v2/service/ssooidc v1.36.6 // indirect - github.com/aws/aws-sdk-go-v2/service/sts v1.43.3 // indirect + github.com/aws/aws-sdk-go-v2/service/sts v1.43.3 github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect github.com/aymerick/douceur v0.2.0 // indirect github.com/beorn7/perks v1.0.1 // indirect diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index fbca501e28736..0203036462289 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -312,23 +312,10 @@ export interface AIProviderBedrockSettings { /** * RoleARN, when set, is the IAM role assumed via STS before calling * Bedrock. The base identity (static keys or the AWS environment, e.g. - * IRSA / instance profile) signs the AssumeRole call, and the resulting - * temporary credentials sign Bedrock requests. Enables cross-account - * access so usage bills to the target account. Non-secret. + * IRSA / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + * call, and the resulting temporary credentials sign Bedrock requests. */ readonly role_arn?: string; - /** - * SessionName is the STS role session name used when assuming RoleARN. - * Auto-generated when empty. Non-secret. - */ - readonly session_name?: string; - /** - * ExternalID is sent on the AssumeRole call for confused-deputy - * protection when assuming a role in an account the deployment does not - * own. Write-only (servers strip it from GET and list responses); a - * pointer for the same PATCH-omit semantics as AccessKeySecret. - */ - readonly external_id?: string; } // From codersdk/aiproviders_bedrock.go @@ -364,14 +351,6 @@ export interface AIProviderConfig { readonly bedrock_region?: string; readonly bedrock_model?: string; readonly bedrock_small_fast_model?: string; - /** - * BedrockRoleARN, when set, is the IAM role assumed via STS before - * calling Bedrock, enabling cross-account access. BedrockExternalID is - * the optional external ID sent on the AssumeRole call (write-only). - * BedrockSessionName is the STS session name, auto-generated when empty. - */ - readonly bedrock_role_arn?: string; - readonly bedrock_session_name?: string; } // From codersdk/aiproviders.go diff --git a/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx b/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx index 99a1d24a92945..9ddc3169663d7 100644 --- a/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx +++ b/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx @@ -25,6 +25,7 @@ export type ProviderFormValues = { smallFastModel: string; accessKey: string; accessKeySecret: string; + roleArn: string; apiKey: string; enabled: boolean; }; @@ -66,6 +67,7 @@ const defaultInitialValues: ProviderFormValues = { smallFastModel: "", accessKey: "", accessKeySecret: "", + roleArn: "", apiKey: "", enabled: true, }; @@ -526,6 +528,17 @@ export const ProviderForm: FC = ({ View docs

+ +

+ Optional. When a role ARN is set, the gateway assumes that role + (using the base identity above) before calling Bedrock, enabling + cross-account access. +

)} diff --git a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts index 1ae786008cc36..82dc1377f9a93 100644 --- a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts +++ b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts @@ -29,6 +29,7 @@ const baseOpenAIFormValues: ProviderFormValues = { smallFastModel: "", accessKey: "", accessKeySecret: "", + roleArn: "", apiKey: "sk-test", enabled: true, }; @@ -42,6 +43,7 @@ const baseBedrockFormValues: ProviderFormValues = { smallFastModel: "anthropic.claude-haiku-4-5", accessKey: "AKIA-test", accessKeySecret: "secret", + roleArn: "", apiKey: "", enabled: true, }; @@ -55,6 +57,7 @@ const baseCopilotFormValues: ProviderFormValues = { smallFastModel: "", accessKey: "", accessKeySecret: "", + roleArn: "", apiKey: "", enabled: true, }; diff --git a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts index 67eec7e4d913a..2eaaa65a4f817 100644 --- a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts +++ b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts @@ -109,6 +109,7 @@ const buildBedrockSettings = ( smallFastModel: string, accessKey: string, accessKeySecret: string, + roleArn: string, ): BedrockSettingsWire => ({ _type: BEDROCK_SETTINGS_TYPE, _version: BEDROCK_SETTINGS_VERSION, @@ -117,6 +118,7 @@ const buildBedrockSettings = ( small_fast_model: smallFastModel, ...(accessKey ? { access_key: accessKey } : {}), ...(accessKeySecret ? { access_key_secret: accessKeySecret } : {}), + ...(roleArn ? { role_arn: roleArn } : {}), }); // Bedrock credentials live in `settings`; openai/anthropic keys go in @@ -141,6 +143,7 @@ export const providerFormValuesToCreate = ( values.smallFastModel.trim(), sanitizeCredential(values.accessKey), sanitizeCredential(values.accessKeySecret), + values.roleArn.trim(), ); return { type: "anthropic", @@ -215,6 +218,7 @@ export const providerFormValuesToUpdate = ( values.smallFastModel.trim(), credentialsChanged ? newAccessKey : "", credentialsChanged ? newAccessKeySecret : "", + values.roleArn.trim(), ); return { ...base, settings: settings as AIProviderSettings }; @@ -238,6 +242,7 @@ export const aiProviderToFormValues = ( smallFastModel: s.small_fast_model ?? "", accessKey: "", accessKeySecret: "", + roleArn: s.role_arn ?? "", enabled: provider.enabled, }; } From 511e209292a732cd8fe08ef4488ad075714784cd Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 17:03:49 +0000 Subject: [PATCH 03/21] chore: remove Bedrock role ARN UI --- .../ProvidersPage/components/ProviderForm.tsx | 13 ------------- .../components/providerFormApiMap.test.ts | 3 --- .../ProvidersPage/components/providerFormApiMap.ts | 5 ----- 3 files changed, 21 deletions(-) diff --git a/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx b/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx index 9ddc3169663d7..99a1d24a92945 100644 --- a/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx +++ b/site/src/pages/AISettingsPage/ProvidersPage/components/ProviderForm.tsx @@ -25,7 +25,6 @@ export type ProviderFormValues = { smallFastModel: string; accessKey: string; accessKeySecret: string; - roleArn: string; apiKey: string; enabled: boolean; }; @@ -67,7 +66,6 @@ const defaultInitialValues: ProviderFormValues = { smallFastModel: "", accessKey: "", accessKeySecret: "", - roleArn: "", apiKey: "", enabled: true, }; @@ -528,17 +526,6 @@ export const ProviderForm: FC = ({ View docs

- -

- Optional. When a role ARN is set, the gateway assumes that role - (using the base identity above) before calling Bedrock, enabling - cross-account access. -

)} diff --git a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts index 82dc1377f9a93..1ae786008cc36 100644 --- a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts +++ b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.test.ts @@ -29,7 +29,6 @@ const baseOpenAIFormValues: ProviderFormValues = { smallFastModel: "", accessKey: "", accessKeySecret: "", - roleArn: "", apiKey: "sk-test", enabled: true, }; @@ -43,7 +42,6 @@ const baseBedrockFormValues: ProviderFormValues = { smallFastModel: "anthropic.claude-haiku-4-5", accessKey: "AKIA-test", accessKeySecret: "secret", - roleArn: "", apiKey: "", enabled: true, }; @@ -57,7 +55,6 @@ const baseCopilotFormValues: ProviderFormValues = { smallFastModel: "", accessKey: "", accessKeySecret: "", - roleArn: "", apiKey: "", enabled: true, }; diff --git a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts index 2eaaa65a4f817..67eec7e4d913a 100644 --- a/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts +++ b/site/src/pages/AISettingsPage/ProvidersPage/components/providerFormApiMap.ts @@ -109,7 +109,6 @@ const buildBedrockSettings = ( smallFastModel: string, accessKey: string, accessKeySecret: string, - roleArn: string, ): BedrockSettingsWire => ({ _type: BEDROCK_SETTINGS_TYPE, _version: BEDROCK_SETTINGS_VERSION, @@ -118,7 +117,6 @@ const buildBedrockSettings = ( small_fast_model: smallFastModel, ...(accessKey ? { access_key: accessKey } : {}), ...(accessKeySecret ? { access_key_secret: accessKeySecret } : {}), - ...(roleArn ? { role_arn: roleArn } : {}), }); // Bedrock credentials live in `settings`; openai/anthropic keys go in @@ -143,7 +141,6 @@ export const providerFormValuesToCreate = ( values.smallFastModel.trim(), sanitizeCredential(values.accessKey), sanitizeCredential(values.accessKeySecret), - values.roleArn.trim(), ); return { type: "anthropic", @@ -218,7 +215,6 @@ export const providerFormValuesToUpdate = ( values.smallFastModel.trim(), credentialsChanged ? newAccessKey : "", credentialsChanged ? newAccessKeySecret : "", - values.roleArn.trim(), ); return { ...base, settings: settings as AIProviderSettings }; @@ -242,7 +238,6 @@ export const aiProviderToFormValues = ( smallFastModel: s.small_fast_model ?? "", accessKey: "", accessKeySecret: "", - roleArn: s.role_arn ?? "", enabled: provider.enabled, }; } From 3e671c57935cba3a5c7b93831313ed90af01c63d Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 18:47:45 +0000 Subject: [PATCH 04/21] refactor: cache only the Bedrock AssumeRole provider --- aibridge/provider/bedrock.go | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go index d7418348d09c6..dc213162ceb86 100644 --- a/aibridge/provider/bedrock.go +++ b/aibridge/provider/bedrock.go @@ -63,15 +63,18 @@ func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.Cr return nil, xerrors.Errorf("failed to load AWS Bedrock config: %w", err) } - // The base identity signs requests directly unless a target RoleARN is + // The base identity signs Bedrock requests directly unless a target role is // configured, in which case it signs the AssumeRole call and the resulting - // temporary credentials sign Bedrock requests. + // temporary credentials sign Bedrock requests. The default credential chain + // is already cache-wrapped, so only the AssumeRoleProvider is wrapped with a + // cache to avoid re-assuming the role on every request. credsProvider := base.Credentials if cfg.RoleARN != "" { credsProvider = stscreds.NewAssumeRoleProvider(sts.NewFromConfig(base), cfg.RoleARN, func(o *stscreds.AssumeRoleOptions) { o.RoleSessionName = bedrockSessionName }) + credsProvider = aws.NewCredentialsCache(credsProvider) } - return aws.NewCredentialsCache(credsProvider), nil + return credsProvider, nil } From 5aece13c5dd47966ff26dc27f4975b4cc849be29 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 19:16:52 +0000 Subject: [PATCH 05/21] test: verify Bedrock AssumeRole result is cached --- aibridge/provider/bedrock_internal_test.go | 47 ++++++++++++++++++++++ 1 file changed, 47 insertions(+) diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index 5e1c1c5fd091d..26cb8e73aac97 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -4,6 +4,7 @@ import ( "context" "net/http" "net/http/httptest" + "sync/atomic" "testing" "github.com/stretchr/testify/require" @@ -203,3 +204,49 @@ func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { require.Equal(t, "arn:aws:iam::123456789012:role/target", gotRoleARN) require.Equal(t, bedrockSessionName, gotSessionName) } + +// TestBuildBedrockCredentialsAssumeRoleCaches verifies the AssumeRole result is +// cached: many credential retrievals, one per LLM request, trigger a single STS +// AssumeRole call rather than re-assuming the role on every request. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleCaches(t *testing.T) { + var stsCalls atomic.Int64 + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + stsCalls.Add(1) + w.Header().Set("Content-Type", "text/xml") + // A far-future expiration keeps the cached credentials valid, so the + // cache serves every retrieval after the first without re-assuming. + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2999-01-01T00:00:00Z + + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + + // Each retrieval stands in for an LLM request resolving credentials from the + // shared provider. Only the first should reach STS. + for range 5 { + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "ASIAASSUMED", got.AccessKeyID) + } + + require.Equal(t, int64(1), stsCalls.Load(), + "AssumeRole should be called once, then served from the credentials cache") +} From 55ee6cca8700b060c2183b50172daf5707e4ce15 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 19:29:52 +0000 Subject: [PATCH 06/21] test: cover Bedrock AssumeRole cache expiry refresh --- aibridge/provider/bedrock_internal_test.go | 44 ++++++++++++++++++++++ 1 file changed, 44 insertions(+) diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index 26cb8e73aac97..516349a4c4c96 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -250,3 +250,47 @@ func TestBuildBedrockCredentialsAssumeRoleCaches(t *testing.T) { require.Equal(t, int64(1), stsCalls.Load(), "AssumeRole should be called once, then served from the credentials cache") } + +// TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry verifies that once the +// assumed credentials expire, the next retrieval re-assumes the role rather than +// serving stale credentials from the cache. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry(t *testing.T) { + var stsCalls atomic.Int64 + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + stsCalls.Add(1) + w.Header().Set("Content-Type", "text/xml") + // An expiration in the past makes the returned credentials immediately + // stale, so the cache cannot reuse them and must re-assume on the next + // retrieval. + _, _ = w.Write([]byte(` + + + ASIAASSUMED + assumed-secret + assumed-token + 2000-01-01T00:00:00Z + + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + + _, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + _, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + + require.Equal(t, int64(2), stsCalls.Load(), + "expired credentials should trigger a fresh AssumeRole on the next retrieval") +} From 0d6cb16be6d6148b72bed9d2b42763f4ab0bee6a Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 20:18:29 +0000 Subject: [PATCH 07/21] fix(aibridge): preserve SDK-resolved region for Bedrock signing --- aibridge/intercept/messages/base.go | 13 ++++++++++++- aibridge/provider/anthropic.go | 4 ++-- aibridge/provider/bedrock.go | 12 +++++++----- aibridge/provider/bedrock_internal_test.go | 12 ++++++------ 4 files changed, 27 insertions(+), 14 deletions(-) diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index dd11ec3e74438..1c3a441ebc5c4 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -69,6 +69,11 @@ var bedrockSupportedBetaFlags = map[string]bool{ type BedrockRuntime struct { Cfg aibconfig.AWSBedrock Creds aws.CredentialsProvider + // ResolvedRegion is the region the AWS SDK resolved at construction (from + // the environment, shared config, or IMDS). It is used for request signing + // when Cfg.Region is empty, e.g. a custom base URL with the region supplied + // via AWS_REGION. + ResolvedRegion string } type interceptionBase struct { @@ -293,8 +298,14 @@ func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context) ([]option. return nil, xerrors.Errorf("no AWS credentials found: %w", err) } + // Fall back to the SDK-resolved region (e.g. from AWS_REGION) when no + // explicit region is configured. + region := cfg.Region + if region == "" { + region = i.bedrock.ResolvedRegion + } awsCfg := aws.Config{ - Region: cfg.Region, + Region: region, Credentials: i.bedrock.Creds, } diff --git a/aibridge/provider/anthropic.go b/aibridge/provider/anthropic.go index d83af7786ab18..01cc587ecff5a 100644 --- a/aibridge/provider/anthropic.go +++ b/aibridge/provider/anthropic.go @@ -63,11 +63,11 @@ func NewAnthropic(ctx context.Context, cfg config.Anthropic, bedrockCfg *config. // so it is cheap to run at construction. var bedrock *messages.BedrockRuntime if bedrockCfg != nil { - creds, err := buildBedrockCredentials(ctx, *bedrockCfg) + creds, region, err := buildBedrockCredentials(ctx, *bedrockCfg) if err != nil { return nil, xerrors.Errorf("build bedrock credentials: %w", err) } - bedrock = &messages.BedrockRuntime{Cfg: *bedrockCfg, Creds: creds} + bedrock = &messages.BedrockRuntime{Cfg: *bedrockCfg, Creds: creds, ResolvedRegion: region} } return &Anthropic{ diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go index dc213162ceb86..0a32cdc3d974d 100644 --- a/aibridge/provider/bedrock.go +++ b/aibridge/provider/bedrock.go @@ -29,9 +29,9 @@ const bedrockSessionName = "coder-aigateway" // so per-request credential retrieval is served from this cache rather than // re-resolving (and re-assuming) on every request. No network call is made here: // the base identity and any AssumeRole are resolved lazily on first retrieval. -func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.CredentialsProvider, error) { +func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.CredentialsProvider, string, error) { if cfg.Region == "" && cfg.BaseURL == "" { - return nil, xerrors.New("region or base url required") + return nil, "", xerrors.New("region or base url required") } var loadOpts []func(*awsconfig.LoadOptions) error @@ -53,14 +53,14 @@ func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.Cr )) // Only one set: misconfiguration. case cfg.AccessKey != "" || cfg.AccessKeySecret != "": - return nil, xerrors.New("both access key and access key secret must be provided together") + return nil, "", xerrors.New("both access key and access key secret must be provided together") // Neither set: SDK default credential chain resolves the base identity. default: } base, err := awsconfig.LoadDefaultConfig(ctx, loadOpts...) if err != nil { - return nil, xerrors.Errorf("failed to load AWS Bedrock config: %w", err) + return nil, "", xerrors.Errorf("failed to load AWS Bedrock config: %w", err) } // The base identity signs Bedrock requests directly unless a target role is @@ -76,5 +76,7 @@ func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.Cr credsProvider = aws.NewCredentialsCache(credsProvider) } - return credsProvider, nil + // base.Region is the region the SDK resolved (explicit config, AWS_REGION / + // AWS_DEFAULT_REGION, shared config, or IMDS). + return credsProvider, base.Region, nil } diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index 516349a4c4c96..b9b986814e778 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -49,7 +49,7 @@ func TestBuildBedrockCredentialsValidation(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - _, err := buildBedrockCredentials(context.Background(), tt.cfg) + _, _, err := buildBedrockCredentials(context.Background(), tt.cfg) require.Error(t, err) require.Contains(t, err.Error(), tt.errorMsg) }) @@ -60,7 +60,7 @@ func TestBuildBedrockCredentialsValidation(t *testing.T) { func TestBuildBedrockCredentialsStatic(t *testing.T) { t.Parallel() - creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -135,7 +135,7 @@ func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { // buildBedrockCredentials only wires up the provider chain; it // does not resolve credentials, so it succeeds regardless of // credential availability. Resolution failures surface on Retrieve. - creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ Region: "us-east-1", }) require.NoError(t, err) @@ -189,7 +189,7 @@ func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { t.Setenv("AWS_ACCESS_KEY_ID", "base-key") t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") - creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ Region: "us-east-1", RoleARN: "arn:aws:iam::123456789012:role/target", }) @@ -233,7 +233,7 @@ func TestBuildBedrockCredentialsAssumeRoleCaches(t *testing.T) { t.Setenv("AWS_ACCESS_KEY_ID", "base-key") t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") - creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ Region: "us-east-1", RoleARN: "arn:aws:iam::123456789012:role/target", }) @@ -280,7 +280,7 @@ func TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry(t *testing.T) { t.Setenv("AWS_ACCESS_KEY_ID", "base-key") t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") - creds, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ Region: "us-east-1", RoleARN: "arn:aws:iam::123456789012:role/target", }) From 7dcf5dbecdb58e48626bbda54fa1552959135162 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 20:45:23 +0000 Subject: [PATCH 08/21] fix(aibridge): clarify Bedrock credential resolution error --- aibridge/intercept/messages/base.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index 1c3a441ebc5c4..4828749130986 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -295,7 +295,7 @@ func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context) ([]option. // the shared cache on most requests (no network); on the cold or refresh // path this performs the actual STS/IMDS call. if _, err := i.bedrock.Creds.Retrieve(ctx); err != nil { - return nil, xerrors.Errorf("no AWS credentials found: %w", err) + return nil, xerrors.Errorf("resolve AWS credentials: %w", err) } // Fall back to the SDK-resolved region (e.g. from AWS_REGION) when no From eec5a9214c54e3ac656b9aef5651ff8415795fdf Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Mon, 22 Jun 2026 23:14:09 +0000 Subject: [PATCH 09/21] feat: validate Bedrock role ARN at write time --- codersdk/aiproviders.go | 21 ++++++++++++++ codersdk/aiproviders_test.go | 53 ++++++++++++++++++++++++++++++++++++ 2 files changed, 74 insertions(+) diff --git a/codersdk/aiproviders.go b/codersdk/aiproviders.go index 7b513340bca62..a2d79348121ab 100644 --- a/codersdk/aiproviders.go +++ b/codersdk/aiproviders.go @@ -11,6 +11,7 @@ import ( "strings" "time" + "github.com/aws/aws-sdk-go-v2/aws/arn" "github.com/google/uuid" "golang.org/x/xerrors" ) @@ -246,6 +247,9 @@ func (req CreateAIProviderRequest) Validate() []ValidationError { Detail: "type=bedrock does not accept api_keys", }) } + if req.Settings.Bedrock != nil { + validations = append(validations, validateAIProviderRoleARN(req.Settings.Bedrock.RoleARN)...) + } if req.Type == AIProviderTypeCopilot && len(req.APIKeys) > 0 { validations = append(validations, ValidationError{ Field: "api_keys", @@ -294,6 +298,9 @@ func (req UpdateAIProviderRequest) Validate() []ValidationError { if req.APIKeys != nil { validations = append(validations, validateAIProviderKeyMutations(*req.APIKeys)...) } + if req.Settings != nil && req.Settings.Bedrock != nil { + validations = append(validations, validateAIProviderRoleARN(req.Settings.Bedrock.RoleARN)...) + } return validations } @@ -316,6 +323,20 @@ func validateAIProviderName(name string) []ValidationError { return validations } +func validateAIProviderRoleARN(roleARN string) []ValidationError { + if roleARN == "" { + return nil + } + parsed, err := arn.Parse(roleARN) + if err != nil || parsed.Service != "iam" || !strings.HasPrefix(parsed.Resource, "role/") { + return []ValidationError{{ + Field: "settings.role_arn", + Detail: "role_arn must be a valid IAM role ARN, e.g. arn:aws:iam::123456789012:role/BedrockRole", + }} + } + return nil +} + func validateRequiredAIProviderBaseURL(raw string) []ValidationError { if raw == "" { return []ValidationError{{Field: "base_url", Detail: "base_url is required"}} diff --git a/codersdk/aiproviders_test.go b/codersdk/aiproviders_test.go index 97baad6535dda..a0ce28ec4ea19 100644 --- a/codersdk/aiproviders_test.go +++ b/codersdk/aiproviders_test.go @@ -159,3 +159,56 @@ func TestAIProviderSettings_Roundtrip(t *testing.T) { require.NoError(t, json.Unmarshal(encoded, &got)) require.Equal(t, orig, got) } + +func TestAIProviderRequest_ValidateRoleARN(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + roleARN string + wantErr bool + }{ + {name: "empty is allowed", roleARN: "", wantErr: false}, + {name: "standard role arn", roleARN: "arn:aws:iam::743809215448:role/bedrock-role", wantErr: false}, + {name: "govcloud partition", roleARN: "arn:aws-us-gov:iam::123456789012:role/bedrock-role", wantErr: false}, + {name: "china partition", roleARN: "arn:aws-cn:iam::123456789012:role/bedrock-role", wantErr: false}, + {name: "role path", roleARN: "arn:aws:iam::123456789012:role/team/bedrock-role", wantErr: false}, + {name: "not an arn", roleARN: "bedrock-role", wantErr: true}, + {name: "wrong resource type", roleARN: "arn:aws:iam::123456789012:user/dave", wantErr: true}, + {name: "wrong service", roleARN: "arn:aws:s3:::my-bucket", wantErr: true}, + {name: "truncated arn", roleARN: "arn:aws:iam::123456789012", wantErr: true}, + } + + hasRoleARNError := func(vs []codersdk.ValidationError) bool { + for _, v := range vs { + if v.Field == "settings.role_arn" { + return true + } + } + return false + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + settings := codersdk.AIProviderSettings{ + Bedrock: &codersdk.AIProviderBedrockSettings{ + Region: "us-east-1", + RoleARN: tc.roleARN, + }, + } + + create := codersdk.CreateAIProviderRequest{ + Type: codersdk.AIProviderTypeBedrock, + Name: "bedrock", + BaseURL: "https://bedrock-runtime.us-east-1.amazonaws.com", + Settings: settings, + } + require.Equal(t, tc.wantErr, hasRoleARNError(create.Validate())) + + update := codersdk.UpdateAIProviderRequest{Settings: &settings} + require.Equal(t, tc.wantErr, hasRoleARNError(update.Validate())) + }) + } +} From a4277e5ce1c3e0c0e312e92b7f29364384db8855 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Tue, 23 Jun 2026 00:31:09 +0000 Subject: [PATCH 10/21] feat(aibridge/provider): require a region when assuming a Bedrock role --- aibridge/provider/bedrock.go | 7 ++++ aibridge/provider/bedrock_internal_test.go | 37 ++++++++++++++++++++++ 2 files changed, 44 insertions(+) diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go index 0a32cdc3d974d..d78af2ce3acc3 100644 --- a/aibridge/provider/bedrock.go +++ b/aibridge/provider/bedrock.go @@ -63,6 +63,13 @@ func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.Cr return nil, "", xerrors.Errorf("failed to load AWS Bedrock config: %w", err) } + // Assuming a role calls STS, which needs a region to resolve its endpoint. + // The region may come from the config or the AWS environment; if neither + // supplies one, fail here. + if cfg.RoleARN != "" && base.Region == "" { + return nil, "", xerrors.New("region is required to assume a role: set it explicitly or via the AWS environment") + } + // The base identity signs Bedrock requests directly unless a target role is // configured, in which case it signs the AssumeRole call and the resulting // temporary credentials sign Bedrock requests. The default credential chain diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index b9b986814e778..7f86281085496 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -294,3 +294,40 @@ func TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry(t *testing.T) { require.Equal(t, int64(2), stsCalls.Load(), "expired credentials should trigger a fresh AssumeRole on the next retrieval") } + +// TestBuildBedrockCredentialsAssumeRoleRequiresRegion verifies that configuring +// a role without a resolvable region fails at construction. STS needs a region +// to resolve its endpoint. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleRequiresRegion(t *testing.T) { + // Ensure no region resolves from the environment, shared config, or IMDS, + // so base.Region ends up empty. + t.Setenv("AWS_REGION", "") + t.Setenv("AWS_DEFAULT_REGION", "") + t.Setenv("AWS_PROFILE", "") + t.Setenv("AWS_CONFIG_FILE", "/dev/null") + t.Setenv("AWS_SHARED_CREDENTIALS_FILE", "/dev/null") + t.Setenv("AWS_EC2_METADATA_DISABLED", "true") + + _, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + BaseURL: "https://bedrock-runtime.example.com", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.ErrorContains(t, err, "region is required to assume a role") +} + +// TestBuildBedrockCredentialsAssumeRoleRegionFromEnv verifies that a role +// configured without an explicit region resolves it from the AWS environment +// (AWS_REGION here). +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleRegionFromEnv(t *testing.T) { + t.Setenv("AWS_REGION", "us-west-2") + + // BaseURL set with no explicit region: the region comes from AWS_REGION. + _, region, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + BaseURL: "https://bedrock-runtime.example.com", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) + require.Equal(t, "us-west-2", region) +} From 8769c2148d1e4c51d77b4d4005b47e598cad0e56 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Tue, 23 Jun 2026 15:35:51 -0400 Subject: [PATCH 11/21] refactor(aibridge): share Anthropic test provider via aibridgetest (#26636) --- aibridge/aibridgetest/aibridgetest.go | 17 ++++++++++++++++ aibridge/bridge_test.go | 20 ++++++------------- .../integrationtest/apidump_internal_test.go | 5 +++-- .../integrationtest/bridge_internal_test.go | 11 +++++----- .../circuit_breaker_internal_test.go | 9 +++++---- .../keypool_failover_internal_test.go | 5 +++-- .../internal/integrationtest/setupbridge.go | 13 +++--------- aibridge/provider/anthropic_internal_test.go | 2 ++ coderd/aibridged/aibridged_test.go | 13 +++--------- enterprise/aibridged_integration_test.go | 11 ++-------- 10 files changed, 50 insertions(+), 56 deletions(-) create mode 100644 aibridge/aibridgetest/aibridgetest.go diff --git a/aibridge/aibridgetest/aibridgetest.go b/aibridge/aibridgetest/aibridgetest.go new file mode 100644 index 0000000000000..19df7524757ab --- /dev/null +++ b/aibridge/aibridgetest/aibridgetest.go @@ -0,0 +1,17 @@ +package aibridgetest + +import ( + "context" + + "github.com/coder/coder/v2/aibridge" +) + +// MustNewAnthropicProvider builds an Anthropic provider for tests, panicking if +// credential resolution fails. +func MustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { + p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) + if err != nil { + panic("build anthropic provider: " + err.Error()) + } + return p +} diff --git a/aibridge/bridge_test.go b/aibridge/bridge_test.go index 6ee283d3e9601..62946617a1e9a 100644 --- a/aibridge/bridge_test.go +++ b/aibridge/bridge_test.go @@ -2,7 +2,6 @@ package aibridge_test import ( "bytes" - "context" "fmt" "io" "net/http" @@ -16,6 +15,7 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/provider" @@ -23,14 +23,6 @@ import ( var bridgeTestTracer = otel.Tracer("bridge_test") -func mustNewAnthropicProvider(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) aibridge.Provider { - p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) - if err != nil { - panic("build anthropic provider: " + err.Error()) - } - return p -} - func TestValidateProviders(t *testing.T) { t.Parallel() @@ -45,7 +37,7 @@ func TestValidateProviders(t *testing.T) { name: "all_supported_providers", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: "https://api.openai.com/v1/"}), - mustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), + aibridgetest.MustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: "https://api.individual.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-business", BaseURL: "https://api.business.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-enterprise", BaseURL: "https://api.enterprise.githubcopilot.com"}), @@ -55,7 +47,7 @@ func TestValidateProviders(t *testing.T) { name: "default_names_and_base_urls", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{}), - mustNewAnthropicProvider(config.Anthropic{}, nil), + aibridgetest.MustNewAnthropicProvider(config.Anthropic{}, nil), aibridge.NewCopilotProvider(config.Copilot{}), }, }, @@ -159,7 +151,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "anthropic_no_base_path", requestPath: "/anthropic/v1/models", provider: func(baseURL string) provider.Provider { - return mustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/models", }, @@ -168,7 +160,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { baseURLPath: "/v1", requestPath: "/anthropic/v1/models", provider: func(baseURL string) provider.Provider { - return mustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/v1/models", }, @@ -226,7 +218,7 @@ func TestRequestBodySizeLimit(t *testing.T) { return aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: baseURL}) } newAnthropic := func(baseURL string) provider.Provider { - return mustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) } newCopilot := func(baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: baseURL}) diff --git a/aibridge/internal/integrationtest/apidump_internal_test.go b/aibridge/internal/integrationtest/apidump_internal_test.go index cff3a10eed01a..81ebe77245aa2 100644 --- a/aibridge/internal/integrationtest/apidump_internal_test.go +++ b/aibridge/internal/integrationtest/apidump_internal_test.go @@ -15,6 +15,7 @@ import ( "github.com/stretchr/testify/require" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/fixtures" "github.com/coder/coder/v2/aibridge/intercept/apidump" @@ -39,7 +40,7 @@ func TestAPIDump(t *testing.T) { name: "anthropic", fixture: fixtures.AntSimple, providerFunc: func(addr, dumpDir string) aibridge.Provider { - return mustNewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.MustNewAnthropicProvider(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, path: pathAnthropicMessages, expectProviderDir: config.ProviderAnthropic, @@ -219,7 +220,7 @@ func TestAPIDumpPassthrough(t *testing.T) { { name: "anthropic", providerFunc: func(addr string, dumpDir string) aibridge.Provider { - return mustNewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.MustNewAnthropicProvider(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, requestPath: "/anthropic/v1/models", expectDumpName: "-v1-models-", diff --git a/aibridge/internal/integrationtest/bridge_internal_test.go b/aibridge/internal/integrationtest/bridge_internal_test.go index 2609a942d5f5d..5da5b0735dd56 100644 --- a/aibridge/internal/integrationtest/bridge_internal_test.go +++ b/aibridge/internal/integrationtest/bridge_internal_test.go @@ -32,6 +32,7 @@ import ( "golang.org/x/xerrors" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/fixtures" "github.com/coder/coder/v2/aibridge/intercept" @@ -342,7 +343,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(mustNewAnthropic(anthropicCfg(upstream.URL, apiKey), bedrockCfg)), + withCustomProvider(aibridgetest.MustNewAnthropicProvider(anthropicCfg(upstream.URL, apiKey), bedrockCfg)), ) // Make API call to aibridge for Anthropic /v1/messages, which will be routed via AWS Bedrock. @@ -472,7 +473,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(mustNewAnthropic(anthropicCfg(upstream.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.MustNewAnthropicProvider(anthropicCfg(upstream.URL, apiKey), bCfg)), ) reqBody, err := sjson.SetBytes(fix.Request(), "stream", streaming) @@ -626,7 +627,7 @@ func TestAWSBedrockIntegration(t *testing.T) { bCfg.Region = region bridgeServer := newBridgeTestServer(ctx, t, mockEgressProxy.URL, - withCustomProvider(mustNewAnthropic(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.MustNewAnthropicProvider(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), ) // Sends a bridge request through a mock egress proxy that @@ -2281,7 +2282,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return mustNewAnthropic(cfg, nil) + return aibridgetest.MustNewAnthropicProvider(cfg, nil) }, fixture: fixtures.AntSimple, streaming: true, @@ -2292,7 +2293,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return mustNewAnthropic(cfg, nil) + return aibridgetest.MustNewAnthropicProvider(cfg, nil) }, fixture: fixtures.AntSimple, streaming: false, diff --git a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go index 6bbd90e9b27ac..2b97102396efe 100644 --- a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go +++ b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go @@ -16,6 +16,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/metrics" @@ -70,7 +71,7 @@ func TestCircuitBreaker_FullRecoveryCycle(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return mustNewAnthropic(config.Anthropic{ + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -237,7 +238,7 @@ func TestCircuitBreaker_HalfOpenFailure(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return mustNewAnthropic(config.Anthropic{ + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -374,7 +375,7 @@ func TestCircuitBreaker_HalfOpenMaxRequests(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return mustNewAnthropic(config.Anthropic{ + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -555,7 +556,7 @@ func TestCircuitBreaker_PerModelIsolation(t *testing.T) { } ctx := t.Context() bridgeServer := newBridgeTestServer(ctx, t, mockUpstream.URL, - withCustomProvider(mustNewAnthropic(config.Anthropic{ + withCustomProvider(aibridgetest.MustNewAnthropicProvider(config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, diff --git a/aibridge/internal/integrationtest/keypool_failover_internal_test.go b/aibridge/internal/integrationtest/keypool_failover_internal_test.go index e5dbda5b98b75..ef21ad711c5a9 100644 --- a/aibridge/internal/integrationtest/keypool_failover_internal_test.go +++ b/aibridge/internal/integrationtest/keypool_failover_internal_test.go @@ -10,6 +10,7 @@ import ( "github.com/tidwall/sjson" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/fixtures" "github.com/coder/coder/v2/aibridge/internal/testutil" @@ -158,7 +159,7 @@ func TestAnthropic_KeyFailover(t *testing.T) { ) bridgeServer := newBridgeTestServer(t.Context(), t, upstream.URL, - withCustomProvider(mustNewAnthropic(config.Anthropic{ + withCustomProvider(aibridgetest.MustNewAnthropicProvider(config.Anthropic{ BaseURL: upstream.URL, KeyPool: pool, }, nil)), @@ -232,7 +233,7 @@ func TestKeyPool_StateSharing(t *testing.T) { name: "anthropic", providerName: config.ProviderAnthropic, newProvider: func(baseURL string, pool *keypool.Pool) aibridge.Provider { - return mustNewAnthropic(config.Anthropic{BaseURL: baseURL, KeyPool: pool}, nil) + return aibridgetest.MustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL, KeyPool: pool}, nil) }, upstreamResponses: []testutil.UpstreamResponse{ testutil.NewErrorResponse(http.StatusTooManyRequests, "60"), diff --git a/aibridge/internal/integrationtest/setupbridge.go b/aibridge/internal/integrationtest/setupbridge.go index f8121589ee36a..ea5d51dafb1f7 100644 --- a/aibridge/internal/integrationtest/setupbridge.go +++ b/aibridge/internal/integrationtest/setupbridge.go @@ -16,6 +16,7 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" aibcontext "github.com/coder/coder/v2/aibridge/context" "github.com/coder/coder/v2/aibridge/fixtures" @@ -249,23 +250,15 @@ func setupInjectedToolTest( return bridgeServer, mockMCP, resp } -func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) aibridge.Provider { - p, err := provider.NewAnthropic(context.Background(), cfg, bedrockCfg) - if err != nil { - panic("build anthropic provider: " + err.Error()) - } - return p -} - // newDefaultProvider creates a Provider with default test configuration. func newDefaultProvider(providerType string, addr string) aibridge.Provider { switch providerType { case config.ProviderAnthropic: - return mustNewAnthropic(anthropicCfg(addr, apiKey), nil) + return aibridgetest.MustNewAnthropicProvider(anthropicCfg(addr, apiKey), nil) case config.ProviderOpenAI: return provider.NewOpenAI(openAICfg(addr, apiKey)) case providerBedrock: - return mustNewAnthropic(anthropicCfg(addr, apiKey), bedrockCfg(addr)) + return aibridgetest.MustNewAnthropicProvider(anthropicCfg(addr, apiKey), bedrockCfg(addr)) default: panic("unknown provider type: " + providerType) } diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 8427a31671c91..87ffda59db89d 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -19,6 +19,8 @@ import ( "github.com/coder/quartz" ) +// mustNewAnthropic is local rather than using aibridgetest.MustNewAnthropicProvider +// to avoid an import cycle. func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) if err != nil { diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index e0477dbe3ed9a..d6b8c3371afe4 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -17,6 +17,7 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/intercept" "github.com/coder/coder/v2/aibridge/keypool" agplaibridge "github.com/coder/coder/v2/coderd/aibridge" @@ -28,14 +29,6 @@ import ( "github.com/coder/quartz" ) -func mustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { - p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) - if err != nil { - panic("build anthropic provider: " + err.Error()) - } - return p -} - // singleKeyPool builds a centralized key pool containing a single key. func singleKeyPool(t *testing.T, name, key string) *keypool.Pool { t.Helper() @@ -667,7 +660,7 @@ func TestServeHTTP_ActorHeaders(t *testing.T) { KeyPool: singleKeyPool(t, "openai", "test-key"), SendActorHeaders: true, }), - mustNewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.MustNewAnthropicProvider(aibridge.AnthropicConfig{ BaseURL: upstreamSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key"), SendActorHeaders: true, @@ -774,7 +767,7 @@ func TestRouting(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL, KeyPool: singleKeyPool(t, "openai", "test-key")}), - mustNewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key")}, nil), + aibridgetest.MustNewAnthropicProvider(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 d9e067bfb9472..fb48ee0836ed9 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -20,6 +20,7 @@ import ( "go.opentelemetry.io/otel/sdk/trace/tracetest" "github.com/coder/coder/v2/aibridge" + "github.com/coder/coder/v2/aibridge/aibridgetest" "github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/keypool" aibtracing "github.com/coder/coder/v2/aibridge/tracing" @@ -47,14 +48,6 @@ func singleKeyPool(t *testing.T, name, key string) *keypool.Pool { var testTracer = otel.Tracer("aibridged_inttest") -func mustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { - p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) - if err != nil { - panic("build anthropic provider: " + err.Error()) - } - return p -} - // TestIntegration is not an exhaustive test against the upstream AI providers' SDKs (see coder/aibridge for those). // This test validates that: // - intercepted requests can be authenticated/authorized @@ -512,7 +505,7 @@ func TestIntegrationCircuitBreaker(t *testing.T) { KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }), - mustNewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.MustNewAnthropicProvider(aibridge.AnthropicConfig{ BaseURL: mockAnthropic.URL, KeyPool: singleKeyPool(t, config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, From 8597d742f5cc675fcb52ae984e611b48def2de54 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Tue, 23 Jun 2026 19:32:06 -0400 Subject: [PATCH 12/21] test: fail via require.NoError in shared Anthropic test provider helper (#26641) --- aibridge/aibridgetest/aibridgetest.go | 14 ++++---- aibridge/bridge_test.go | 36 +++++++++---------- .../integrationtest/apidump_internal_test.go | 4 +-- .../integrationtest/bridge_internal_test.go | 10 +++--- .../circuit_breaker_internal_test.go | 8 ++--- .../keypool_failover_internal_test.go | 4 +-- .../internal/integrationtest/setupbridge.go | 20 +++++------ aibridge/provider/anthropic_internal_test.go | 24 ++++++------- coderd/aibridged/aibridged_test.go | 4 +-- enterprise/aibridged_integration_test.go | 2 +- 10 files changed, 64 insertions(+), 62 deletions(-) diff --git a/aibridge/aibridgetest/aibridgetest.go b/aibridge/aibridgetest/aibridgetest.go index 19df7524757ab..c86bee683ab6a 100644 --- a/aibridge/aibridgetest/aibridgetest.go +++ b/aibridge/aibridgetest/aibridgetest.go @@ -2,16 +2,18 @@ package aibridgetest import ( "context" + "testing" + + "github.com/stretchr/testify/require" "github.com/coder/coder/v2/aibridge" ) -// MustNewAnthropicProvider builds an Anthropic provider for tests, panicking if -// credential resolution fails. -func MustNewAnthropicProvider(cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { +// NewAnthropicProvider builds an Anthropic provider for tests, failing the test +// if credential resolution fails. +func NewAnthropicProvider(t testing.TB, cfg aibridge.AnthropicConfig, bedrockCfg *aibridge.AWSBedrockConfig) aibridge.Provider { + t.Helper() p, err := aibridge.NewAnthropicProvider(context.Background(), cfg, bedrockCfg) - if err != nil { - panic("build anthropic provider: " + err.Error()) - } + require.NoError(t, err) return p } diff --git a/aibridge/bridge_test.go b/aibridge/bridge_test.go index 62946617a1e9a..76242f20be280 100644 --- a/aibridge/bridge_test.go +++ b/aibridge/bridge_test.go @@ -37,7 +37,7 @@ func TestValidateProviders(t *testing.T) { name: "all_supported_providers", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: "https://api.openai.com/v1/"}), - aibridgetest.MustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), + aibridgetest.NewAnthropicProvider(t, config.Anthropic{Name: "anthropic", BaseURL: "https://api.anthropic.com/"}, nil), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: "https://api.individual.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-business", BaseURL: "https://api.business.githubcopilot.com"}), aibridge.NewCopilotProvider(config.Copilot{Name: "copilot-enterprise", BaseURL: "https://api.enterprise.githubcopilot.com"}), @@ -47,7 +47,7 @@ func TestValidateProviders(t *testing.T) { name: "default_names_and_base_urls", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{}), - aibridgetest.MustNewAnthropicProvider(config.Anthropic{}, nil), + aibridgetest.NewAnthropicProvider(t, config.Anthropic{}, nil), aibridge.NewCopilotProvider(config.Copilot{}), }, }, @@ -127,13 +127,13 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name string baseURLPath string requestPath string - provider func(string) provider.Provider + provider func(*testing.T, string) provider.Provider expectPath string }{ { name: "openAI_no_base_path", requestPath: "/openai/v1/conversations", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewOpenAIProvider(config.OpenAI{BaseURL: baseURL}) }, expectPath: "/conversations", @@ -142,7 +142,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "openAI_with_base_path", baseURLPath: "/v1", requestPath: "/openai/v1/conversations", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewOpenAIProvider(config.OpenAI{BaseURL: baseURL}) }, expectPath: "/v1/conversations", @@ -150,8 +150,8 @@ func TestPassthroughRoutesForProviders(t *testing.T) { { name: "anthropic_no_base_path", requestPath: "/anthropic/v1/models", - provider: func(baseURL string) provider.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + provider: func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/models", }, @@ -159,15 +159,15 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "anthropic_with_base_path", baseURLPath: "/v1", requestPath: "/anthropic/v1/models", - provider: func(baseURL string) provider.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL}, nil) + provider: func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/v1/models", }, { name: "copilot_no_base_path", requestPath: "/copilot/models", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL}) }, expectPath: "/models", @@ -176,7 +176,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { name: "copilot_with_base_path", baseURLPath: "/v1", requestPath: "/copilot/models", - provider: func(baseURL string) provider.Provider { + provider: func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL}) }, expectPath: "/v1/models", @@ -197,7 +197,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) { t.Cleanup(upstream.Close) rec := testutil.MockRecorder{} - prov := tc.provider(upstream.URL + tc.baseURLPath) + prov := tc.provider(t, upstream.URL+tc.baseURLPath) bridge, err := aibridge.NewRequestBridge(t.Context(), []provider.Provider{prov}, &rec, nil, logger, nil, bridgeTestTracer) require.NoError(t, err) @@ -214,13 +214,13 @@ func TestPassthroughRoutesForProviders(t *testing.T) { func TestRequestBodySizeLimit(t *testing.T) { t.Parallel() - newOpenAI := func(baseURL string) provider.Provider { + newOpenAI := func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewOpenAIProvider(config.OpenAI{Name: "openai", BaseURL: baseURL}) } - newAnthropic := func(baseURL string) provider.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) + newAnthropic := func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{Name: "anthropic", BaseURL: baseURL}, nil) } - newCopilot := func(baseURL string) provider.Provider { + newCopilot := func(_ *testing.T, baseURL string) provider.Provider { return aibridge.NewCopilotProvider(config.Copilot{Name: "copilot", BaseURL: baseURL}) } @@ -233,7 +233,7 @@ func TestRequestBodySizeLimit(t *testing.T) { tests := []struct { name string - provider func(baseURL string) provider.Provider + provider func(*testing.T, string) provider.Provider path string body []byte }{ @@ -259,7 +259,7 @@ func TestRequestBodySizeLimit(t *testing.T) { })) t.Cleanup(upstream.Close) - prov := tc.provider(upstream.URL) + prov := tc.provider(t, upstream.URL) bridge, err := aibridge.NewRequestBridge( t.Context(), []provider.Provider{prov}, diff --git a/aibridge/internal/integrationtest/apidump_internal_test.go b/aibridge/internal/integrationtest/apidump_internal_test.go index 81ebe77245aa2..b48af4c3e7ff4 100644 --- a/aibridge/internal/integrationtest/apidump_internal_test.go +++ b/aibridge/internal/integrationtest/apidump_internal_test.go @@ -40,7 +40,7 @@ func TestAPIDump(t *testing.T) { name: "anthropic", fixture: fixtures.AntSimple, providerFunc: func(addr, dumpDir string) aibridge.Provider { - return aibridgetest.MustNewAnthropicProvider(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, path: pathAnthropicMessages, expectProviderDir: config.ProviderAnthropic, @@ -220,7 +220,7 @@ func TestAPIDumpPassthrough(t *testing.T) { { name: "anthropic", providerFunc: func(addr string, dumpDir string) aibridge.Provider { - return aibridgetest.MustNewAnthropicProvider(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) }, requestPath: "/anthropic/v1/models", expectDumpName: "-v1-models-", diff --git a/aibridge/internal/integrationtest/bridge_internal_test.go b/aibridge/internal/integrationtest/bridge_internal_test.go index 5da5b0735dd56..024652928e7fd 100644 --- a/aibridge/internal/integrationtest/bridge_internal_test.go +++ b/aibridge/internal/integrationtest/bridge_internal_test.go @@ -343,7 +343,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(aibridgetest.MustNewAnthropicProvider(anthropicCfg(upstream.URL, apiKey), bedrockCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(upstream.URL, apiKey), bedrockCfg)), ) // Make API call to aibridge for Anthropic /v1/messages, which will be routed via AWS Bedrock. @@ -473,7 +473,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(aibridgetest.MustNewAnthropicProvider(anthropicCfg(upstream.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(upstream.URL, apiKey), bCfg)), ) reqBody, err := sjson.SetBytes(fix.Request(), "stream", streaming) @@ -627,7 +627,7 @@ func TestAWSBedrockIntegration(t *testing.T) { bCfg.Region = region bridgeServer := newBridgeTestServer(ctx, t, mockEgressProxy.URL, - withCustomProvider(aibridgetest.MustNewAnthropicProvider(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), ) // Sends a bridge request through a mock egress proxy that @@ -2282,7 +2282,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return aibridgetest.MustNewAnthropicProvider(cfg, nil) + return aibridgetest.NewAnthropicProvider(t, cfg, nil) }, fixture: fixtures.AntSimple, streaming: true, @@ -2293,7 +2293,7 @@ func TestActorHeaders(t *testing.T) { createProviderFn: func(url, key string, sendHeaders bool) aibridge.Provider { cfg := anthropicCfg(url, key) cfg.SendActorHeaders = sendHeaders - return aibridgetest.MustNewAnthropicProvider(cfg, nil) + return aibridgetest.NewAnthropicProvider(t, cfg, nil) }, fixture: fixtures.AntSimple, streaming: false, diff --git a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go index 2b97102396efe..57f9b27df31e3 100644 --- a/aibridge/internal/integrationtest/circuit_breaker_internal_test.go +++ b/aibridge/internal/integrationtest/circuit_breaker_internal_test.go @@ -71,7 +71,7 @@ func TestCircuitBreaker_FullRecoveryCycle(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -238,7 +238,7 @@ func TestCircuitBreaker_HalfOpenFailure(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -375,7 +375,7 @@ func TestCircuitBreaker_HalfOpenMaxRequests(t *testing.T) { }, path: pathAnthropicMessages, createProvider: func(baseURL string, cbConfig *config.CircuitBreaker) provider.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: baseURL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, @@ -556,7 +556,7 @@ func TestCircuitBreaker_PerModelIsolation(t *testing.T) { } ctx := t.Context() bridgeServer := newBridgeTestServer(ctx, t, mockUpstream.URL, - withCustomProvider(aibridgetest.MustNewAnthropicProvider(config.Anthropic{ + withCustomProvider(aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, diff --git a/aibridge/internal/integrationtest/keypool_failover_internal_test.go b/aibridge/internal/integrationtest/keypool_failover_internal_test.go index ef21ad711c5a9..cb6d3c24ef2b6 100644 --- a/aibridge/internal/integrationtest/keypool_failover_internal_test.go +++ b/aibridge/internal/integrationtest/keypool_failover_internal_test.go @@ -159,7 +159,7 @@ func TestAnthropic_KeyFailover(t *testing.T) { ) bridgeServer := newBridgeTestServer(t.Context(), t, upstream.URL, - withCustomProvider(aibridgetest.MustNewAnthropicProvider(config.Anthropic{ + withCustomProvider(aibridgetest.NewAnthropicProvider(t, config.Anthropic{ BaseURL: upstream.URL, KeyPool: pool, }, nil)), @@ -233,7 +233,7 @@ func TestKeyPool_StateSharing(t *testing.T) { name: "anthropic", providerName: config.ProviderAnthropic, newProvider: func(baseURL string, pool *keypool.Pool) aibridge.Provider { - return aibridgetest.MustNewAnthropicProvider(config.Anthropic{BaseURL: baseURL, KeyPool: pool}, nil) + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL, KeyPool: pool}, nil) }, upstreamResponses: []testutil.UpstreamResponse{ testutil.NewErrorResponse(http.StatusTooManyRequests, "60"), diff --git a/aibridge/internal/integrationtest/setupbridge.go b/aibridge/internal/integrationtest/setupbridge.go index ea5d51dafb1f7..d2e8c0929b9fd 100644 --- a/aibridge/internal/integrationtest/setupbridge.go +++ b/aibridge/internal/integrationtest/setupbridge.go @@ -46,7 +46,7 @@ const ( var defaultTracer = otel.Tracer("integrationtest") type bridgeConfig struct { - providerBuilders []func(upstreamURL string) aibridge.Provider + providerBuilders []func(t *testing.T, upstreamURL string) aibridge.Provider metrics *metrics.Metrics tracer trace.Tracer mcpProxy mcp.ServerProxier @@ -88,8 +88,8 @@ type bridgeOption func(*bridgeConfig) // When any provider option is used, the default "all providers" set is not created. func withProvider(providerType string) bridgeOption { return func(c *bridgeConfig) { - c.providerBuilders = append(c.providerBuilders, func(addr string) aibridge.Provider { - return newDefaultProvider(providerType, addr) + c.providerBuilders = append(c.providerBuilders, func(t *testing.T, addr string) aibridge.Provider { + return newDefaultProvider(t, providerType, addr) }) } } @@ -99,7 +99,7 @@ func withProvider(providerType string) bridgeOption { // When any provider option is used, the default "all providers" set is not created. func withCustomProvider(p aibridge.Provider) bridgeOption { return func(c *bridgeConfig) { - c.providerBuilders = append(c.providerBuilders, func(string) aibridge.Provider { + c.providerBuilders = append(c.providerBuilders, func(*testing.T, string) aibridge.Provider { return p }) } @@ -159,12 +159,12 @@ func newBridgeTestServer( var providers []aibridge.Provider if len(cfg.providerBuilders) > 0 { for _, b := range cfg.providerBuilders { - providers = append(providers, b(upstreamURL)) + providers = append(providers, b(t, upstreamURL)) } } else { providers = []aibridge.Provider{ - newDefaultProvider(config.ProviderAnthropic, upstreamURL), - newDefaultProvider(config.ProviderOpenAI, upstreamURL), + newDefaultProvider(t, config.ProviderAnthropic, upstreamURL), + newDefaultProvider(t, config.ProviderOpenAI, upstreamURL), } } @@ -251,14 +251,14 @@ func setupInjectedToolTest( } // newDefaultProvider creates a Provider with default test configuration. -func newDefaultProvider(providerType string, addr string) aibridge.Provider { +func newDefaultProvider(t *testing.T, providerType string, addr string) aibridge.Provider { switch providerType { case config.ProviderAnthropic: - return aibridgetest.MustNewAnthropicProvider(anthropicCfg(addr, apiKey), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfg(addr, apiKey), nil) case config.ProviderOpenAI: return provider.NewOpenAI(openAICfg(addr, apiKey)) case providerBedrock: - return aibridgetest.MustNewAnthropicProvider(anthropicCfg(addr, apiKey), bedrockCfg(addr)) + return aibridgetest.NewAnthropicProvider(t, anthropicCfg(addr, apiKey), bedrockCfg(addr)) default: panic("unknown provider type: " + providerType) } diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 87ffda59db89d..87ac73caeeb26 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -19,13 +19,13 @@ import ( "github.com/coder/quartz" ) -// mustNewAnthropic is local rather than using aibridgetest.MustNewAnthropicProvider -// to avoid an import cycle. -func mustNewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { +// newAnthropic is local (not aibridgetest.NewAnthropicProvider) because these +// white-box tests need the concrete *Anthropic, and importing aibridgetest here +// would create an import cycle. +func newAnthropic(t testing.TB, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { + t.Helper() p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) - if err != nil { - panic("build anthropic provider: " + err.Error()) - } + require.NoError(t, err) return p } @@ -56,7 +56,7 @@ func TestAnthropic_TypeAndName(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := mustNewAnthropic(tc.cfg, nil) + p := newAnthropic(t, tc.cfg, nil) assert.Equal(t, tc.expectType, p.Type()) assert.Equal(t, tc.expectName, p.Name()) }) @@ -92,7 +92,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := mustNewAnthropic(tc.cfg, nil) + p := newAnthropic(t, tc.cfg, nil) if tc.expectedKeys == nil { assert.Nil(t, p.cfg.KeyPool, "expected no KeyPool") @@ -117,7 +117,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { func TestAnthropic_CreateInterceptor(t *testing.T) { t.Parallel() - provider := mustNewAnthropic(config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) + provider := newAnthropic(t, config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) t.Run("Messages_NonStreamingRequest_BlockingInterceptor", func(t *testing.T) { t.Parallel() @@ -175,7 +175,7 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := mustNewAnthropic(config.Anthropic{ + provider := newAnthropic(t, config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), }, nil) @@ -342,7 +342,7 @@ func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { bedrock.AccessKeySecret = "wJalrXUtnFEMI-secret-value" } } - provider := mustNewAnthropic(acfg, bedrock) + provider := newAnthropic(t, 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)) @@ -386,7 +386,7 @@ func TestAnthropic_KeyFailoverConfig(t *testing.T) { pool, err := keypool.New(config.ProviderAnthropic, []string{"k0", "k1"}, quartz.NewMock(t), nil) require.NoError(t, err) - p := mustNewAnthropic(config.Anthropic{KeyPool: pool}, nil) + p := newAnthropic(t, config.Anthropic{KeyPool: pool}, nil) cfg := p.KeyFailoverConfig(slog.Make()) diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index d6b8c3371afe4..e1fd056075372 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -660,7 +660,7 @@ func TestServeHTTP_ActorHeaders(t *testing.T) { KeyPool: singleKeyPool(t, "openai", "test-key"), SendActorHeaders: true, }), - aibridgetest.MustNewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{ BaseURL: upstreamSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key"), SendActorHeaders: true, @@ -767,7 +767,7 @@ func TestRouting(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL, KeyPool: singleKeyPool(t, "openai", "test-key")}), - aibridgetest.MustNewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key")}, nil), + aibridgetest.NewAnthropicProvider(t, 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 fb48ee0836ed9..f370d8a2a6de0 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -505,7 +505,7 @@ func TestIntegrationCircuitBreaker(t *testing.T) { KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }), - aibridgetest.MustNewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{ BaseURL: mockAnthropic.URL, KeyPool: singleKeyPool(t, config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, From 167ad2dfe8dd35a1e41caf4a206a86578510698e Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Tue, 23 Jun 2026 23:40:19 +0000 Subject: [PATCH 13/21] chore: update error message --- aibridge/provider/bedrock.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aibridge/provider/bedrock.go b/aibridge/provider/bedrock.go index d78af2ce3acc3..93bbf7038a032 100644 --- a/aibridge/provider/bedrock.go +++ b/aibridge/provider/bedrock.go @@ -67,7 +67,7 @@ func buildBedrockCredentials(ctx context.Context, cfg config.AWSBedrock) (aws.Cr // The region may come from the config or the AWS environment; if neither // supplies one, fail here. if cfg.RoleARN != "" && base.Region == "" { - return nil, "", xerrors.New("region is required to assume a role: set it explicitly or via the AWS environment") + return nil, "", xerrors.New("region is required to assume a role, but was not specified") } // The base identity signs Bedrock requests directly unless a target role is From 9bc41467bf36e72633eb8c2b40c4f9d1000542b7 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 00:01:03 +0000 Subject: [PATCH 14/21] docs: add comments --- aibridge/provider/bedrock_internal_test.go | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index 7f86281085496..f067b28e15e23 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -160,6 +160,8 @@ func TestBuildBedrockCredentialsDefaultChain(t *testing.T) { // NOTE: no t.Parallel() because it uses t.Setenv. func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { var gotRoleARN, gotSessionName string + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { require.NoError(t, r.ParseForm()) gotRoleARN = r.Form.Get("RoleArn") @@ -211,6 +213,8 @@ func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { // NOTE: no t.Parallel() because it uses t.Setenv. func TestBuildBedrockCredentialsAssumeRoleCaches(t *testing.T) { var stsCalls atomic.Int64 + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { stsCalls.Add(1) w.Header().Set("Content-Type", "text/xml") @@ -257,6 +261,8 @@ func TestBuildBedrockCredentialsAssumeRoleCaches(t *testing.T) { // NOTE: no t.Parallel() because it uses t.Setenv. func TestBuildBedrockCredentialsAssumeRoleRefreshesOnExpiry(t *testing.T) { var stsCalls atomic.Int64 + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { stsCalls.Add(1) w.Header().Set("Content-Type", "text/xml") From 37529bd6be1b4705f8a35712cf0489914a2247c8 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 00:44:22 +0000 Subject: [PATCH 15/21] fix(codersdk/aiproviders): report what an invalid role_arn resolved to --- codersdk/aiproviders.go | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/codersdk/aiproviders.go b/codersdk/aiproviders.go index a2d79348121ab..6cf8fb3359e70 100644 --- a/codersdk/aiproviders.go +++ b/codersdk/aiproviders.go @@ -327,12 +327,19 @@ func validateAIProviderRoleARN(roleARN string) []ValidationError { if roleARN == "" { return nil } + const exampleRoleARN = "arn:aws:iam::123456789012:role/BedrockRole" + invalid := func(detail string) []ValidationError { + return []ValidationError{{Field: "settings.role_arn", Detail: detail}} + } parsed, err := arn.Parse(roleARN) - if err != nil || parsed.Service != "iam" || !strings.HasPrefix(parsed.Resource, "role/") { - return []ValidationError{{ - Field: "settings.role_arn", - Detail: "role_arn must be a valid IAM role ARN, e.g. arn:aws:iam::123456789012:role/BedrockRole", - }} + if err != nil { + return invalid(fmt.Sprintf("role_arn %q is not a valid ARN, e.g. %s", roleARN, exampleRoleARN)) + } + if parsed.Service != "iam" { + return invalid(fmt.Sprintf("role_arn must be an IAM ARN, but resolved to service %q, e.g. %s", parsed.Service, exampleRoleARN)) + } + if !strings.HasPrefix(parsed.Resource, "role/") { + return invalid(fmt.Sprintf("role_arn must reference an IAM role, but resolved to resource %q, e.g. %s", parsed.Resource, exampleRoleARN)) } return nil } From 2fd0b14c7be6ae014943249a84a750312fcb26b4 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 01:05:39 +0000 Subject: [PATCH 16/21] test(aibridge/provider): cover Bedrock AssumeRole failure path --- aibridge/provider/bedrock_internal_test.go | 39 ++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/aibridge/provider/bedrock_internal_test.go b/aibridge/provider/bedrock_internal_test.go index f067b28e15e23..0ecae6c780b26 100644 --- a/aibridge/provider/bedrock_internal_test.go +++ b/aibridge/provider/bedrock_internal_test.go @@ -207,6 +207,45 @@ func TestBuildBedrockCredentialsAssumeRole(t *testing.T) { require.Equal(t, bedrockSessionName, gotSessionName) } +// TestBuildBedrockCredentialsAssumeRoleError verifies that when STS rejects the +// AssumeRole call (e.g. a trust-policy or IAM denial), the failure surfaces to +// the caller on Retrieve with enough detail to diagnose it, rather than being +// swallowed. The base identity resolved fine; only the role assumption failed. +// NOTE: no t.Parallel() because it uses t.Setenv. +func TestBuildBedrockCredentialsAssumeRoleError(t *testing.T) { + // Mock the AWS STS AssumeRole API. + // https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html + sts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/xml") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(` + + Sender + AccessDenied + User arn:aws:iam::123456789012:user/base is not authorized to perform sts:AssumeRole on arn:aws:iam::123456789012:role/target + +`)) + })) + defer sts.Close() + + t.Setenv("AWS_ENDPOINT_URL_STS", sts.URL) + t.Setenv("AWS_ACCESS_KEY_ID", "base-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "base-secret") + + creds, _, err := buildBedrockCredentials(context.Background(), config.AWSBedrock{ + Region: "us-east-1", + RoleARN: "arn:aws:iam::123456789012:role/target", + }) + require.NoError(t, err) // Build is lazy; the STS call happens on Retrieve. + + _, err = creds.Retrieve(context.Background()) + require.Error(t, err) + // The error must carry the STS operation and failure code so operators can + // tell this is an AssumeRole authorization problem, not missing credentials. + require.ErrorContains(t, err, "AssumeRole") + require.ErrorContains(t, err, "AccessDenied") +} + // TestBuildBedrockCredentialsAssumeRoleCaches verifies the AssumeRole result is // cached: many credential retrievals, one per LLM request, trigger a single STS // AssumeRole call rather than re-assuming the role on every request. From 5ef859fcdfe353385946017253eac3850c429ae6 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 14:19:22 +0000 Subject: [PATCH 17/21] test(aibridge/intercept): assert credential hints fit the DB column --- aibridge/intercept/credential_test.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/aibridge/intercept/credential_test.go b/aibridge/intercept/credential_test.go index 6c147380aba5c..84e27e083a9f3 100644 --- a/aibridge/intercept/credential_test.go +++ b/aibridge/intercept/credential_test.go @@ -19,6 +19,9 @@ import ( func TestCredential(t *testing.T) { t.Parallel() + // Matches the VARCHAR(15) DB constraint. + const maxCredentialHintLength = 15 + tests := []struct { name string newCred func(t *testing.T) intercept.Credential @@ -118,6 +121,8 @@ func TestCredential(t *testing.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.LessOrEqual(t, len(cred.Hint()), maxCredentialHintLength, + "Hint must fit the credential_hint column") assert.Equal(t, tc.expectLength, cred.Length(), "Length") credBYOK, credBYOKOK := intercept.AsBYOK(cred) From b07ef4403dfdd9930ca33435970120833ed9a1e9 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 14:31:27 +0000 Subject: [PATCH 18/21] test(aibridge/provider): rename newAnthropic to newTestAnthropic --- aibridge/provider/anthropic_internal_test.go | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/aibridge/provider/anthropic_internal_test.go b/aibridge/provider/anthropic_internal_test.go index 87ac73caeeb26..9df6f0371843f 100644 --- a/aibridge/provider/anthropic_internal_test.go +++ b/aibridge/provider/anthropic_internal_test.go @@ -19,10 +19,10 @@ import ( "github.com/coder/quartz" ) -// newAnthropic is local (not aibridgetest.NewAnthropicProvider) because these +// newTestAnthropic is local (not aibridgetest.NewAnthropicProvider) because these // white-box tests need the concrete *Anthropic, and importing aibridgetest here // would create an import cycle. -func newAnthropic(t testing.TB, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { +func newTestAnthropic(t testing.TB, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { t.Helper() p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) require.NoError(t, err) @@ -56,7 +56,7 @@ func TestAnthropic_TypeAndName(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := newAnthropic(t, tc.cfg, nil) + p := newTestAnthropic(t, tc.cfg, nil) assert.Equal(t, tc.expectType, p.Type()) assert.Equal(t, tc.expectName, p.Name()) }) @@ -92,7 +92,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := newAnthropic(t, tc.cfg, nil) + p := newTestAnthropic(t, tc.cfg, nil) if tc.expectedKeys == nil { assert.Nil(t, p.cfg.KeyPool, "expected no KeyPool") @@ -117,7 +117,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { func TestAnthropic_CreateInterceptor(t *testing.T) { t.Parallel() - provider := newAnthropic(t, config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) + provider := newTestAnthropic(t, config.Anthropic{KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key")}, nil) t.Run("Messages_NonStreamingRequest_BlockingInterceptor", func(t *testing.T) { t.Parallel() @@ -175,7 +175,7 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := newAnthropic(t, config.Anthropic{ + provider := newTestAnthropic(t, config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), }, nil) @@ -342,7 +342,7 @@ func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { bedrock.AccessKeySecret = "wJalrXUtnFEMI-secret-value" } } - provider := newAnthropic(t, acfg, bedrock) + provider := newTestAnthropic(t, 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)) @@ -386,7 +386,7 @@ func TestAnthropic_KeyFailoverConfig(t *testing.T) { pool, err := keypool.New(config.ProviderAnthropic, []string{"k0", "k1"}, quartz.NewMock(t), nil) require.NoError(t, err) - p := newAnthropic(t, config.Anthropic{KeyPool: pool}, nil) + p := newTestAnthropic(t, config.Anthropic{KeyPool: pool}, nil) cfg := p.KeyFailoverConfig(slog.Make()) From 9b436d6ca84f911207742e1842d2553afce7623f Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 15:03:11 +0000 Subject: [PATCH 19/21] test(cli/aibridged): cover RoleARN in bedrockConfigFromRow --- cli/aibridged_internal_test.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/cli/aibridged_internal_test.go b/cli/aibridged_internal_test.go index dee00f841f2d8..d431b063b7eb1 100644 --- a/cli/aibridged_internal_test.go +++ b/cli/aibridged_internal_test.go @@ -267,6 +267,7 @@ func TestBuildProviders(t *testing.T) { Name: "anthropic-bedrock", BaseUrl: "https://bedrock-runtime.us-west-2.amazonaws.com/", } + roleARN := "arn:aws:iam::123456789012:role/BedrockRole" settings := codersdk.AIProviderSettings{ Bedrock: &codersdk.AIProviderBedrockSettings{ Region: "us-west-2", @@ -274,6 +275,7 @@ func TestBuildProviders(t *testing.T) { AccessKeySecret: &secret, Model: model, SmallFastModel: smallModel, + RoleARN: roleARN, }, } got := bedrockConfigFromRow(row, settings) @@ -284,6 +286,7 @@ func TestBuildProviders(t *testing.T) { assert.Equal(t, secret, got.AccessKeySecret) assert.Equal(t, model, got.Model) assert.Equal(t, smallModel, got.SmallFastModel) + assert.Equal(t, roleARN, got.RoleARN) }) t.Run("BedrockSettingsEmpty", func(t *testing.T) { From 4ea714e2eddd0799e5bfcc367f355c2cdbf4d9e2 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 15:08:27 +0000 Subject: [PATCH 20/21] test(coderd/ai_providers): assert secret cleared in MigrateStaticToRole --- coderd/ai_providers_test.go | 2 ++ 1 file changed, 2 insertions(+) diff --git a/coderd/ai_providers_test.go b/coderd/ai_providers_test.go index ae3d4c577e0c3..4fcc1dd63bc6e 100644 --- a/coderd/ai_providers_test.go +++ b/coderd/ai_providers_test.go @@ -1550,5 +1550,7 @@ func TestAIProviderSettingsMerge(t *testing.T) { require.Equal(t, "arn:aws:iam::123456789012:role/target", persisted.Bedrock.RoleARN) require.NotNil(t, persisted.Bedrock.AccessKey) require.Equal(t, "", *persisted.Bedrock.AccessKey) + require.NotNil(t, persisted.Bedrock.AccessKeySecret) + require.Equal(t, "", *persisted.Bedrock.AccessKeySecret) }) } From ed11c8c2754ce2271f1c7325faa47263de6eab43 Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 15:10:26 +0000 Subject: [PATCH 21/21] docs: minor fix --- aibridge/intercept/messages/base.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index 4828749130986..7c895595e9799 100644 --- a/aibridge/intercept/messages/base.go +++ b/aibridge/intercept/messages/base.go @@ -274,8 +274,8 @@ func (i *interceptionBase) withBody() option.RequestOption { // withAWSBedrockOptions returns request options for authenticating with AWS Bedrock. // -// Credentials come from i.bedrock.Creds, it is a shared credentials cache, so the per-request -// Retrieve below is served from that cache and does not re-resolve or re-assume on every request. +// Credentials come from i.bedrock.Creds. It is a shared credentials cache, so the per-request Retrieve() +// below is served from that cache and does not re-resolve or re-assume on every request. func (i *interceptionBase) withAWSBedrockOptions(ctx context.Context) ([]option.RequestOption, error) { if i.bedrock == nil { return nil, xerrors.New("nil bedrock runtime")