diff --git a/coderd/aibridgedtest/aibridgedtest.go b/coderd/aibridgedtest/aibridgedtest.go new file mode 100644 index 0000000000..9ae3511566 --- /dev/null +++ b/coderd/aibridgedtest/aibridgedtest.go @@ -0,0 +1,94 @@ +//go:build !slim + +// Package aibridgedtest provides helpers for starting an in-process +// aibridged daemon in tests. +package aibridgedtest + +import ( + "context" + "testing" + + "github.com/prometheus/client_golang/prometheus" + "go.opentelemetry.io/otel" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/cli" + "github.com/coder/coder/v2/coderd" + "github.com/coder/coder/v2/coderd/aibridged" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/codersdk" +) + +// StartTestAIBridgeDaemon wires an in-process aibridged daemon onto the +// supplied API, mirroring what cli/server.go does in production. Tests that +// create AI provider rows with BaseURL pointing at fake upstream HTTP servers +// (e.g. chattest.NewOpenAI) will have their requests proxied through the real +// aibridged stack as they would in production. +// +// metrics is the registry the daemon reports provider reload events to. +// The caller owns the metrics instance and can assert on it after the daemon +// runs. Use [aibridged.NewMetrics] to create one, or nil for a throwaway. +func StartTestAIBridgeDaemon( + ctx context.Context, + t testing.TB, + api *coderd.API, + metrics *aibridged.Metrics, +) { + t.Helper() + + logger := api.Logger.Named("aibridged").Leveled(slog.LevelDebug) + cfg := api.DeploymentValues.AI.BridgeConfig + tracer := otel.Tracer("aibridge-test") + + providers, _, err := cli.BuildProviders(ctx, api.Database, cfg, logger, nil) + if err != nil { + t.Fatalf("build providers: %v", err) + } + + pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger.Named("pool"), nil, tracer) + if err != nil { + t.Fatalf("create bridge pool: %v", err) + } + t.Cleanup(func() { _ = pool.Shutdown(context.Background()) }) + + if metrics == nil { + metrics = aibridged.NewMetrics(prometheus.NewRegistry()) + } + reloader := &testPoolReloader{pool: pool, db: api.Database, cfg: cfg, logger: logger.Named("reloader"), metrics: metrics} + unsubscribe, err := aibridged.SubscribeProviderReload(ctx, api.Pubsub, reloader, logger.Named("subscriber")) + if err != nil { + t.Fatalf("subscribe provider reload: %v", err) + } + t.Cleanup(unsubscribe) + + srv, err := aibridged.New(ctx, pool, func(dialCtx context.Context) (aibridged.DRPCClient, error) { + return api.CreateInMemoryAIBridgeServer(dialCtx) + }, logger, tracer) + if err != nil { + t.Fatalf("create aibridged server: %v", err) + } + t.Cleanup(func() { _ = srv.Close() }) + + api.RegisterInMemoryAIBridgedHTTPHandler(srv) +} + +type testPoolReloader struct { + pool *aibridged.CachedBridgePool + db database.Store + cfg codersdk.AIBridgeConfig + logger slog.Logger + metrics *aibridged.Metrics +} + +func (r *testPoolReloader) Reload(ctx context.Context) error { + // Stamp the attempt before building providers so the gap between + // attempt and success timestamps reveals a mid-reload hang. + r.metrics.RecordReloadAttempt() + providers, outcomes, err := cli.BuildProviders(ctx, r.db, r.cfg, r.logger, nil) + if err != nil { + return err + } + r.pool.ReplaceProviders(providers) + r.metrics.RecordReloadSuccess(outcomes) + return nil +} diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index acf1609ecc..c36aea4efb 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -11,7 +11,6 @@ import ( "io" "net/http" "net/http/httptest" - "net/url" "os" "path/filepath" "slices" @@ -78,87 +77,6 @@ func testAPIKeyID(t testing.TB, db database.Store, userID uuid.UUID) string { return key.ID } -type chatAIGatewayRecordedRequest struct { - ProviderName string - Source aibridge.Source - APIKeyID string - Path string - Authorization string - XAPIKey string - CoderToken string -} - -type chatAIGatewayTestFactory struct { - target *url.URL - transport http.RoundTripper - preservePath bool - mu sync.Mutex - requests []chatAIGatewayRecordedRequest -} - -func newChatAIGatewayTestFactory(t testing.TB, targetBaseURL string) *chatAIGatewayTestFactory { - t.Helper() - - target, err := url.Parse(targetBaseURL) - require.NoError(t, err) - return &chatAIGatewayTestFactory{target: target, transport: http.DefaultTransport} -} - -func newChatAIGatewayPreservePathTestFactory(t testing.TB, targetBaseURL string) *chatAIGatewayTestFactory { - t.Helper() - - target, err := url.Parse(targetBaseURL) - require.NoError(t, err) - return &chatAIGatewayTestFactory{target: target, transport: http.DefaultTransport, preservePath: true} -} - -func (f *chatAIGatewayTestFactory) TransportFor(providerName string, source aibridge.Source) (http.RoundTripper, error) { - return chatAIGatewayRoundTripper{factory: f, providerName: providerName, source: source}, nil -} - -func (f *chatAIGatewayTestFactory) requestsSnapshot() []chatAIGatewayRecordedRequest { - f.mu.Lock() - defer f.mu.Unlock() - return slices.Clone(f.requests) -} - -type chatAIGatewayRoundTripper struct { - factory *chatAIGatewayTestFactory - providerName string - source aibridge.Source -} - -func (t chatAIGatewayRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { - apiKeyID, _ := aibridge.DelegatedAPIKeyIDFromContext(req.Context()) - t.factory.mu.Lock() - t.factory.requests = append(t.factory.requests, chatAIGatewayRecordedRequest{ - ProviderName: t.providerName, - Source: t.source, - APIKeyID: apiKeyID, - Path: req.URL.Path, - Authorization: req.Header.Get("Authorization"), - XAPIKey: req.Header.Get("X-Api-Key"), - CoderToken: req.Header.Get(aibridge.HeaderCoderToken), - }) - t.factory.mu.Unlock() - - targetURL := *t.factory.target - if t.factory.preservePath { - targetURL.Path = req.URL.Path - } else { - targetURL.Path = strings.TrimPrefix(req.URL.Path, "/v1") - if targetURL.Path == "" { - targetURL.Path = "/" - } - } - targetURL.RawQuery = req.URL.RawQuery - - cloned := req.Clone(req.Context()) - cloned.URL = &targetURL - cloned.Host = t.factory.target.Host - return t.factory.transport.RoundTrip(cloned) -} - func chatAIGatewayTransportFactoryPointer(factory aibridge.TransportFactory) *atomic.Pointer[aibridge.TransportFactory] { var ptr atomic.Pointer[aibridge.TransportFactory] ptr.Store(&factory) @@ -5401,7 +5319,7 @@ func TestActiveServer_AIGatewayRoutingPreservesAPIKeyAfterCompaction(t *testing. return chattest.AnthropicStreamingResponse() } }) - factory := newChatAIGatewayPreservePathTestFactory(t, anthropicURL) + factory := chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()) user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL) model = updateChatModelCompressionThreshold(t, db, model, contextLimit, thresholdPercent) provider, err := db.GetAIProviderByID(ctx, model.AIProviderID.UUID) @@ -5480,14 +5398,14 @@ func TestActiveServer_AIGatewayRoutingPreservesAPIKeyAfterCompaction(t *testing. require.True(t, compressed.summaries[0].APIKeyID.Valid) require.Equal(t, apiKey.ID, compressed.summaries[0].APIKeyID.String) - requests := factory.requestsSnapshot() + requests := factory.RequestsSnapshot() require.NotEmpty(t, requests) for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) require.Equal(t, apiKey.ID, req.APIKeyID) - require.Equal(t, "sk-user-aibridge", req.XAPIKey) - require.Equal(t, "delegated", req.CoderToken) + require.Equal(t, "sk-user-aibridge", req.Request.Header.Get("X-Api-Key")) + require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) } } @@ -9790,7 +9708,7 @@ func TestProcessChat_AIGatewayRoutingUsesDelegatedAPIKey(t *testing.T) { } return chattest.OpenAINonStreamingResponse(`{"title":"AI Gateway Chat"}`) }) - factory := newChatAIGatewayTestFactory(t, openAIURL) + factory := chattest.NewMockAIBridgeTransport(t, openAIURL) user, org, provider, model, apiKey := seedAIGatewayOpenAITestDependencies(t, db, openAIURL) @@ -9824,24 +9742,19 @@ func TestProcessChat_AIGatewayRoutingUsesDelegatedAPIKey(t *testing.T) { require.Equal(t, database.ChatStatusWaiting, chatResult.Status) require.False(t, chatResult.LastError.Valid) - requests := factory.requestsSnapshot() + requests := factory.RequestsSnapshot() require.NotEmpty(t, requests) - require.Contains(t, requests, chatAIGatewayRecordedRequest{ - ProviderName: provider.Name, - Source: aibridge.SourceAgents, - APIKeyID: apiKey.ID, - Path: "/v1/responses", - Authorization: "Bearer sk-user-aibridge", - CoderToken: "delegated", - }) + require.True(t, slices.ContainsFunc(requests, func(req chattest.RecordedRequest) bool { + return req.Request.URL.Path == "/v1/responses" + }), "no request to /v1/responses found") for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) require.Equal(t, apiKey.ID, req.APIKeyID) - require.Equal(t, "Bearer sk-user-aibridge", req.Authorization) - require.Empty(t, req.XAPIKey) - require.Equal(t, "delegated", req.CoderToken) - require.True(t, strings.HasPrefix(req.Path, "/v1/"), "unexpected aibridge path %q", req.Path) + require.Equal(t, "Bearer sk-user-aibridge", req.Request.Header.Get("Authorization")) + require.Empty(t, req.Request.Header.Get("X-Api-Key")) + require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) + require.True(t, strings.HasPrefix(req.Request.URL.Path, "/v1/"), "unexpected aibridge path %q", req.Request.URL.Path) } } @@ -9859,7 +9772,7 @@ func TestProcessChat_AIGatewayRoutingPreservesAPIKeyAfterWorkspaceContext(t *tes } return chattest.OpenAINonStreamingResponse(`{"title":"AI Gateway Workspace"}`) }) - factory := newChatAIGatewayTestFactory(t, openAIURL) + factory := chattest.NewMockAIBridgeTransport(t, openAIURL) user, org, provider, model, apiKey := seedAIGatewayOpenAITestDependencies(t, db, openAIURL) ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID) @@ -9908,14 +9821,14 @@ func TestProcessChat_AIGatewayRoutingPreservesAPIKeyAfterWorkspaceContext(t *tes require.NoError(t, err) require.NotEmpty(t, pinned, "workspace context should be pinned to the chat") - requests := factory.requestsSnapshot() + requests := factory.RequestsSnapshot() require.NotEmpty(t, requests) for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) require.Equal(t, apiKey.ID, req.APIKeyID) - require.Equal(t, "Bearer sk-user-aibridge", req.Authorization) - require.Equal(t, "delegated", req.CoderToken) + require.Equal(t, "Bearer sk-user-aibridge", req.Request.Header.Get("Authorization")) + require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) } } diff --git a/coderd/x/chatd/chattest/mock_aibridge_transport.go b/coderd/x/chatd/chattest/mock_aibridge_transport.go new file mode 100644 index 0000000000..25b1b1c9ea --- /dev/null +++ b/coderd/x/chatd/chattest/mock_aibridge_transport.go @@ -0,0 +1,125 @@ +package chattest + +import ( + "net/http" + "net/url" + "slices" + "strings" + "sync" + "testing" + + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/coderd/aibridge" +) + +// RecordedRequest captures metadata from a single request that passed +// through the mock transport factory. Fields not already available on +// [http.Request] are included here; tests can access headers and path +// via [RecordedRequest.Request]. +type RecordedRequest struct { + // Request is a clone of the original [http.Request]. + Request *http.Request + // ProviderName is the AI provider instance name passed to + // [TransportFactory.TransportFor]. + ProviderName string + // Source is the aibridge source passed to TransportFor. + Source aibridge.Source + // APIKeyID is the delegated API key ID attached to request ctx. + APIKeyID string +} + +// MockAIBridgeTransportOption configures a [MockAIBridgeTransport]. +type MockAIBridgeTransportOption func(*MockAIBridgeTransport) + +// WithPreservePath disables the default "/v1" path stripping so the +// target server receives the full original request path. +func WithPreservePath() MockAIBridgeTransportOption { + return func(f *MockAIBridgeTransport) { f.preservePath = true } +} + +// MockAIBridgeTransport is a test [aibridge.TransportFactory] that +// redirects requests to a target URL (typically a [chattest.NewOpenAI] +// or [chattest.NewAnthropic] server) and records each request for +// later inspection. +// +// By default it strips the leading "/v1" path segment before +// forwarding, matching how the real AI Gateway transport rewrites +// upstream-shaped requests. Pass [WithPreservePath] when the target +// server expects the full original path. +type MockAIBridgeTransport struct { + target *url.URL + transport http.RoundTripper + preservePath bool + mu sync.Mutex + requests []RecordedRequest +} + +// NewMockAIBridgeTransport creates a [MockAIBridgeTransport] that +// forwards to targetBaseURL. +func NewMockAIBridgeTransport(t testing.TB, targetBaseURL string, opts ...MockAIBridgeTransportOption) *MockAIBridgeTransport { + t.Helper() + target, err := url.Parse(targetBaseURL) + if err != nil { + t.Fatalf("parse target URL: %v", err) + } + f := &MockAIBridgeTransport{target: target, transport: http.DefaultTransport} + for _, opt := range opts { + opt(f) + } + return f +} + +// TransportFor implements [aibridge.TransportFactory]. +func (f *MockAIBridgeTransport) TransportFor(providerName string, source aibridge.Source) (http.RoundTripper, error) { + if len(providerName) == 0 { + return nil, xerrors.New("provider name is required") + } + return mockRoundTripper{factory: f, providerName: providerName, source: source}, nil +} + +// RequestsSnapshot returns a copy of all recorded requests. +func (f *MockAIBridgeTransport) RequestsSnapshot() []RecordedRequest { + f.mu.Lock() + defer f.mu.Unlock() + return slices.Clone(f.requests) +} + +type mockRoundTripper struct { + factory *MockAIBridgeTransport + providerName string + source aibridge.Source +} + +func (rt mockRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + // Mirror the real aibridged transport: a delegated API key must be + // on the context, otherwise aibridged has no identity to act under. + apiKeyID, ok := aibridge.DelegatedAPIKeyIDFromContext(req.Context()) + if !ok { + return nil, xerrors.New("mock aibridged transport requires WithDelegatedAPIKeyID on the request context") + } + rt.factory.mu.Lock() + rt.factory.requests = append(rt.factory.requests, RecordedRequest{ + Request: req.Clone(req.Context()), + ProviderName: rt.providerName, + Source: rt.source, + APIKeyID: apiKeyID, + }) + rt.factory.mu.Unlock() + + targetURL := *rt.factory.target + if rt.factory.preservePath { + targetURL.Path = req.URL.Path + } else { + targetURL.Path = strings.TrimPrefix(req.URL.Path, "/v1") + if targetURL.Path == "" { + targetURL.Path = "/" + } + } + targetURL.RawQuery = req.URL.RawQuery + + cloned := req.Clone(req.Context()) + cloned.URL = &targetURL + cloned.Host = rt.factory.target.Host + return rt.factory.transport.RoundTrip(cloned) +} diff --git a/enterprise/coderd/aibridge_reload_test.go b/enterprise/coderd/aibridge_reload_test.go index d911d3c05a..b4e605dede 100644 --- a/enterprise/coderd/aibridge_reload_test.go +++ b/enterprise/coderd/aibridge_reload_test.go @@ -1,7 +1,6 @@ package coderd_test import ( - "context" "encoding/json" "io" "net/http" @@ -13,15 +12,10 @@ import ( promtest "github.com/prometheus/client_golang/prometheus/testutil" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "go.opentelemetry.io/otel" - "cdr.dev/slog/v3" - "cdr.dev/slog/v3/sloggers/slogtest" - "github.com/coder/coder/v2/cli" - "github.com/coder/coder/v2/coderd" "github.com/coder/coder/v2/coderd/aibridged" + "github.com/coder/coder/v2/coderd/aibridgedtest" "github.com/coder/coder/v2/coderd/coderdtest" - "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" "github.com/coder/coder/v2/enterprise/coderd/license" @@ -52,60 +46,6 @@ func newMockUpstream(t *testing.T, name string) *mockUpstream { return m } -// startTestAIBridgeDaemon wires an in-process aibridged daemon onto -// the supplied API and subscribes it to ai_providers change events. -// This mirrors what cli/server.go does in production so /api/v2/ai-gateway -// requests dispatch through the real pool and reloader. -func startTestAIBridgeDaemon(t *testing.T, api *coderd.API) *aibridged.Metrics { - t.Helper() - - ctx := context.Background() - logger := slogtest.Make(t, nil).Named("aibridged").Leveled(slog.LevelDebug) - cfg := api.DeploymentValues.AI.BridgeConfig - tracer := otel.Tracer("aibridge-reload-test") - - providers, _, err := cli.BuildProviders(ctx, api.Database, cfg, logger, nil) - require.NoError(t, err) - - pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger.Named("pool"), nil, tracer) - require.NoError(t, err) - t.Cleanup(func() { _ = pool.Shutdown(context.Background()) }) - - metrics := aibridged.NewMetrics(prometheus.NewRegistry()) - reloader := &testPoolReloader{pool: pool, db: api.Database, cfg: cfg, logger: logger.Named("reloader"), metrics: metrics} - unsubscribe, err := aibridged.SubscribeProviderReload(ctx, api.Pubsub, reloader, logger.Named("subscriber")) - require.NoError(t, err) - t.Cleanup(unsubscribe) - - srv, err := aibridged.New(ctx, pool, func(dialCtx context.Context) (aibridged.DRPCClient, error) { - return api.CreateInMemoryAIBridgeServer(dialCtx) - }, logger, tracer) - require.NoError(t, err) - t.Cleanup(func() { _ = srv.Close() }) - - api.RegisterInMemoryAIBridgedHTTPHandler(srv) - return metrics -} - -type testPoolReloader struct { - pool *aibridged.CachedBridgePool - db database.Store - cfg codersdk.AIBridgeConfig - logger slog.Logger - metrics *aibridged.Metrics -} - -func (r *testPoolReloader) Reload(ctx context.Context) error { - defer r.metrics.RecordReloadAttempt() - providers, outcomes, err := cli.BuildProviders(ctx, r.db, r.cfg, r.logger, nil) - if err != nil { - return err - } - r.pool.ReplaceProviders(providers) - r.metrics.RecordReloadSuccess(outcomes) - return nil -} - // TestAIBridgeProviderHotReload exercises the end-to-end CRUD -> // reload -> routing path: every provider mutation made through codersdk // must, within a short window, change the routing observed at @@ -131,7 +71,8 @@ func TestAIBridgeProviderHotReload(t *testing.T) { }, }) - metrics := startTestAIBridgeDaemon(t, api.AGPL) + metrics := aibridged.NewMetrics(prometheus.NewRegistry()) + aibridgedtest.StartTestAIBridgeDaemon(testutil.Context(t, testutil.WaitLong), t, api.AGPL, metrics) // requireProviderStatus polls until the provider_info series for // (name, status) settles to value 1. Reloads happen via pubsub, so