From 8bf6f430161449ccb0927072343f6b061f8b425b Mon Sep 17 00:00:00 2001 From: Yevhenii Shcherbina Date: Wed, 24 Jun 2026 12:03:27 -0400 Subject: [PATCH] feat: support cross-account Bedrock AssumeRole in AI Bridge (#26527) # Support IAM role assumption for AWS Bedrock in AI Bridge ## Summary Implements https://linear.app/codercom/issue/AIGOV-371/support-dynamic-bedrock-assumerole-across-aws-accounts-for-ai-gateway A Bedrock provider can now be configured with an IAM role to assume. Before calling Bedrock, the gateway assumes that role via STS and signs requests with the resulting temporary credentials. Whether the role lives in the same account or another one is entirely a matter of the role's trust policy. ## Problem Many organizations prohibit long-lived AWS access keys and expect workloads to authenticate through assumed IAM roles instead. A common case is an organization that runs Bedrock across several AWS accounts, one per business unit, and needs each unit's usage billed to its own account by assuming a role there. AI Bridge previously authenticated a Bedrock provider only with static keys or the gateway's own ambient AWS identity, which is shared by every provider, with no way to assume a role. These deployments had no clean path. ## How it works When a provider is configured with a role ARN, the gateway uses its base identity to assume that role via STS and signs Bedrock requests with the temporary credentials it returns. The base identity is whatever the AWS default credential chain resolves, IRSA, EKS Pod Identity, EC2 Instance Profile, or static keys. Credentials are resolved once when the provider is set up and are then cached and rotated, so individual requests are served from the cache rather than triggering a new STS call. A deployment that needs several roles configures several providers, each pointing at its own role. ## Configuration The role ARN is part of the Bedrock provider settings and is set through the AI provider API. It is optional: a provider with no role ARN behaves exactly as before. ## Scope and trade-offs - This PR is backend only. The settings UI for the role ARN ships in a follow-up. - Configuration is not exposed through environment variables. Environment-based provider configuration is being phased out in favor of database-managed providers, so the role ARN is intentionally database and API only. Follow-up PR: https://github.com/coder/coder/pull/26578 --- aibridge/aibridgetest/aibridgetest.go | 19 + aibridge/api.go | 4 +- aibridge/bridge_test.go | 37 +- aibridge/config/config.go | 5 + aibridge/intercept/credential.go | 2 +- aibridge/intercept/credential_test.go | 7 +- aibridge/intercept/messages/base.go | 91 ++--- .../intercept/messages/base_internal_test.go | 152 ++----- aibridge/intercept/messages/blocking.go | 5 +- aibridge/intercept/messages/streaming.go | 5 +- .../integrationtest/apidump_internal_test.go | 5 +- .../integrationtest/bridge_internal_test.go | 26 +- .../circuit_breaker_internal_test.go | 9 +- .../keypool_failover_internal_test.go | 5 +- .../internal/integrationtest/setupbridge.go | 21 +- aibridge/passthrough_internal_test.go | 5 +- aibridge/provider/anthropic.go | 35 +- aibridge/provider/anthropic_internal_test.go | 25 +- aibridge/provider/bedrock.go | 89 +++++ aibridge/provider/bedrock_internal_test.go | 378 ++++++++++++++++++ cli/aibridged.go | 8 +- cli/aibridged_internal_test.go | 3 + cli/server_aibridge_internal_test.go | 4 +- coderd/ai_providers_test.go | 52 +++ coderd/aibridged/aibridged_test.go | 5 +- codersdk/aiproviders.go | 28 ++ codersdk/aiproviders_bedrock.go | 8 + codersdk/aiproviders_test.go | 53 +++ enterprise/aibridged_integration_test.go | 3 +- go.mod | 2 +- site/src/api/typesGenerated.ts | 7 + 31 files changed, 837 insertions(+), 261 deletions(-) create mode 100644 aibridge/aibridgetest/aibridgetest.go create mode 100644 aibridge/provider/bedrock.go create mode 100644 aibridge/provider/bedrock_internal_test.go diff --git a/aibridge/aibridgetest/aibridgetest.go b/aibridge/aibridgetest/aibridgetest.go new file mode 100644 index 0000000000..c86bee683a --- /dev/null +++ b/aibridge/aibridgetest/aibridgetest.go @@ -0,0 +1,19 @@ +package aibridgetest + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/aibridge" +) + +// 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) + require.NoError(t, err) + return p +} diff --git a/aibridge/api.go b/aibridge/api.go index 34dce84ef8..b74912d565 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 9ac7ea9ec3..76242f20be 100644 --- a/aibridge/bridge_test.go +++ b/aibridge/bridge_test.go @@ -15,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" @@ -36,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/"}), - aibridge.NewAnthropicProvider(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"}), @@ -46,7 +47,7 @@ func TestValidateProviders(t *testing.T) { name: "default_names_and_base_urls", providers: []provider.Provider{ aibridge.NewOpenAIProvider(config.OpenAI{}), - aibridge.NewAnthropicProvider(config.Anthropic{}, nil), + aibridgetest.NewAnthropicProvider(t, config.Anthropic{}, nil), aibridge.NewCopilotProvider(config.Copilot{}), }, }, @@ -126,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", @@ -141,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", @@ -149,8 +150,8 @@ 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) + provider: func(t *testing.T, baseURL string) provider.Provider { + return aibridgetest.NewAnthropicProvider(t, config.Anthropic{BaseURL: baseURL}, nil) }, expectPath: "/v1/models", }, @@ -158,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 aibridge.NewAnthropicProvider(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", @@ -175,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", @@ -196,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) @@ -213,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 aibridge.NewAnthropicProvider(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}) } @@ -232,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 }{ @@ -258,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/config/config.go b/aibridge/config/config.go index 9fdeb32d57..3dc76841aa 100644 --- a/aibridge/config/config.go +++ b/aibridge/config/config.go @@ -33,6 +33,11 @@ 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 / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + // call, and the resulting temporary credentials sign Bedrock requests. + RoleARN string } // OpenAI carries configuration for an OpenAI provider. diff --git a/aibridge/intercept/credential.go b/aibridge/intercept/credential.go index 008c43463d..24016b6950 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 0f78d8dee0..84e27e083a 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 @@ -72,7 +75,7 @@ func TestCredential(t *testing.T) { }, expectKind: intercept.CredentialKindCentralized, expectAuthHeader: "", - expectHint: "", + expectHint: "", expectLength: 0, }, { @@ -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) diff --git a/aibridge/intercept/messages/base.go b/aibridge/intercept/messages/base.go index dff58c1cad..7c895595e9 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" @@ -65,13 +64,26 @@ 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 + // 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 { id uuid.UUID reqPayload RequestPayload - cfg intercept.Config - cred intercept.Credential - bedrockCfg *aibconfig.AWSBedrock + 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 @@ -107,10 +119,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 } @@ -126,7 +138,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), } } @@ -238,10 +250,10 @@ func (i *interceptionBase) newMessagesService(ctx context.Context, opts ...optio opts = append(opts, option.WithMiddleware(mw)) } - if i.bedrockCfg != nil { + if i.bedrock != nil { ctx, cancel := context.WithTimeout(ctx, time.Second*30) defer cancel() - bedrockOpts, err := i.withAWSBedrockOptions(ctx, i.bedrockCfg) + bedrockOpts, err := i.withAWSBedrockOptions(ctx) if err != nil { return anthropic.MessageService{}, err } @@ -262,14 +274,13 @@ 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) { - 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") } @@ -280,38 +291,22 @@ func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconf return nil, xerrors.New("small fast model required") } - loadOpts := []func(*config.LoadOptions) error{ - config.WithRegion(cfg.Region), + // Fail fast: ensure credentials can be resolved before signing. Served from + // 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("resolve AWS credentials: %w", err) } - // 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: + // 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, 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 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: region, + Credentials: i.bedrock.Creds, } var out []option.RequestOption @@ -336,7 +331,7 @@ func (*interceptionBase) withAWSBedrockOptions(ctx context.Context, cfg *aibconf // 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 5a4c5e6f30..0b2ee4d77b 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" @@ -87,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", @@ -104,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", @@ -119,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", @@ -130,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", @@ -140,32 +141,9 @@ 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{ + cfg: config.AWSBedrock{ Region: "us-east-1", AccessKey: "test-key", AccessKeySecret: "test-secret", @@ -177,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", @@ -189,24 +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() - base := &interceptionBase{} - opts, err := base.withAWSBedrockOptions(context.Background(), tt.cfg) + base := &interceptionBase{ + bedrock: &BedrockRuntime{ + Cfg: tt.cfg, + Creds: credentials.NewStaticCredentialsProvider("test-key", "test-secret", ""), + }, + } + opts, err := base.withAWSBedrockOptions(context.Background()) if tt.expectError { require.Error(t, err) @@ -219,87 +196,16 @@ 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", - }, - } +// 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() - 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) - - 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()) + require.Error(t, err) + require.Contains(t, err.Error(), "nil bedrock runtime") } func TestAccumulateUsage(t *testing.T) { @@ -878,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 e6b5a8a596..8b3d9cc613 100644 --- a/aibridge/intercept/messages/blocking.go +++ b/aibridge/intercept/messages/blocking.go @@ -17,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" @@ -36,7 +35,7 @@ func NewBlockingInterceptor( reqPayload RequestPayload, cfg intercept.Config, cred intercept.Credential, - bedrockCfg *aibconfig.AWSBedrock, + bedrock *BedrockRuntime, clientHeaders http.Header, tracer trace.Tracer, ) *BlockingInterception { @@ -45,7 +44,7 @@ func NewBlockingInterceptor( reqPayload: reqPayload, cfg: cfg, cred: cred, - bedrockCfg: bedrockCfg, + bedrock: bedrock, clientHeaders: clientHeaders, tracer: tracer, }} diff --git a/aibridge/intercept/messages/streaming.go b/aibridge/intercept/messages/streaming.go index 442c4291b0..2369c17654 100644 --- a/aibridge/intercept/messages/streaming.go +++ b/aibridge/intercept/messages/streaming.go @@ -21,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" @@ -41,7 +40,7 @@ func NewStreamingInterceptor( reqPayload RequestPayload, cfg intercept.Config, cred intercept.Credential, - bedrockCfg *aibconfig.AWSBedrock, + bedrock *BedrockRuntime, clientHeaders http.Header, tracer trace.Tracer, ) *StreamingInterception { @@ -50,7 +49,7 @@ func NewStreamingInterceptor( reqPayload: reqPayload, cfg: cfg, cred: cred, - bedrockCfg: bedrockCfg, + bedrock: bedrock, clientHeaders: clientHeaders, tracer: tracer, }} diff --git a/aibridge/internal/integrationtest/apidump_internal_test.go b/aibridge/internal/integrationtest/apidump_internal_test.go index 42811cb362..b48af4c3e7 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 provider.NewAnthropic(anthropicCfgWithAPIDump(addr, apiKey, dumpDir), nil) + return aibridgetest.NewAnthropicProvider(t, 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 provider.NewAnthropic(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 ef226db2b9..024652928e 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" @@ -316,19 +317,8 @@ func TestAWSBedrockIntegration(t *testing.T) { SmallFastModel: "test-haiku", } - bridgeServer := newBridgeTestServer(ctx, t, "http://unused", - withCustomProvider(provider.NewAnthropic(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) { @@ -353,7 +343,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(provider.NewAnthropic(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. @@ -483,7 +473,7 @@ func TestAWSBedrockIntegration(t *testing.T) { } bridgeServer := newBridgeTestServer(ctx, t, upstream.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(upstream.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(upstream.URL, apiKey), bCfg)), ) reqBody, err := sjson.SetBytes(fix.Request(), "stream", streaming) @@ -637,7 +627,7 @@ func TestAWSBedrockIntegration(t *testing.T) { bCfg.Region = region bridgeServer := newBridgeTestServer(ctx, t, mockEgressProxy.URL, - withCustomProvider(provider.NewAnthropic(anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), + withCustomProvider(aibridgetest.NewAnthropicProvider(t, anthropicCfg(mockEgressProxy.URL, apiKey), bCfg)), ) // Sends a bridge request through a mock egress proxy that @@ -2292,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 provider.NewAnthropic(cfg, nil) + return aibridgetest.NewAnthropicProvider(t, cfg, nil) }, fixture: fixtures.AntSimple, streaming: true, @@ -2303,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 provider.NewAnthropic(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 e9909b5ae3..57f9b27df3 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 provider.NewAnthropic(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, 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 provider.NewAnthropic(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, 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 provider.NewAnthropic(config.Anthropic{ + return aibridgetest.NewAnthropicProvider(t, 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(provider.NewAnthropic(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 af92056c87..cb6d3c24ef 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(provider.NewAnthropic(config.Anthropic{ + withCustomProvider(aibridgetest.NewAnthropicProvider(t, 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 provider.NewAnthropic(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 e63f554a00..d2e8c0929b 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" @@ -45,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 @@ -87,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) }) } } @@ -98,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 }) } @@ -158,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), } } @@ -250,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 provider.NewAnthropic(anthropicCfg(addr, apiKey), nil) + return aibridgetest.NewAnthropicProvider(t, anthropicCfg(addr, apiKey), nil) case config.ProviderOpenAI: return provider.NewOpenAI(openAICfg(addr, apiKey)) case providerBedrock: - return provider.NewAnthropic(anthropicCfg(addr, apiKey), bedrockCfg(addr)) + return aibridgetest.NewAnthropicProvider(t, 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 5a0aeb3cb6..76ca6c1731 100644 --- a/aibridge/passthrough_internal_test.go +++ b/aibridge/passthrough_internal_test.go @@ -1,6 +1,7 @@ package aibridge import ( + "context" "crypto/tls" "maps" "net" @@ -316,10 +317,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 2b04bdf8b1..01cc587ecf 100644 --- a/aibridge/provider/anthropic.go +++ b/aibridge/provider/anthropic.go @@ -1,6 +1,7 @@ package provider import ( + "context" "fmt" "io" "net/http" @@ -25,8 +26,9 @@ var _ Provider = &Anthropic{} // Anthropic allows for interactions with the Anthropic API. type Anthropic struct { - cfg config.Anthropic - bedrockCfg *config.AWSBedrock + cfg config.Anthropic + // bedrock is nil for non-Bedrock providers. + bedrock *messages.BedrockRuntime } const routeMessages = "/v1/messages" // https://docs.anthropic.com/en/api/messages @@ -43,7 +45,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 } @@ -55,10 +57,23 @@ 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 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 { + 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, ResolvedRegion: region} } + + return &Anthropic{ + cfg: cfg, + bedrock: bedrock, + }, nil } func (*Anthropic) Type() string { @@ -123,9 +138,9 @@ func (p *Anthropic) CreateInterceptor(_ http.ResponseWriter, r *http.Request, tr var interceptor intercept.Interceptor if reqPayload.Stream() { - interceptor = messages.NewStreamingInterceptor(id, reqPayload, cfg, cred, p.bedrockCfg, 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, r.Header, tracer) + interceptor = messages.NewBlockingInterceptor(id, reqPayload, cfg, cred, p.bedrock, r.Header, tracer) } span.SetAttributes(interceptor.TraceAttributes(r)...) return interceptor, nil @@ -153,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 679469fdc6..9df6f03718 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,16 @@ import ( "github.com/coder/quartz" ) +// 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 newTestAnthropic(t testing.TB, cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropic { + t.Helper() + p, err := NewAnthropic(context.Background(), cfg, bedrockCfg) + require.NoError(t, err) + return p +} + func TestAnthropic_TypeAndName(t *testing.T) { t.Parallel() @@ -45,7 +56,7 @@ func TestAnthropic_TypeAndName(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - p := NewAnthropic(tc.cfg, nil) + p := newTestAnthropic(t, tc.cfg, nil) assert.Equal(t, tc.expectType, p.Type()) assert.Equal(t, tc.expectName, p.Name()) }) @@ -81,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(tc.cfg, nil) + p := newTestAnthropic(t, tc.cfg, nil) if tc.expectedKeys == nil { assert.Nil(t, p.cfg.KeyPool, "expected no KeyPool") @@ -106,7 +117,7 @@ func TestNewAnthropic_KeyResolution(t *testing.T) { func TestAnthropic_CreateInterceptor(t *testing.T) { t.Parallel() - provider := NewAnthropic(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() @@ -164,7 +175,7 @@ func TestAnthropic_CreateInterceptor(t *testing.T) { })) t.Cleanup(mockUpstream.Close) - provider := NewAnthropic(config.Anthropic{ + provider := newTestAnthropic(t, config.Anthropic{ BaseURL: mockUpstream.URL, KeyPool: testutil.SingleKeyPool(config.ProviderAnthropic, "test-key"), }, nil) @@ -286,7 +297,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. @@ -331,7 +342,7 @@ func TestAnthropic_CreateInterceptor_Credential(t *testing.T) { bedrock.AccessKeySecret = "wJalrXUtnFEMI-secret-value" } } - provider := NewAnthropic(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)) @@ -375,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(config.Anthropic{KeyPool: pool}, nil) + p := newTestAnthropic(t, 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 0000000000..93bbf7038a --- /dev/null +++ b/aibridge/provider/bedrock.go @@ -0,0 +1,89 @@ +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" +) + +// 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 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. 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, string, 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) + } + + // 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, but was not specified") + } + + // 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 + // 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) + } + + // 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 new file mode 100644 index 0000000000..0ecae6c780 --- /dev/null +++ b/aibridge/provider/bedrock_internal_test.go @@ -0,0 +1,378 @@ +package provider + +import ( + "context" + "net/http" + "net/http/httptest" + "sync/atomic" + "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. +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. +// 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", + envVars: map[string]string{ + "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", + envVars: map[string]string{ + "AWS_ACCESS_KEY_ID": "test-key", + "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", + 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 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) + + 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 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, 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") + 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 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", + }) + 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, 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. +// 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") + // 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") +} + +// 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 + // 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") + // 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") +} + +// 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) +} diff --git a/cli/aibridged.go b/cli/aibridged.go index e838c1db5a..9aa2ea5c27 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 @@ -348,6 +349,7 @@ func bedrockConfigFromRow(row database.AIProvider, settings codersdk.AIProviderS AccessKeySecret: accessKeySecret, Model: bedrockSettings.Model, SmallFastModel: bedrockSettings.SmallFastModel, + RoleARN: bedrockSettings.RoleARN, } } diff --git a/cli/aibridged_internal_test.go b/cli/aibridged_internal_test.go index dee00f841f..d431b063b7 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) { diff --git a/cli/server_aibridge_internal_test.go b/cli/server_aibridge_internal_test.go index 21711b0289..cb08ec530b 100644 --- a/cli/server_aibridge_internal_test.go +++ b/cli/server_aibridge_internal_test.go @@ -741,7 +741,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 +755,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_test.go b/coderd/ai_providers_test.go index b9bfd283f1..4fcc1dd63b 100644 --- a/coderd/ai_providers_test.go +++ b/coderd/ai_providers_test.go @@ -1501,4 +1501,56 @@ func TestAIProviderSettingsMerge(t *testing.T) { require.NotNil(t, persisted.Bedrock.AccessKeySecret) require.Equal(t, "secret-new", *persisted.Bedrock.AccessKeySecret) }) + + t.Run("MigrateStaticToRole", func(t *testing.T) { + t.Parallel() + // 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) + + //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", + }, + }, + }) + require.NoError(t, err) + + require.NotNil(t, updated.Settings.Bedrock) + require.Equal(t, "arn:aws:iam::123456789012:role/target", updated.Settings.Bedrock.RoleARN) + + //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.NotNil(t, persisted.Bedrock.AccessKey) + require.Equal(t, "", *persisted.Bedrock.AccessKey) + require.NotNil(t, persisted.Bedrock.AccessKeySecret) + require.Equal(t, "", *persisted.Bedrock.AccessKeySecret) + }) } diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index c0592e7390..e1fd056075 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" @@ -659,7 +660,7 @@ func TestServeHTTP_ActorHeaders(t *testing.T) { KeyPool: singleKeyPool(t, "openai", "test-key"), SendActorHeaders: true, }), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{ BaseURL: upstreamSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key"), SendActorHeaders: true, @@ -766,7 +767,7 @@ func TestRouting(t *testing.T) { providers := []aibridge.Provider{ aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL, KeyPool: singleKeyPool(t, "openai", "test-key")}), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL, KeyPool: singleKeyPool(t, "anthropic", "test-key")}, nil), + 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/codersdk/aiproviders.go b/codersdk/aiproviders.go index 7b513340bc..6cf8fb3359 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,27 @@ func validateAIProviderName(name string) []ValidationError { return validations } +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 { + 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 +} + func validateRequiredAIProviderBaseURL(raw string) []ValidationError { if raw == "" { return []ValidationError{{Field: "base_url", Detail: "base_url is required"}} diff --git a/codersdk/aiproviders_bedrock.go b/codersdk/aiproviders_bedrock.go index 88edcb0017..360b763333 100644 --- a/codersdk/aiproviders_bedrock.go +++ b/codersdk/aiproviders_bedrock.go @@ -30,6 +30,11 @@ 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 / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + // call, and the resulting temporary credentials sign Bedrock requests. + RoleARN string `json:"role_arn,omitempty"` } // IsConfigured reports whether any load-bearing Bedrock field is set, @@ -47,6 +52,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/aiproviders_test.go b/codersdk/aiproviders_test.go index 97baad6535..a0ce28ec4e 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())) + }) + } +} diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 791c7323a5..f370d8a2a6 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" @@ -504,7 +505,7 @@ func TestIntegrationCircuitBreaker(t *testing.T) { KeyPool: singleKeyPool(t, config.ProviderOpenAI, "test-key"), CircuitBreaker: cbConfig, }), - aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{ + aibridgetest.NewAnthropicProvider(t, aibridge.AnthropicConfig{ BaseURL: mockAnthropic.URL, KeyPool: singleKeyPool(t, config.ProviderAnthropic, "test-key"), CircuitBreaker: cbConfig, diff --git a/go.mod b/go.mod index c3ea78e67a..ede6490189 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 9301a9eb79..6ecb349f45 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -322,6 +322,13 @@ 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 / EKS Pod Identity / EC2 Instance Profile) signs the AssumeRole + * call, and the resulting temporary credentials sign Bedrock requests. + */ + readonly role_arn?: string; } // From codersdk/aiproviders_bedrock.go