test: extract AI Gateway test helpers for chatd (#26639)

Extracts test infrastructure for AI Gateway routing into shared helpers
under a new package `coderd/aibridgedtest` so both AGPL and enterprise
tests can use them.

- aibridgedtest.StartTestAIBridgeDaemon` spins up a real in-process 
  aibridged daemon wired to fake upstream providers.
- `chattest.MockAIBridgeTransport` is a mock `aibridge.TransportFactory`
   for the 3 bare-chatd tests that use `newActiveTestServer`.

> 🤖 Generated by Coder Agents under the eyes of a human.
This commit is contained in:
Cian Johnston
2026-06-26 14:33:26 +01:00
committed by GitHub
parent 0135f29cd8
commit 387011d725
4 changed files with 239 additions and 166 deletions
+94
View File
@@ -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
}
+17 -104
View File
@@ -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))
}
}
@@ -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)
}