mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user