refactor: move aibridged out of enterprise to AGPL (#25570)

In order to allow Coder Agents to use AI Gateway in OSS, we need to rehome the `aibridged`\-related code into the AGPL path.

The HTTP API is only registered under enterprise so will still require the AI Governance Add-on to be present in order to use it, whereas Coder Agents uses an in-memory pipe to the same handlers.
This commit is contained in:
Danny Kopping
2026-05-22 09:11:37 +02:00
committed by GitHub
parent c50b0e84b9
commit ddec110b0e
32 changed files with 631 additions and 603 deletions
+111
View File
@@ -0,0 +1,111 @@
package coderd
import (
"context"
"errors"
"io"
"net/http"
"golang.org/x/xerrors"
"storj.io/drpc/drpcmux"
"storj.io/drpc/drpcserver"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/aibridged"
aibridgedproto "github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/coderd/aibridgedserver"
"github.com/coder/coder/v2/coderd/tracing"
"github.com/coder/coder/v2/codersdk/drpcsdk"
)
// GetAIBridgedHandler returns the in-memory aibridge HTTP handler set by
// [API.RegisterInMemoryAIBridgedHTTPHandler], or nil if the daemon has not
// been wired in. Used by the enterprise /api/v2/aibridge route (license-gated)
// to forward requests into the same in-memory handler that chatd dispatches
// to in-process.
func (api *API) GetAIBridgedHandler() http.Handler {
return api.aibridgedHandler
}
// RegisterInMemoryAIBridgedHTTPHandler mounts [aibridged.Server]'s HTTP router onto
// [API]'s router, so that requests to aibridged will be relayed from Coder's API server
// to the in-memory aibridged.
func (api *API) RegisterInMemoryAIBridgedHTTPHandler(srv http.Handler) {
if srv == nil {
panic("aibridged cannot be nil")
}
api.aibridgedHandler = srv
}
// CreateInMemoryAIBridgeServer creates a [aibridged.DRPCServer] and returns a
// [aibridged.DRPCClient] to it, connected over an in-memory transport.
// This server is responsible for all the Coder-specific functionality that aibridged
// requires such as persistence and retrieving configuration.
func (api *API) CreateInMemoryAIBridgeServer(dialCtx context.Context) (client aibridged.DRPCClient, err error) {
// TODO(dannyk): implement options.
// TODO(dannyk): implement tracing.
// TODO(dannyk): implement API versioning.
clientSession, serverSession := drpcsdk.MemTransportPipe()
defer func() {
if err != nil {
_ = clientSession.Close()
_ = serverSession.Close()
}
}()
mux := drpcmux.New()
srv, err := aibridgedserver.NewServer(api.ctx, api.Database, api.Logger.Named("aibridgedserver"),
api.AccessURL.String(), api.DeploymentValues.AI.BridgeConfig, api.ExternalAuthConfigs, api.Experiments, api.AISeatTracker)
if err != nil {
return nil, err
}
err = aibridgedproto.DRPCRegisterRecorder(mux, srv)
if err != nil {
return nil, xerrors.Errorf("register recorder service: %w", err)
}
err = aibridgedproto.DRPCRegisterMCPConfigurator(mux, srv)
if err != nil {
return nil, xerrors.Errorf("register MCP configurator service: %w", err)
}
err = aibridgedproto.DRPCRegisterAuthorizer(mux, srv)
if err != nil {
return nil, xerrors.Errorf("register key validator service: %w", err)
}
server := drpcserver.NewWithOptions(&tracing.DRPCHandler{Handler: mux},
drpcserver.Options{
Manager: drpcsdk.DefaultDRPCOptions(nil),
Log: func(err error) {
if errors.Is(err, io.EOF) {
return
}
api.Logger.Debug(dialCtx, "aibridged drpc server error", slog.Error(err))
},
},
)
// in-mem pipes aren't technically "websockets" but they have the same properties as far as the
// API is concerned: they are long-lived connections that we need to close before completing
// shutdown of the API.
api.WebsocketWaitMutex.Lock()
api.WebsocketWaitGroup.Add(1)
api.WebsocketWaitMutex.Unlock()
go func() {
defer api.WebsocketWaitGroup.Done()
// Here we pass the background context, since we want the server to keep serving until the
// client hangs up. The aibridged is local, in-mem, so there isn't a danger of losing contact with it and
// having a dead connection we don't know the status of.
err := server.Serve(context.Background(), serverSession)
api.Logger.Info(dialCtx, "aibridge daemon disconnected", slog.Error(err))
// Close the sessions, so we don't leak goroutines serving them.
_ = clientSession.Close()
_ = serverSession.Close()
}()
return &aibridged.Client{
Conn: clientSession,
DRPCRecorderClient: aibridgedproto.NewDRPCRecorderClient(clientSession),
DRPCMCPConfiguratorClient: aibridgedproto.NewDRPCMCPConfiguratorClient(clientSession),
DRPCAuthorizerClient: aibridgedproto.NewDRPCAuthorizerClient(clientSession),
}, nil
}
+199
View File
@@ -0,0 +1,199 @@
package aibridged
import (
"context"
"errors"
"io"
"net/http"
"sync"
"time"
"go.opentelemetry.io/otel/trace"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/retry"
)
var _ io.Closer = &Server{}
// Server provides the AI Bridge functionality.
// It is responsible for:
// - receiving requests on /api/v2/aibridged/*
// - manipulating the requests
// - relaying requests to upstream AI services and relaying responses to caller
//
// It requires a [Dialer] to provide a [DRPCClient] implementation to
// communicate with a [DRPCServer] implementation, to persist state and perform other functions.
type Server struct {
clientDialer Dialer
clientCh chan DRPCClient
// A pool of [aibridge.RequestBridge] instances, which service incoming requests.
requestBridgePool Pooler
logger slog.Logger
tracer trace.Tracer
wg sync.WaitGroup
// initConnectionCh will receive when the daemon connects to coderd for the
// first time.
initConnectionCh chan struct{}
initConnectionOnce sync.Once
// lifecycleCtx is canceled when we start closing.
lifecycleCtx context.Context
// cancelFn closes the lifecycleCtx.
cancelFn func()
shutdownOnce sync.Once
}
func New(ctx context.Context, pool Pooler, rpcDialer Dialer, logger slog.Logger, tracer trace.Tracer) (*Server, error) {
if rpcDialer == nil {
return nil, xerrors.Errorf("nil rpcDialer given")
}
ctx, cancel := context.WithCancel(ctx)
daemon := &Server{
logger: logger,
tracer: tracer,
clientDialer: rpcDialer,
clientCh: make(chan DRPCClient),
lifecycleCtx: ctx,
cancelFn: cancel,
initConnectionCh: make(chan struct{}),
requestBridgePool: pool,
}
daemon.wg.Add(1)
go daemon.connect()
return daemon, nil
}
// Connect establishes a connection to coderd.
func (s *Server) connect() {
defer s.logger.Debug(s.lifecycleCtx, "connect loop exited")
defer s.wg.Done()
logConnect := s.logger.With(slog.F("context", "aibridged.server")).Debug
// An exponential back-off occurs when the connection is failing to dial.
// This is to prevent server spam in case of a coderd outage.
connectLoop:
for retrier := retry.New(50*time.Millisecond, 10*time.Second); retrier.Wait(s.lifecycleCtx); {
// It's possible for the aibridge daemon to be shut down
// before the wait is complete!
if s.isShutdown() {
return
}
s.logger.Debug(s.lifecycleCtx, "dialing coderd")
client, err := s.clientDialer(s.lifecycleCtx)
if err != nil {
if errors.Is(err, context.Canceled) {
return
}
var sdkErr *codersdk.Error
// If something is wrong with our auth, stop trying to connect.
if errors.As(err, &sdkErr) && sdkErr.StatusCode() == http.StatusForbidden {
s.logger.Error(s.lifecycleCtx, "not authorized to dial coderd", slog.Error(err))
return
}
if s.isShutdown() {
return
}
s.logger.Warn(s.lifecycleCtx, "coderd client failed to dial", slog.Error(err))
continue
}
// TODO: log this with INFO level when we implement external aibridge daemons.
logConnect(s.lifecycleCtx, "successfully connected to coderd")
retrier.Reset()
s.initConnectionOnce.Do(func() {
close(s.initConnectionCh)
})
// Serve the client until we are closed or it disconnects.
for {
select {
case <-s.lifecycleCtx.Done():
client.DRPCConn().Close()
return
case <-client.DRPCConn().Closed():
logConnect(s.lifecycleCtx, "connection to coderd closed")
continue connectLoop
case s.clientCh <- client:
continue
}
}
}
}
func (s *Server) Client() (DRPCClient, error) {
select {
case <-s.lifecycleCtx.Done():
return nil, xerrors.New("context closed")
case client := <-s.clientCh:
return client, nil
}
}
// GetRequestHandler retrieves a (possibly reused) [*aibridge.RequestBridge] from the pool, for the given user.
func (s *Server) GetRequestHandler(ctx context.Context, req Request) (http.Handler, error) {
if s.requestBridgePool == nil {
return nil, xerrors.New("nil requestBridgePool")
}
reqBridge, err := s.requestBridgePool.Acquire(ctx, req, s.Client, NewMCPProxyFactory(s.logger, s.tracer, s.Client))
if err != nil {
return nil, xerrors.Errorf("acquire request bridge: %w", err)
}
return reqBridge, nil
}
// isShutdown returns whether the Server is shutdown or not.
func (s *Server) isShutdown() bool {
select {
case <-s.lifecycleCtx.Done():
return true
default:
return false
}
}
// Shutdown waits for all exiting in-flight requests to complete, or the context to expire, whichever comes first.
func (s *Server) Shutdown(ctx context.Context) error {
var err error
s.shutdownOnce.Do(func() {
s.cancelFn()
// Wait for any outstanding connections to terminate.
s.wg.Wait()
select {
case <-ctx.Done():
s.logger.Warn(ctx, "graceful shutdown failed", slog.Error(ctx.Err()))
err = ctx.Err()
return
default:
}
s.logger.Info(ctx, "shutting down request pool")
if err = s.requestBridgePool.Shutdown(ctx); err != nil {
s.logger.Error(ctx, "request pool shutdown failed with error", slog.Error(err))
}
s.logger.Info(ctx, "gracefully shutdown")
})
return err
}
// Close shuts down the server with a timeout of 5s.
func (s *Server) Close() error {
ctx, cancel := context.WithTimeout(context.Background(), time.Second*5)
defer cancel()
return s.Shutdown(ctx)
}
+638
View File
@@ -0,0 +1,638 @@
package aibridged_test
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"testing"
"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"golang.org/x/xerrors"
"storj.io/drpc"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/aibridge/intercept"
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/aibridged"
mock "github.com/coder/coder/v2/coderd/aibridged/aibridgedmock"
"github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
)
func newTestServer(t *testing.T) (*aibridged.Server, *mock.MockDRPCClient, *mock.MockPooler) {
t.Helper()
logger := slogtest.Make(t, nil)
ctrl := gomock.NewController(t)
client := mock.NewMockDRPCClient(ctrl)
pool := mock.NewMockPooler(ctrl)
conn := &mockDRPCConn{}
client.EXPECT().DRPCConn().AnyTimes().Return(conn)
pool.EXPECT().Shutdown(gomock.Any()).MinTimes(1).Return(nil)
srv, err := aibridged.New(
t.Context(),
pool,
func(ctx context.Context) (aibridged.DRPCClient, error) {
return client, nil
}, logger, testTracer)
require.NoError(t, err, "create new aibridged")
t.Cleanup(func() {
srv.Shutdown(context.Background())
})
return srv, client, pool
}
// mockDRPCConn is a mock implementation of drpc.Conn
type mockDRPCConn struct{}
func (*mockDRPCConn) Close() error { return nil }
func (*mockDRPCConn) Closed() <-chan struct{} { ch := make(chan struct{}); return ch }
func (*mockDRPCConn) Transport() drpc.Transport { return nil }
func (*mockDRPCConn) Invoke(ctx context.Context, rpc string, enc drpc.Encoding, in, out drpc.Message) error {
return nil
}
func (*mockDRPCConn) NewStream(ctx context.Context, rpc string, enc drpc.Encoding) (drpc.Stream, error) {
// nolint:nilnil // Chillchill.
return nil, nil
}
func TestServeHTTP_FailureModes(t *testing.T) {
t.Parallel()
defaultHeaders := map[string]string{"Authorization": "Bearer key"}
httpClient := &http.Client{}
cases := []struct {
name string
reqHeaders map[string]string
applyMocksFn func(client *mock.MockDRPCClient, pool *mock.MockPooler)
dialerFn aibridged.Dialer
contextFn func() context.Context
expectedErr error
expectedStatus int
}{
// Authnz-related failures.
{
name: "no auth key",
reqHeaders: make(map[string]string),
expectedErr: aibridged.ErrNoAuthKey,
expectedStatus: http.StatusBadRequest,
},
{
name: "unrecognized header",
reqHeaders: map[string]string{
codersdk.SessionTokenHeader: "key", // Coder-Session-Token is not supported; requests originate with AI clients, not coder CLI.
},
applyMocksFn: func(client *mock.MockDRPCClient, _ *mock.MockPooler) {},
expectedErr: aibridged.ErrNoAuthKey,
expectedStatus: http.StatusBadRequest,
},
{
name: "unauthorized",
applyMocksFn: func(client *mock.MockDRPCClient, _ *mock.MockPooler) {
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(nil, xerrors.New("not authorized"))
},
expectedErr: aibridged.ErrUnauthorized,
expectedStatus: http.StatusForbidden,
},
{
name: "invalid key owner ID",
applyMocksFn: func(client *mock.MockDRPCClient, _ *mock.MockPooler) {
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.IsAuthorizedResponse{OwnerId: "oops"}, nil)
},
expectedErr: aibridged.ErrUnauthorized,
expectedStatus: http.StatusForbidden,
},
// TODO: coderd connection-related failures.
// Pool-related failures.
{
name: "pool instance",
applyMocksFn: func(client *mock.MockDRPCClient, pool *mock.MockPooler) {
// Should pass authorization.
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.IsAuthorizedResponse{OwnerId: uuid.NewString()}, nil)
// But fail when acquiring a pool instance.
pool.EXPECT().Acquire(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes().Return(nil, xerrors.New("oops"))
},
expectedErr: aibridged.ErrAcquireRequestHandler,
expectedStatus: http.StatusInternalServerError,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
srv, client, pool := newTestServer(t)
conn := &mockDRPCConn{}
client.EXPECT().DRPCConn().AnyTimes().Return(conn)
if tc.applyMocksFn != nil {
tc.applyMocksFn(client, pool)
}
httpSrv := httptest.NewServer(srv)
ctx := testutil.Context(t, testutil.WaitShort)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, httpSrv.URL+"/openai/v1/chat/completions", nil)
require.NoError(t, err, "make request to test server")
headers := defaultHeaders
if tc.reqHeaders != nil {
headers = tc.reqHeaders
}
for k, v := range headers {
req.Header.Set(k, v)
}
resp, err := httpClient.Do(req)
t.Cleanup(func() {
if resp == nil || resp.Body == nil {
return
}
resp.Body.Close()
})
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
require.NoError(t, err, "read response body")
require.Contains(t, string(body), tc.expectedErr.Error())
require.Equal(t, tc.expectedStatus, resp.StatusCode)
})
}
}
func TestServeHTTP_StripCoderToken(t *testing.T) {
t.Parallel()
cases := []struct {
name string
reqHeaders map[string]string
expectPresent map[string]string // header → expected value
expectAbsent []string // headers that must be gone
}{
{
// Centralized: the client sets Authorization and X-Api-Key,
// but does not include HeaderCoderToken.
// All auth headers are stripped.
name: "centralized",
reqHeaders: map[string]string{
"Authorization": "Bearer coder-token",
"X-Api-Key": "sk-ant-api03-user-key",
},
expectAbsent: []string{
"Authorization",
"X-Api-Key",
agplaibridge.HeaderCoderToken,
},
},
{
// BYOK with access token: Coder token in BYOK header,
// user's access token in Authorization. Only the
// BYOK header is stripped.
name: "byok bearer token",
reqHeaders: map[string]string{
agplaibridge.HeaderCoderToken: "coder-token",
"Authorization": "Bearer sk-ant-oat01-user-oauth-token",
},
expectPresent: map[string]string{
"Authorization": "Bearer sk-ant-oat01-user-oauth-token",
},
expectAbsent: []string{
agplaibridge.HeaderCoderToken,
},
},
{
// BYOK with personal API key: Coder token in BYOK header,
// user's API key in X-Api-Key. Only the BYOK header is
// stripped.
name: "byok api key",
reqHeaders: map[string]string{
agplaibridge.HeaderCoderToken: "coder-token",
"X-Api-Key": "sk-ant-api03-user-key",
},
expectPresent: map[string]string{
"X-Api-Key": "sk-ant-api03-user-key",
},
expectAbsent: []string{
agplaibridge.HeaderCoderToken,
},
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
mockH := &mockHandler{}
srv, client, pool := newTestServer(t)
conn := &mockDRPCConn{}
client.EXPECT().DRPCConn().AnyTimes().Return(conn)
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.IsAuthorizedResponse{OwnerId: uuid.NewString()}, nil)
pool.EXPECT().Acquire(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes().Return(mockH, nil)
httpSrv := httptest.NewServer(srv)
t.Cleanup(httpSrv.Close)
ctx := testutil.Context(t, testutil.WaitShort)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, httpSrv.URL+"/openai/v1/chat/completions", nil)
require.NoError(t, err)
for k, v := range tc.reqHeaders {
req.Header.Set(k, v)
}
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode)
require.NotNil(t, mockH.headersReceived)
for header, expected := range tc.expectPresent {
require.Equal(t, expected, mockH.headersReceived.Get(header),
"header %q should be preserved with value %q", header, expected)
}
for _, header := range tc.expectAbsent {
require.Empty(t, mockH.headersReceived.Get(header),
"header %q should be stripped", header)
}
// HeaderCoderToken should always be stripped
require.Empty(t, mockH.headersReceived.Get(agplaibridge.HeaderCoderToken),
"header %q should be stripped", agplaibridge.HeaderCoderToken)
})
}
}
func TestExtractAuthToken(t *testing.T) {
t.Parallel()
cases := []struct {
name string
headers map[string]string
expectedKey string
}{
{
name: "none",
},
{
name: "authorization/invalid",
headers: map[string]string{"authorization": "invalid"},
},
{
name: "authorization/bearer empty",
headers: map[string]string{"authorization": "bearer"},
},
{
name: "authorization/bearer ok",
headers: map[string]string{"authorization": "bearer key"},
expectedKey: "key",
},
{
name: "authorization/case",
headers: map[string]string{"AUTHORIZATION": "BEARer key"},
expectedKey: "key",
},
{
name: "authorization/priority over x-api-key",
headers: map[string]string{
"Authorization": "Bearer auth-token",
"X-Api-Key": "api-key",
},
expectedKey: "auth-token",
},
{
name: "x-api-key/empty",
headers: map[string]string{"X-Api-Key": ""},
},
{
name: "x-api-key/ok",
headers: map[string]string{"X-Api-Key": "key"},
expectedKey: "key",
},
// BYOK: X-Coder-AI-Governance-Token carries the Coder
// token and has the highest priority.
{
name: "byok/empty",
headers: map[string]string{agplaibridge.HeaderCoderToken: ""},
},
{
name: "byok/ok",
headers: map[string]string{agplaibridge.HeaderCoderToken: "coder-token"},
expectedKey: "coder-token",
},
{
name: "byok/priority over all",
headers: map[string]string{
agplaibridge.HeaderCoderToken: "coder-token",
"Authorization": "Bearer oauth-token",
"X-Api-Key": "api-key",
},
expectedKey: "coder-token",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
headers := make(http.Header, len(tc.headers))
for k, v := range tc.headers {
headers.Add(k, v)
}
key := agplaibridge.ExtractAuthToken(headers)
require.Equal(t, tc.expectedKey, key)
})
}
}
var _ http.Handler = &mockHandler{}
type mockHandler struct {
headersReceived http.Header
}
func (h *mockHandler) ServeHTTP(rw http.ResponseWriter, r *http.Request) {
h.headersReceived = r.Header.Clone()
rw.WriteHeader(http.StatusOK)
_, _ = rw.Write([]byte(r.URL.Path))
}
// TestServeHTTP_ActorHeaders validates that actor headers are correctly forwarded to
// upstream AI providers when SendActorHeaders is enabled in the provider configuration.
// These headers allow upstream providers to identify the user making the request for
// tracking and auditing purposes.
func TestServeHTTP_ActorHeaders(t *testing.T) {
t.Parallel()
testUsername := "testuser"
testUserID := uuid.New()
cases := []struct {
path string
}{
// Not a complete set of paths; we're not testing the specific APIs - just the provider configs.
{
path: "/openai/v1/chat/completions",
},
{
path: "/anthropic/v1/messages",
},
}
for _, tc := range cases {
t.Run(tc.path, func(t *testing.T) {
t.Parallel()
// Setup mock upstream AI server that captures headers.
var receivedHeaders http.Header
upstreamSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
receivedHeaders = r.Header.Clone()
w.WriteHeader(http.StatusTeapot)
_, _ = w.Write([]byte(`i am a teapot`))
}))
t.Cleanup(upstreamSrv.Close)
// Setup with SendActorHeaders enabled.
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ctrl := gomock.NewController(t)
client := mock.NewMockDRPCClient(ctrl)
// Create providers with SendActorHeaders=true.
providers := []aibridge.Provider{
aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{
BaseURL: upstreamSrv.URL,
SendActorHeaders: true,
}),
aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{
BaseURL: upstreamSrv.URL,
SendActorHeaders: true,
}, nil),
}
pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, nil, testTracer)
require.NoError(t, err)
conn := &mockDRPCConn{}
client.EXPECT().DRPCConn().AnyTimes().Return(conn)
// Return authorization response with user ID and username.
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.IsAuthorizedResponse{
OwnerId: testUserID.String(),
Username: testUsername,
}, nil)
client.EXPECT().GetMCPServerConfigs(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.GetMCPServerConfigsResponse{}, nil)
client.EXPECT().RecordInterception(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.RecordInterceptionResponse{}, nil)
client.EXPECT().RecordInterceptionEnded(gomock.Any(), gomock.Any()).AnyTimes()
// Given: aibridged is started.
srv, err := aibridged.New(t.Context(), pool, func(ctx context.Context) (aibridged.DRPCClient, error) {
return client, nil
}, logger, testTracer)
require.NoError(t, err, "create new aibridged")
t.Cleanup(func() {
_ = srv.Shutdown(testutil.Context(t, testutil.WaitShort))
})
// When: a request is made to aibridged.
ctx := testutil.Context(t, testutil.WaitShort)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, tc.path, bytes.NewBufferString(`{}`))
require.NoError(t, err, "make request to test server")
req.Header.Add("Authorization", "Bearer key")
req.Header.Add("Accept", "application/json")
// When: aibridged handles the request.
rec := httptest.NewRecorder()
srv.ServeHTTP(rec, req)
// Then: the actor headers should be present in the upstream request.
require.NotEmpty(t, receivedHeaders, "upstream server should have received headers")
// Verify the actor ID header is present with the correct value.
actorIDHeader := receivedHeaders.Get(intercept.ActorIDHeader())
assert.Equal(t, testUserID.String(), actorIDHeader, "actor ID header should contain user ID")
// Verify the actor metadata header for username is present.
usernameHeader := receivedHeaders.Get(intercept.ActorMetadataHeader("Username"))
assert.Equal(t, testUsername, usernameHeader, "actor metadata username header should contain username")
})
}
}
// TestRouting validates that a request which originates with aibridged will be handled
// by coder/aibridge's handling logic in a provider-specific manner.
// We must validate that logic that pertains to coder/coder is exercised.
// aibridge will only handle certain routes; we don't need to test these exhaustively
// (that's coder/aibridge's responsibility), but we do need to validate that it handles
// requests correctly.
func TestRouting(t *testing.T) {
t.Parallel()
cases := []struct {
name string
path string
expectedStatus int
expectedHits int // Expected hits to the upstream server.
}{
{
name: "unsupported",
path: "/this-route-does-not-exist",
expectedStatus: http.StatusNotFound,
expectedHits: 0,
},
{
name: "openai chat completions",
path: "/openai/v1/chat/completions",
expectedStatus: http.StatusTeapot, // Nonsense status to indicate server was hit.
expectedHits: 1,
},
{
name: "anthropic messages",
path: "/anthropic/v1/messages",
expectedStatus: http.StatusTeapot, // Nonsense status to indicate server was hit.
expectedHits: 1,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
// Setup mock upstream AI server.
upstreamSrv := &mockAIUpstreamServer{}
openaiSrv := httptest.NewServer(upstreamSrv)
antSrv := httptest.NewServer(upstreamSrv)
t.Cleanup(openaiSrv.Close)
t.Cleanup(antSrv.Close)
// Setup.
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
ctrl := gomock.NewController(t)
client := mock.NewMockDRPCClient(ctrl)
providers := []aibridge.Provider{
aibridge.NewOpenAIProvider(aibridge.OpenAIConfig{BaseURL: openaiSrv.URL}),
aibridge.NewAnthropicProvider(aibridge.AnthropicConfig{BaseURL: antSrv.URL}, nil),
}
pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger, nil, testTracer)
require.NoError(t, err)
conn := &mockDRPCConn{}
client.EXPECT().DRPCConn().AnyTimes().Return(conn)
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.IsAuthorizedResponse{OwnerId: uuid.NewString()}, nil)
client.EXPECT().GetMCPServerConfigs(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.GetMCPServerConfigsResponse{}, nil)
// This is the only recording we really care about in this test. This is called before the provider-specific logic processes
// the incoming request, and anything beyond that is the responsibility of coder/aibridge to test.
var interceptionID string
client.EXPECT().RecordInterception(gomock.Any(), gomock.Any()).Times(tc.expectedHits).DoAndReturn(func(ctx context.Context, in *proto.RecordInterceptionRequest) (*proto.RecordInterceptionResponse, error) {
interceptionID = in.GetId()
return &proto.RecordInterceptionResponse{}, nil
})
client.EXPECT().RecordInterceptionEnded(gomock.Any(), gomock.Any()).Times(tc.expectedHits)
// Given: aibridged is started.
srv, err := aibridged.New(t.Context(), pool, func(ctx context.Context) (aibridged.DRPCClient, error) {
return client, nil
}, logger, testTracer)
require.NoError(t, err, "create new aibridged")
t.Cleanup(func() {
_ = srv.Shutdown(testutil.Context(t, testutil.WaitShort))
})
// When: a request is made to aibridged.
ctx := testutil.Context(t, testutil.WaitShort)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, tc.path, bytes.NewBufferString(`{}`))
require.NoError(t, err, "make request to test server")
req.Header.Add("Authorization", "Bearer key")
req.Header.Add("Accept", "application/json")
// When: aibridged handles the request.
rec := httptest.NewRecorder()
srv.ServeHTTP(rec, req)
// Then: the upstream server will have received a number of hits.
// NOTE: we *expect* the interceptions to fail because [mockAIUpstreamServer] returns a nonsense status code.
// We only need to test that the request was routed, NOT processed.
require.Equal(t, tc.expectedStatus, rec.Code)
assert.EqualValues(t, tc.expectedHits, upstreamSrv.Hits())
if tc.expectedHits > 0 {
_, err = uuid.Parse(interceptionID)
require.NoError(t, err, "parse interception ID")
}
})
}
}
// TestServeHTTP_StripInternalHeaders verifies that internal X-Coder-*
// headers are never forwarded to upstream LLM providers.
func TestServeHTTP_StripInternalHeaders(t *testing.T) {
t.Parallel()
cases := []struct {
name string
header string
value string
}{
{
name: "X-Coder-AI-Governance-Token",
header: agplaibridge.HeaderCoderToken,
value: "coder-token",
},
{
name: "X-Coder-AI-Governance-Request-Id",
header: agplaibridge.HeaderCoderRequestID,
value: uuid.NewString(),
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
mockH := &mockHandler{}
srv, client, pool := newTestServer(t)
conn := &mockDRPCConn{}
client.EXPECT().DRPCConn().AnyTimes().Return(conn)
client.EXPECT().IsAuthorized(gomock.Any(), gomock.Any()).AnyTimes().Return(&proto.IsAuthorizedResponse{OwnerId: uuid.NewString()}, nil)
pool.EXPECT().Acquire(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes().Return(mockH, nil)
httpSrv := httptest.NewServer(srv)
t.Cleanup(httpSrv.Close)
ctx := testutil.Context(t, testutil.WaitShort)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, httpSrv.URL+"/anthropic/v1/messages", nil)
require.NoError(t, err)
// Always set a valid auth token so the request reaches
// the upstream handler.
req.Header.Set("Authorization", "Bearer coder-token")
req.Header.Set(tc.header, tc.value)
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode)
require.NotNil(t, mockH.headersReceived)
// Assert no X-Coder-* headers were forwarded upstream.
for name := range mockH.headersReceived {
require.NotContains(t, name, "X-Coder-",
"internal header %q must not be forwarded to upstream providers", name)
}
})
}
}
@@ -0,0 +1,192 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: github.com/coder/coder/v2/coderd/aibridged (interfaces: DRPCClient)
//
// Generated by this command:
//
// mockgen -destination ./clientmock.go -package aibridgedmock github.com/coder/coder/v2/coderd/aibridged DRPCClient
//
// Package aibridgedmock is a generated GoMock package.
package aibridgedmock
import (
context "context"
reflect "reflect"
proto "github.com/coder/coder/v2/coderd/aibridged/proto"
gomock "go.uber.org/mock/gomock"
drpc "storj.io/drpc"
)
// MockDRPCClient is a mock of DRPCClient interface.
type MockDRPCClient struct {
ctrl *gomock.Controller
recorder *MockDRPCClientMockRecorder
isgomock struct{}
}
// MockDRPCClientMockRecorder is the mock recorder for MockDRPCClient.
type MockDRPCClientMockRecorder struct {
mock *MockDRPCClient
}
// NewMockDRPCClient creates a new mock instance.
func NewMockDRPCClient(ctrl *gomock.Controller) *MockDRPCClient {
mock := &MockDRPCClient{ctrl: ctrl}
mock.recorder = &MockDRPCClientMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockDRPCClient) EXPECT() *MockDRPCClientMockRecorder {
return m.recorder
}
// DRPCConn mocks base method.
func (m *MockDRPCClient) DRPCConn() drpc.Conn {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DRPCConn")
ret0, _ := ret[0].(drpc.Conn)
return ret0
}
// DRPCConn indicates an expected call of DRPCConn.
func (mr *MockDRPCClientMockRecorder) DRPCConn() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DRPCConn", reflect.TypeOf((*MockDRPCClient)(nil).DRPCConn))
}
// GetMCPServerAccessTokensBatch mocks base method.
func (m *MockDRPCClient) GetMCPServerAccessTokensBatch(ctx context.Context, in *proto.GetMCPServerAccessTokensBatchRequest) (*proto.GetMCPServerAccessTokensBatchResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetMCPServerAccessTokensBatch", ctx, in)
ret0, _ := ret[0].(*proto.GetMCPServerAccessTokensBatchResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetMCPServerAccessTokensBatch indicates an expected call of GetMCPServerAccessTokensBatch.
func (mr *MockDRPCClientMockRecorder) GetMCPServerAccessTokensBatch(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerAccessTokensBatch", reflect.TypeOf((*MockDRPCClient)(nil).GetMCPServerAccessTokensBatch), ctx, in)
}
// GetMCPServerConfigs mocks base method.
func (m *MockDRPCClient) GetMCPServerConfigs(ctx context.Context, in *proto.GetMCPServerConfigsRequest) (*proto.GetMCPServerConfigsResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetMCPServerConfigs", ctx, in)
ret0, _ := ret[0].(*proto.GetMCPServerConfigsResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetMCPServerConfigs indicates an expected call of GetMCPServerConfigs.
func (mr *MockDRPCClientMockRecorder) GetMCPServerConfigs(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetMCPServerConfigs", reflect.TypeOf((*MockDRPCClient)(nil).GetMCPServerConfigs), ctx, in)
}
// IsAuthorized mocks base method.
func (m *MockDRPCClient) IsAuthorized(ctx context.Context, in *proto.IsAuthorizedRequest) (*proto.IsAuthorizedResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "IsAuthorized", ctx, in)
ret0, _ := ret[0].(*proto.IsAuthorizedResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// IsAuthorized indicates an expected call of IsAuthorized.
func (mr *MockDRPCClientMockRecorder) IsAuthorized(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "IsAuthorized", reflect.TypeOf((*MockDRPCClient)(nil).IsAuthorized), ctx, in)
}
// RecordInterception mocks base method.
func (m *MockDRPCClient) RecordInterception(ctx context.Context, in *proto.RecordInterceptionRequest) (*proto.RecordInterceptionResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "RecordInterception", ctx, in)
ret0, _ := ret[0].(*proto.RecordInterceptionResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// RecordInterception indicates an expected call of RecordInterception.
func (mr *MockDRPCClientMockRecorder) RecordInterception(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RecordInterception", reflect.TypeOf((*MockDRPCClient)(nil).RecordInterception), ctx, in)
}
// RecordInterceptionEnded mocks base method.
func (m *MockDRPCClient) RecordInterceptionEnded(ctx context.Context, in *proto.RecordInterceptionEndedRequest) (*proto.RecordInterceptionEndedResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "RecordInterceptionEnded", ctx, in)
ret0, _ := ret[0].(*proto.RecordInterceptionEndedResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// RecordInterceptionEnded indicates an expected call of RecordInterceptionEnded.
func (mr *MockDRPCClientMockRecorder) RecordInterceptionEnded(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RecordInterceptionEnded", reflect.TypeOf((*MockDRPCClient)(nil).RecordInterceptionEnded), ctx, in)
}
// RecordModelThought mocks base method.
func (m *MockDRPCClient) RecordModelThought(ctx context.Context, in *proto.RecordModelThoughtRequest) (*proto.RecordModelThoughtResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "RecordModelThought", ctx, in)
ret0, _ := ret[0].(*proto.RecordModelThoughtResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// RecordModelThought indicates an expected call of RecordModelThought.
func (mr *MockDRPCClientMockRecorder) RecordModelThought(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RecordModelThought", reflect.TypeOf((*MockDRPCClient)(nil).RecordModelThought), ctx, in)
}
// RecordPromptUsage mocks base method.
func (m *MockDRPCClient) RecordPromptUsage(ctx context.Context, in *proto.RecordPromptUsageRequest) (*proto.RecordPromptUsageResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "RecordPromptUsage", ctx, in)
ret0, _ := ret[0].(*proto.RecordPromptUsageResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// RecordPromptUsage indicates an expected call of RecordPromptUsage.
func (mr *MockDRPCClientMockRecorder) RecordPromptUsage(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RecordPromptUsage", reflect.TypeOf((*MockDRPCClient)(nil).RecordPromptUsage), ctx, in)
}
// RecordTokenUsage mocks base method.
func (m *MockDRPCClient) RecordTokenUsage(ctx context.Context, in *proto.RecordTokenUsageRequest) (*proto.RecordTokenUsageResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "RecordTokenUsage", ctx, in)
ret0, _ := ret[0].(*proto.RecordTokenUsageResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// RecordTokenUsage indicates an expected call of RecordTokenUsage.
func (mr *MockDRPCClientMockRecorder) RecordTokenUsage(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RecordTokenUsage", reflect.TypeOf((*MockDRPCClient)(nil).RecordTokenUsage), ctx, in)
}
// RecordToolUsage mocks base method.
func (m *MockDRPCClient) RecordToolUsage(ctx context.Context, in *proto.RecordToolUsageRequest) (*proto.RecordToolUsageResponse, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "RecordToolUsage", ctx, in)
ret0, _ := ret[0].(*proto.RecordToolUsageResponse)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// RecordToolUsage indicates an expected call of RecordToolUsage.
func (mr *MockDRPCClientMockRecorder) RecordToolUsage(ctx, in any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RecordToolUsage", reflect.TypeOf((*MockDRPCClient)(nil).RecordToolUsage), ctx, in)
}
+4
View File
@@ -0,0 +1,4 @@
package aibridgedmock
//go:generate go tool mockgen -destination ./clientmock.go -package aibridgedmock github.com/coder/coder/v2/coderd/aibridged DRPCClient
//go:generate go tool mockgen -destination ./poolmock.go -package aibridgedmock github.com/coder/coder/v2/coderd/aibridged Pooler
@@ -0,0 +1,72 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: github.com/coder/coder/v2/coderd/aibridged (interfaces: Pooler)
//
// Generated by this command:
//
// mockgen -destination ./poolmock.go -package aibridgedmock github.com/coder/coder/v2/coderd/aibridged Pooler
//
// Package aibridgedmock is a generated GoMock package.
package aibridgedmock
import (
context "context"
http "net/http"
reflect "reflect"
aibridged "github.com/coder/coder/v2/coderd/aibridged"
gomock "go.uber.org/mock/gomock"
)
// MockPooler is a mock of Pooler interface.
type MockPooler struct {
ctrl *gomock.Controller
recorder *MockPoolerMockRecorder
isgomock struct{}
}
// MockPoolerMockRecorder is the mock recorder for MockPooler.
type MockPoolerMockRecorder struct {
mock *MockPooler
}
// NewMockPooler creates a new mock instance.
func NewMockPooler(ctrl *gomock.Controller) *MockPooler {
mock := &MockPooler{ctrl: ctrl}
mock.recorder = &MockPoolerMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockPooler) EXPECT() *MockPoolerMockRecorder {
return m.recorder
}
// Acquire mocks base method.
func (m *MockPooler) Acquire(ctx context.Context, req aibridged.Request, clientFn aibridged.ClientFunc, mcpBootstrapper aibridged.MCPProxyBuilder) (http.Handler, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Acquire", ctx, req, clientFn, mcpBootstrapper)
ret0, _ := ret[0].(http.Handler)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// Acquire indicates an expected call of Acquire.
func (mr *MockPoolerMockRecorder) Acquire(ctx, req, clientFn, mcpBootstrapper any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Acquire", reflect.TypeOf((*MockPooler)(nil).Acquire), ctx, req, clientFn, mcpBootstrapper)
}
// Shutdown mocks base method.
func (m *MockPooler) Shutdown(ctx context.Context) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Shutdown", ctx)
ret0, _ := ret[0].(error)
return ret0
}
// Shutdown indicates an expected call of Shutdown.
func (mr *MockPoolerMockRecorder) Shutdown(ctx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Shutdown", reflect.TypeOf((*MockPooler)(nil).Shutdown), ctx)
}
+34
View File
@@ -0,0 +1,34 @@
package aibridged
import (
"context"
"storj.io/drpc"
"github.com/coder/coder/v2/coderd/aibridged/proto"
)
type Dialer func(ctx context.Context) (DRPCClient, error)
type ClientFunc func() (DRPCClient, error)
// DRPCClient is the union of various service interfaces the client must support.
type DRPCClient interface {
proto.DRPCRecorderClient
proto.DRPCMCPConfiguratorClient
proto.DRPCAuthorizerClient
}
var _ DRPCClient = &Client{}
type Client struct {
proto.DRPCRecorderClient
proto.DRPCMCPConfiguratorClient
proto.DRPCAuthorizerClient
Conn drpc.Conn
}
func (c *Client) DRPCConn() drpc.Conn {
return c.Conn
}
+132
View File
@@ -0,0 +1,132 @@
package aibridged
import (
"net/http"
"strings"
"github.com/google/uuid"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/aibridge/recorder"
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/aibridged/proto"
)
var _ http.Handler = &Server{}
var (
ErrNoAuthKey = xerrors.New("no authentication key provided")
ErrConnect = xerrors.New("could not connect to coderd")
ErrUnauthorized = xerrors.New("unauthorized")
ErrAcquireRequestHandler = xerrors.New("failed to acquire request handler")
)
// ServeHTTP is the entrypoint for requests which will be intercepted by AI Bridge.
// This function will validate that the given API key may be used to perform the request.
//
// An [aibridge.RequestBridge] instance is acquired from a pool based on the API key's
// owner (referred to as the "initiator"); this instance is responsible for the
// AI Bridge-specific handling of the request.
//
// A [DRPCClient] is provided to the [aibridge.RequestBridge] instance so that data can
// be passed up to a [DRPCServer] for persistence.
func (s *Server) ServeHTTP(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
logger := s.logger.With(
slog.F("method", r.Method),
slog.F("path", r.URL.Path),
)
// Extract and strip proxy request ID for cross-service log
// correlation. Absent for direct requests not routed through
// aibridgeproxyd.
if proxyReqID := r.Header.Get(agplaibridge.HeaderCoderRequestID); proxyReqID != "" {
// Inject into context so downstream loggers include it.
ctx = slog.With(ctx, slog.F("aibridgeproxy_id", proxyReqID))
logger = logger.With(slog.F("aibridgeproxy_id", proxyReqID))
}
r.Header.Del(agplaibridge.HeaderCoderRequestID)
byok := agplaibridge.IsBYOK(r.Header)
authMode := "centralized"
if byok {
authMode = "byok"
}
key := strings.TrimSpace(agplaibridge.ExtractAuthToken(r.Header))
if key == "" {
// Some clients (e.g. Claude) send a HEAD request
// without credentials to check connectivity.
if r.Method == http.MethodHead {
logger.Info(ctx, "unauthenticated HEAD request")
} else {
logger.Warn(ctx, "no auth key provided")
}
http.Error(rw, ErrNoAuthKey.Error(), http.StatusBadRequest)
return
}
// Strip every header that may carry the Coder token so it is
// never forwarded to upstream providers. After stripping, the
// aibridge library can treat the request as a normal LLM API call
// with no Coder-specific information.
if byok {
// In BYOK mode the token is in X-Coder-AI-Governance-Token;
// Authorization and X-Api-Key carry the user's own LLM credentials
// and must be preserved.
r.Header.Del(agplaibridge.HeaderCoderToken)
} else {
// In centralized mode the token may be in Authorization (the
// documented path) or X-Api-Key (legacy clients that set
// ANTHROPIC_API_KEY to their Coder token). Both are
// stripped.
r.Header.Del("Authorization")
r.Header.Del("X-Api-Key")
}
client, err := s.Client()
if err != nil {
logger.Warn(ctx, "failed to connect to coderd", slog.Error(err))
http.Error(rw, ErrConnect.Error(), http.StatusServiceUnavailable)
return
}
resp, err := client.IsAuthorized(ctx, &proto.IsAuthorizedRequest{Key: key})
if err != nil {
logger.Warn(ctx, "key authorization check failed", slog.Error(err), slog.F("auth_mode", authMode))
http.Error(rw, ErrUnauthorized.Error(), http.StatusForbidden)
return
}
// Rewire request context to include actor.
//
// [NOTE]
// The metadata provided here must NOT be sensitive as it could be included
// in requests to upstream services.
r = r.WithContext(aibridge.AsActor(ctx, resp.GetOwnerId(), recorder.Metadata{
"Username": resp.GetUsername(),
}))
id, err := uuid.Parse(resp.GetOwnerId())
if err != nil {
logger.Warn(ctx, "failed to parse user ID", slog.Error(err), slog.F("id", resp.GetOwnerId()))
http.Error(rw, ErrUnauthorized.Error(), http.StatusForbidden)
return
}
handler, err := s.GetRequestHandler(ctx, Request{
SessionKey: key,
APIKeyID: resp.ApiKeyId,
InitiatorID: id,
})
if err != nil {
logger.Warn(ctx, "failed to acquire request handler", slog.Error(err))
http.Error(rw, ErrAcquireRequestHandler.Error(), http.StatusInternalServerError)
return
}
handler.ServeHTTP(rw, r)
}
+197
View File
@@ -0,0 +1,197 @@
package aibridged
import (
"context"
"fmt"
"regexp"
"time"
"go.opentelemetry.io/otel/trace"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/aibridge/mcp"
"github.com/coder/coder/v2/coderd/aibridged/proto"
)
var (
ErrEmptyConfig = xerrors.New("empty config given")
ErrCompileRegex = xerrors.New("compile tool regex")
)
const (
InternalMCPServerID = "coder"
)
// Deprecated: Injected MCP in AI Bridge is deprecated and will be removed in a future release.
type MCPProxyBuilder interface {
// Build creates a [mcp.ServerProxier] for the given request initiator.
// At minimum, the Coder MCP server will be proxied.
// The SessionKey from [Request] is used to authenticate against the Coder MCP server.
//
// NOTE: the [mcp.ServerProxier] instance may be proxying one or more MCP servers.
Build(ctx context.Context, req Request, tracer trace.Tracer) (mcp.ServerProxier, error)
}
var _ MCPProxyBuilder = &MCPProxyFactory{}
// Deprecated: Injected MCP in AI Bridge is deprecated and will be removed in a future release.
type MCPProxyFactory struct {
logger slog.Logger
tracer trace.Tracer
clientFn ClientFunc
}
func NewMCPProxyFactory(logger slog.Logger, tracer trace.Tracer, clientFn ClientFunc) *MCPProxyFactory {
return &MCPProxyFactory{
logger: logger,
tracer: tracer,
clientFn: clientFn,
}
}
func (m *MCPProxyFactory) Build(ctx context.Context, req Request, tracer trace.Tracer) (mcp.ServerProxier, error) {
proxiers, err := m.retrieveMCPServerConfigs(ctx, req)
if err != nil {
return nil, xerrors.Errorf("resolve configs: %w", err)
}
return mcp.NewServerProxyManager(proxiers, tracer), nil
}
func (m *MCPProxyFactory) retrieveMCPServerConfigs(ctx context.Context, req Request) (map[string]mcp.ServerProxier, error) {
client, err := m.clientFn()
if err != nil {
return nil, xerrors.Errorf("acquire client: %w", err)
}
srvCfgCtx, srvCfgCancel := context.WithTimeout(ctx, time.Second*10)
defer srvCfgCancel()
// Fetch MCP server configs.
mcpSrvCfgs, err := client.GetMCPServerConfigs(srvCfgCtx, &proto.GetMCPServerConfigsRequest{
UserId: req.InitiatorID.String(),
})
if err != nil {
return nil, xerrors.Errorf("get MCP server configs: %w", err)
}
proxiers := make(map[string]mcp.ServerProxier, len(mcpSrvCfgs.GetExternalAuthMcpConfigs())+1) // Extra one for Coder MCP server.
if mcpSrvCfgs.GetCoderMcpConfig() != nil {
// Setup the Coder MCP server proxy.
coderMCPProxy, err := m.newStreamableHTTPServerProxy(mcpSrvCfgs.GetCoderMcpConfig(), req.SessionKey) // The session key is used to auth against our internal MCP server.
if err != nil {
m.logger.Warn(ctx, "failed to create MCP server proxy", slog.F("mcp_server_id", mcpSrvCfgs.GetCoderMcpConfig().GetId()), slog.Error(err))
} else {
proxiers[InternalMCPServerID] = coderMCPProxy
}
}
if len(mcpSrvCfgs.GetExternalAuthMcpConfigs()) == 0 {
return proxiers, nil
}
serverIDs := make([]string, 0, len(mcpSrvCfgs.GetExternalAuthMcpConfigs()))
for _, cfg := range mcpSrvCfgs.GetExternalAuthMcpConfigs() {
serverIDs = append(serverIDs, cfg.GetId())
}
accTokCtx, accTokCancel := context.WithTimeout(ctx, time.Second*10)
defer accTokCancel()
// Request a batch of access tokens, one per given server ID.
resp, err := client.GetMCPServerAccessTokensBatch(accTokCtx, &proto.GetMCPServerAccessTokensBatchRequest{
UserId: req.InitiatorID.String(),
McpServerConfigIds: serverIDs,
})
if err != nil {
m.logger.Warn(ctx, "failed to retrieve access token(s)", slog.F("server_ids", serverIDs), slog.Error(err))
}
if resp == nil {
m.logger.Warn(ctx, "nil response given to mcp access tokens call")
return proxiers, nil
}
tokens := resp.GetAccessTokens()
if len(tokens) == 0 {
return proxiers, nil
}
// Iterate over all External Auth configurations which are configured for MCP and attempt to setup
// a [mcp.ServerProxier] for it using the access token retrieved above.
for _, cfg := range mcpSrvCfgs.GetExternalAuthMcpConfigs() {
if err, ok := resp.GetErrors()[cfg.GetId()]; ok {
m.logger.Debug(ctx, "failed to get access token", slog.F("mcp_server_id", cfg.GetId()), slog.F("error", err))
continue
}
token, ok := tokens[cfg.GetId()]
if !ok {
m.logger.Warn(ctx, "no access token found", slog.F("mcp_server_id", cfg.GetId()))
continue
}
proxy, err := m.newStreamableHTTPServerProxy(cfg, token)
if err != nil {
m.logger.Warn(ctx, "failed to create MCP server proxy", slog.F("mcp_server_id", cfg.GetId()), slog.Error(err))
continue
}
proxiers[cfg.Id] = proxy
}
return proxiers, nil
}
// newStreamableHTTPServerProxy creates an MCP server capable of proxying requests using the Streamable HTTP transport.
//
// TODO: support SSE transport.
func (m *MCPProxyFactory) newStreamableHTTPServerProxy(cfg *proto.MCPServerConfig, accessToken string) (mcp.ServerProxier, error) {
if cfg == nil {
return nil, ErrEmptyConfig
}
var (
allowlist, denylist *regexp.Regexp
err error
)
if cfg.GetToolAllowRegex() != "" {
allowlist, err = regexp.Compile(cfg.GetToolAllowRegex())
if err != nil {
return nil, ErrCompileRegex
}
}
if cfg.GetToolDenyRegex() != "" {
denylist, err = regexp.Compile(cfg.GetToolDenyRegex())
if err != nil {
return nil, ErrCompileRegex
}
}
// TODO: future improvement:
//
// The access token provided here may expire at any time, or the connection to the MCP server could be severed.
// Instead of passing through an access token directly, rather provide an interface through which to retrieve
// an access token imperatively. In the event of a tool call failing, we could Ping() the MCP server to establish
// whether the connection is still active. If not, this indicates that the access token is probably expired/revoked.
// (It could also mean the server has a problem, which we should account for.)
// The proxy could then use its interface to retrieve a new access token and re-establish a connection.
// For now though, the short TTL of this cache should mostly mask this problem.
srv, err := mcp.NewStreamableHTTPServerProxy(
cfg.GetId(),
cfg.GetUrl(),
// See https://modelcontextprotocol.io/specification/2025-06-18/basic/authorization#token-requirements.
map[string]string{
"Authorization": fmt.Sprintf("Bearer %s", accessToken),
},
allowlist,
denylist,
m.logger.Named(fmt.Sprintf("mcp-server-proxy-%s", cfg.GetId())),
m.tracer,
)
if err != nil {
return nil, xerrors.Errorf("create streamable HTTP MCP server proxy: %w", err)
}
return srv, nil
}
+62
View File
@@ -0,0 +1,62 @@
package aibridged
import (
"testing"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel"
"github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/testutil"
)
func TestMCPRegex(t *testing.T) {
t.Parallel()
cases := []struct {
name string
allowRegex, denyRegex string
expectedErr error
}{
{
name: "invalid allow regex",
allowRegex: `\`,
expectedErr: ErrCompileRegex,
},
{
name: "invalid deny regex",
denyRegex: `+`,
expectedErr: ErrCompileRegex,
},
{
name: "valid empty",
},
{
name: "valid",
allowRegex: "(allowed|allowed2)",
denyRegex: ".*disallowed.*",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
logger := testutil.Logger(t)
f := NewMCPProxyFactory(logger, otel.Tracer("aibridged_test"), nil)
_, err := f.newStreamableHTTPServerProxy(&proto.MCPServerConfig{
Id: "mock",
Url: "mock/mcp",
ToolAllowRegex: tc.allowRegex,
ToolDenyRegex: tc.denyRegex,
}, "")
if tc.expectedErr == nil {
require.NoError(t, err)
} else {
require.ErrorIs(t, err, tc.expectedErr)
}
})
}
}
+205
View File
@@ -0,0 +1,205 @@
package aibridged
import (
"context"
"net/http"
"sync"
"time"
"github.com/dgraph-io/ristretto/v2"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
"golang.org/x/xerrors"
"tailscale.com/util/singleflight"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/aibridge/mcp"
"github.com/coder/coder/v2/aibridge/tracing"
)
const (
cacheCost = 1 // We can't know the actual size in bytes of the value (it'll change over time).
)
// Pooler describes a pool of [*aibridge.RequestBridge] instances from which instances can be retrieved.
// One [*aibridge.RequestBridge] instance is created per given key.
type Pooler interface {
Acquire(ctx context.Context, req Request, clientFn ClientFunc, mcpBootstrapper MCPProxyBuilder) (http.Handler, error)
Shutdown(ctx context.Context) error
}
type PoolMetrics interface {
Hits() uint64
Misses() uint64
KeysAdded() uint64
KeysEvicted() uint64
}
type PoolOptions struct {
MaxItems int64
TTL time.Duration
}
var DefaultPoolOptions = PoolOptions{MaxItems: 5000, TTL: time.Minute * 15}
var _ Pooler = &CachedBridgePool{}
type CachedBridgePool struct {
cache *ristretto.Cache[string, *aibridge.RequestBridge]
providers []aibridge.Provider
logger slog.Logger
options PoolOptions
singleflight *singleflight.Group[string, *aibridge.RequestBridge]
metrics *aibridge.Metrics
tracer trace.Tracer
shutDownOnce sync.Once
shuttingDownCh chan struct{}
}
func NewCachedBridgePool(options PoolOptions, providers []aibridge.Provider, logger slog.Logger, metrics *aibridge.Metrics, tracer trace.Tracer) (*CachedBridgePool, error) {
cache, err := ristretto.NewCache(&ristretto.Config[string, *aibridge.RequestBridge]{
NumCounters: options.MaxItems * 10, // Docs suggest setting this 10x number of keys.
MaxCost: options.MaxItems * cacheCost, // Up to n instances.
IgnoreInternalCost: true, // Don't try estimate cost using bytes (ristretto does this naïvely anyway, just using the size of the value struct not the REAL memory usage).
BufferItems: 64, // Sticking with recommendation from docs.
Metrics: true, // Collect metrics (only used in tests, for now).
OnEvict: func(item *ristretto.Item[*aibridge.RequestBridge]) {
if item == nil || item.Value == nil {
return
}
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), time.Second*5)
defer shutdownCancel()
// Run the eviction in the background since ristretto blocks sets until a free slot is available.
go func() {
_ = item.Value.Shutdown(shutdownCtx)
}()
},
})
if err != nil {
return nil, xerrors.Errorf("create cache: %w", err)
}
return &CachedBridgePool{
cache: cache,
providers: providers,
options: options,
metrics: metrics,
tracer: tracer,
logger: logger,
singleflight: &singleflight.Group[string, *aibridge.RequestBridge]{},
shuttingDownCh: make(chan struct{}),
}, nil
}
// Acquire retrieves or creates a [*aibridge.RequestBridge] instance per given key.
//
// Each returned [*aibridge.RequestBridge] is safe for concurrent use.
// Each [*aibridge.RequestBridge] is stateful because it has MCP clients which maintain sessions to the configured MCP server.
func (p *CachedBridgePool) Acquire(ctx context.Context, req Request, clientFn ClientFunc, mcpProxyFactory MCPProxyBuilder) (_ http.Handler, outErr error) {
spanAttrs := []attribute.KeyValue{
attribute.String(tracing.InitiatorID, req.InitiatorID.String()),
attribute.String(tracing.APIKeyID, req.APIKeyID),
}
ctx, span := p.tracer.Start(ctx, "CachedBridgePool.Acquire", trace.WithAttributes(spanAttrs...))
defer tracing.EndSpanErr(span, &outErr)
ctx = tracing.WithRequestBridgeAttributesInContext(ctx, spanAttrs)
if err := ctx.Err(); err != nil {
return nil, xerrors.Errorf("acquire: %w", err)
}
select {
case <-p.shuttingDownCh:
return nil, xerrors.New("pool shutting down")
default:
}
// Wait for all buffered writes to be applied, otherwise multiple calls in quick succession
// may visit the slow path unnecessarily.
defer p.cache.Wait()
// Fast path.
cacheKey := req.InitiatorID.String() + "|" + req.APIKeyID
bridge, ok := p.cache.Get(cacheKey)
if ok && bridge != nil {
// TODO: future improvement:
// Once we can detect token expiry against an MCP server, we no longer need to let these instances
// expire after the original TTL; we can extend the TTL on each Acquire() call.
// For now, we need to let the instance expiry to keep the MCP connections fresh.
span.AddEvent("cache_hit")
return bridge, nil
}
span.AddEvent("cache_miss")
recorder := aibridge.NewRecorder(p.logger.Named("recorder"), p.tracer, func() (aibridge.Recorder, error) {
client, err := clientFn()
if err != nil {
return nil, xerrors.Errorf("acquire client: %w", err)
}
return &recorderTranslation{apiKeyID: req.APIKeyID, client: client}, nil
})
// Slow path.
// Creating an *aibridge.RequestBridge may take some time, so gate all subsequent callers behind the initial request and return the resulting value.
// TODO: track startup time since it adds latency to first request (histogram count will also help us see how often this occurs).
instance, err, _ := p.singleflight.Do(req.InitiatorID.String(), func() (*aibridge.RequestBridge, error) {
var (
mcpServers mcp.ServerProxier
err error
)
mcpServers, err = mcpProxyFactory.Build(ctx, req, p.tracer)
if err != nil {
p.logger.Warn(ctx, "failed to create MCP server proxiers", slog.Error(err))
// Don't fail here; MCP server injection can gracefully degrade.
}
if mcpServers != nil {
// This will block while connections are established with upstream MCP server(s), and tools are listed.
if err := mcpServers.Init(ctx); err != nil {
p.logger.Warn(ctx, "failed to initialize MCP server proxier(s)", slog.Error(err))
}
}
bridge, err := aibridge.NewRequestBridge(ctx, p.providers, recorder, mcpServers, p.logger, p.metrics, p.tracer)
if err != nil {
return nil, xerrors.Errorf("create new request bridge: %w", err)
}
p.cache.SetWithTTL(cacheKey, bridge, cacheCost, p.options.TTL)
return bridge, nil
})
return instance, err
}
func (p *CachedBridgePool) CacheMetrics() PoolMetrics {
if p.cache == nil {
return nil
}
return p.cache.Metrics
}
// Shutdown will close the cache which will trigger eviction of all the Bridge entries.
func (p *CachedBridgePool) Shutdown(_ context.Context) error {
p.shutDownOnce.Do(func() {
// Prevent new requests from being served.
close(p.shuttingDownCh)
p.cache.Close()
})
return nil
}
+181
View File
@@ -0,0 +1,181 @@
package aibridged_test
import (
"context"
"testing"
"testing/synctest"
"time"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/trace"
"go.uber.org/mock/gomock"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/aibridge/mcp"
"github.com/coder/coder/v2/aibridge/mcpmock"
"github.com/coder/coder/v2/coderd/aibridged"
mock "github.com/coder/coder/v2/coderd/aibridged/aibridgedmock"
)
// TestPool validates the published behavior of [aibridged.CachedBridgePool].
// It is not meant to be an exhaustive test of the internal cache's functionality,
// since that is already covered by its library.
func TestPool(t *testing.T) {
t.Parallel()
logger := slogtest.Make(t, nil)
ctrl := gomock.NewController(t)
client := mock.NewMockDRPCClient(ctrl)
mcpProxy := mcpmock.NewMockServerProxier(ctrl)
opts := aibridged.PoolOptions{MaxItems: 1, TTL: time.Second}
pool, err := aibridged.NewCachedBridgePool(opts, nil, logger, nil, testTracer)
require.NoError(t, err)
t.Cleanup(func() { pool.Shutdown(context.Background()) })
id, id2, apiKeyID1, apiKeyID2 := uuid.New(), uuid.New(), uuid.New(), uuid.New()
clientFn := func() (aibridged.DRPCClient, error) {
return client, nil
}
// Once a pool instance is initialized, it will try setup its MCP proxier(s).
// This is called exactly once since the instance below is only created once.
mcpProxy.EXPECT().Init(gomock.Any()).Times(1).Return(nil)
// This is part of the lifecycle.
mcpProxy.EXPECT().Shutdown(gomock.Any()).AnyTimes().Return(nil)
// Acquiring a pool instance will create one the first time it sees an
// initiator ID...
inst, err := pool.Acquire(t.Context(), aibridged.Request{
SessionKey: "key",
InitiatorID: id,
APIKeyID: apiKeyID1.String(),
}, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err, "acquire pool instance")
// ...and it will return it when acquired again.
instB, err := pool.Acquire(t.Context(), aibridged.Request{
SessionKey: "key",
InitiatorID: id,
APIKeyID: apiKeyID1.String(),
}, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err, "acquire pool instance")
require.Same(t, inst, instB)
cacheMetrics := pool.CacheMetrics()
require.EqualValues(t, 1, cacheMetrics.KeysAdded())
require.EqualValues(t, 0, cacheMetrics.KeysEvicted())
require.EqualValues(t, 1, cacheMetrics.Hits())
require.EqualValues(t, 1, cacheMetrics.Misses())
// This will get called again because a new instance will be created.
mcpProxy.EXPECT().Init(gomock.Any()).Times(1).Return(nil)
// But that key will be evicted when a new initiator is seen (maxItems=1):
inst2, err := pool.Acquire(t.Context(), aibridged.Request{
SessionKey: "key",
InitiatorID: id2,
APIKeyID: apiKeyID1.String(),
}, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err, "acquire pool instance")
require.NotSame(t, inst, inst2)
cacheMetrics = pool.CacheMetrics()
require.EqualValues(t, 2, cacheMetrics.KeysAdded())
require.EqualValues(t, 1, cacheMetrics.KeysEvicted())
require.EqualValues(t, 1, cacheMetrics.Hits())
require.EqualValues(t, 2, cacheMetrics.Misses())
// This will get called again because a new instance will be created.
mcpProxy.EXPECT().Init(gomock.Any()).Times(1).Return(nil)
// New instance is created for different api key id
inst2B, err := pool.Acquire(t.Context(), aibridged.Request{
SessionKey: "key",
InitiatorID: id2,
APIKeyID: apiKeyID2.String(),
}, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err, "acquire pool instance 2B")
require.NotSame(t, inst2, inst2B)
cacheMetrics = pool.CacheMetrics()
require.EqualValues(t, 3, cacheMetrics.KeysAdded())
require.EqualValues(t, 2, cacheMetrics.KeysEvicted())
require.EqualValues(t, 1, cacheMetrics.Hits())
require.EqualValues(t, 3, cacheMetrics.Misses())
}
func TestPool_Expiry(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
logger := slogtest.Make(t, nil)
ctrl := gomock.NewController(t)
client := mock.NewMockDRPCClient(ctrl)
mcpProxy := mcpmock.NewMockServerProxier(ctrl)
mcpProxy.EXPECT().Init(gomock.Any()).AnyTimes().Return(nil)
mcpProxy.EXPECT().Shutdown(gomock.Any()).AnyTimes().Return(nil)
const ttl = time.Second
opts := aibridged.PoolOptions{MaxItems: 1, TTL: ttl}
pool, err := aibridged.NewCachedBridgePool(opts, nil, logger, nil, testTracer)
require.NoError(t, err)
t.Cleanup(func() { pool.Shutdown(context.Background()) })
req := aibridged.Request{
SessionKey: "key",
InitiatorID: uuid.New(),
APIKeyID: uuid.New().String(),
}
clientFn := func() (aibridged.DRPCClient, error) {
return client, nil
}
ctx := t.Context()
// First acquire is a cache miss.
_, err = pool.Acquire(ctx, req, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err)
// Second acquire is a cache hit.
_, err = pool.Acquire(ctx, req, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err)
metrics := pool.CacheMetrics()
require.EqualValues(t, 1, metrics.Misses())
require.EqualValues(t, 1, metrics.Hits())
// TTL expires
time.Sleep(ttl + time.Millisecond)
// Third acquire is a cache miss because the entry expired.
_, err = pool.Acquire(ctx, req, clientFn, newMockMCPFactory(mcpProxy))
require.NoError(t, err)
metrics = pool.CacheMetrics()
require.EqualValues(t, 2, metrics.Misses())
require.EqualValues(t, 1, metrics.Hits())
// Wait for all eviction goroutines to complete before gomock's ctrl.Finish()
// runs in test cleanup. ristretto's OnEvict callback spawns goroutines that
// need to finish calling mcpProxy.Shutdown() before ctrl.finish clears the
// expectations.
synctest.Wait()
})
}
var _ aibridged.MCPProxyBuilder = &mockMCPFactory{}
type mockMCPFactory struct {
proxy *mcpmock.MockServerProxier
}
func newMockMCPFactory(proxy *mcpmock.MockServerProxier) *mockMCPFactory {
return &mockMCPFactory{proxy: proxy}
}
func (m *mockMCPFactory) Build(ctx context.Context, req aibridged.Request, tracer trace.Tracer) (mcp.ServerProxier, error) {
return m.proxy, nil
}
File diff suppressed because it is too large Load Diff
+144
View File
@@ -0,0 +1,144 @@
syntax = "proto3";
option go_package = "github.com/coder/coder/v2/coderd/aibridged/proto";
package proto;
import "google/protobuf/any.proto";
import "google/protobuf/timestamp.proto";
// Recorder is responsible for persisting AI usage records along with their related interception.
service Recorder {
// RecordInterception creates a new interception record to which all other sub-resources
// (token, prompt, tool uses, model thoughts) will be related.
rpc RecordInterception(RecordInterceptionRequest) returns (RecordInterceptionResponse);
rpc RecordInterceptionEnded(RecordInterceptionEndedRequest) returns (RecordInterceptionEndedResponse);
rpc RecordTokenUsage(RecordTokenUsageRequest) returns (RecordTokenUsageResponse);
rpc RecordPromptUsage(RecordPromptUsageRequest) returns (RecordPromptUsageResponse);
rpc RecordToolUsage(RecordToolUsageRequest) returns (RecordToolUsageResponse);
rpc RecordModelThought(RecordModelThoughtRequest) returns (RecordModelThoughtResponse);
}
// MCPConfigurator is responsible for retrieving any relevant data required for configuring MCP clients
// against remote servers.
service MCPConfigurator {
// GetMCPServerConfigs will retrieve MCP server configurations.
rpc GetMCPServerConfigs(GetMCPServerConfigsRequest) returns (GetMCPServerConfigsResponse);
// GetMCPServerAccessTokensBatch will retrieve an access token for a given list of MCP servers, which may involve
// acquiring, validating, or refreshing tokens synchronously. The server should make every effort to
// parallelise this work.
rpc GetMCPServerAccessTokensBatch(GetMCPServerAccessTokensBatchRequest) returns (GetMCPServerAccessTokensBatchResponse);
}
// Authorizer handles all Coder-related authorization functions.
service Authorizer {
// IsAuthorized validates that a given Coder key is valid and the user is authorized to use AI Bridge.
// TODO: add authorization; currently only key validation takes place.
rpc IsAuthorized(IsAuthorizedRequest) returns (IsAuthorizedResponse);
}
message RecordInterceptionRequest {
string id = 1; // UUID.
string initiator_id = 2; // UUID.
string provider = 3;
string model = 4;
map<string, google.protobuf.Any> metadata = 5;
google.protobuf.Timestamp started_at = 6;
string api_key_id = 7;
string client = 8;
string user_agent = 9;
optional string correlating_tool_call_id = 10;
optional string client_session_id = 11;
string provider_name = 12;
string credential_kind = 13;
string credential_hint = 14;
}
message RecordInterceptionResponse {}
message RecordInterceptionEndedRequest {
string id = 1; // UUID.
google.protobuf.Timestamp ended_at = 2;
}
message RecordInterceptionEndedResponse {}
message RecordTokenUsageRequest {
string interception_id = 1; // UUID.
string msg_id = 2; // ID provided by provider.
int64 input_tokens = 3;
int64 output_tokens = 4;
map<string, google.protobuf.Any> metadata = 5;
google.protobuf.Timestamp created_at = 6;
int64 cache_read_input_tokens = 7;
int64 cache_write_input_tokens = 8;
}
message RecordTokenUsageResponse {}
message RecordPromptUsageRequest {
string interception_id = 1; // UUID.
string msg_id = 2; // ID provided by provider.
string prompt = 3;
map<string, google.protobuf.Any> metadata = 4;
google.protobuf.Timestamp created_at = 5;
}
message RecordPromptUsageResponse {}
message RecordToolUsageRequest {
string interception_id = 1; // UUID.
string msg_id = 2; // ID provided by provider.
optional string server_url = 3; // The URL of the MCP server.
string tool = 4;
string input = 5;
bool injected = 6;
optional string invocation_error = 7; // Only injected tools are invoked.
map<string, google.protobuf.Any> metadata = 8;
google.protobuf.Timestamp created_at = 9;
string tool_call_id = 10; // The ID of the tool call provided by the AI provider.
}
message RecordToolUsageResponse {}
message RecordModelThoughtRequest {
string interception_id = 1; // UUID.
string content = 2;
map<string, google.protobuf.Any> metadata = 3;
google.protobuf.Timestamp created_at = 4;
}
message RecordModelThoughtResponse {}
message GetMCPServerConfigsRequest {
string user_id = 1; // UUID. // Not used yet, will be necessary for later RBAC purposes.
}
message GetMCPServerConfigsResponse {
MCPServerConfig coder_mcp_config = 1;
repeated MCPServerConfig external_auth_mcp_configs = 2;
}
message MCPServerConfig {
string id = 1; // Maps to the ID of the External Auth; this ID is unique.
string url = 2;
string tool_allow_regex = 3;
string tool_deny_regex = 4;
}
message GetMCPServerAccessTokensBatchRequest {
string user_id = 1; // UUID.
repeated string mcp_server_config_ids = 2;
}
// GetMCPServerAccessTokensBatchResponse returns a map for resulting tokens or errors, indexed
// by server ID.
message GetMCPServerAccessTokensBatchResponse{
map<string, string> access_tokens = 1;
map<string, string> errors = 2;
}
message IsAuthorizedRequest {
string key = 1;
}
message IsAuthorizedResponse {
string owner_id = 1;
string api_key_id = 2;
string username = 3;
}
+501
View File
@@ -0,0 +1,501 @@
// Code generated by protoc-gen-go-drpc. DO NOT EDIT.
// protoc-gen-go-drpc version: v0.0.34
// source: coderd/aibridged/proto/aibridged.proto
package proto
import (
context "context"
errors "errors"
protojson "google.golang.org/protobuf/encoding/protojson"
proto "google.golang.org/protobuf/proto"
drpc "storj.io/drpc"
drpcerr "storj.io/drpc/drpcerr"
)
type drpcEncoding_File_coderd_aibridged_proto_aibridged_proto struct{}
func (drpcEncoding_File_coderd_aibridged_proto_aibridged_proto) Marshal(msg drpc.Message) ([]byte, error) {
return proto.Marshal(msg.(proto.Message))
}
func (drpcEncoding_File_coderd_aibridged_proto_aibridged_proto) MarshalAppend(buf []byte, msg drpc.Message) ([]byte, error) {
return proto.MarshalOptions{}.MarshalAppend(buf, msg.(proto.Message))
}
func (drpcEncoding_File_coderd_aibridged_proto_aibridged_proto) Unmarshal(buf []byte, msg drpc.Message) error {
return proto.Unmarshal(buf, msg.(proto.Message))
}
func (drpcEncoding_File_coderd_aibridged_proto_aibridged_proto) JSONMarshal(msg drpc.Message) ([]byte, error) {
return protojson.Marshal(msg.(proto.Message))
}
func (drpcEncoding_File_coderd_aibridged_proto_aibridged_proto) JSONUnmarshal(buf []byte, msg drpc.Message) error {
return protojson.Unmarshal(buf, msg.(proto.Message))
}
type DRPCRecorderClient interface {
DRPCConn() drpc.Conn
RecordInterception(ctx context.Context, in *RecordInterceptionRequest) (*RecordInterceptionResponse, error)
RecordInterceptionEnded(ctx context.Context, in *RecordInterceptionEndedRequest) (*RecordInterceptionEndedResponse, error)
RecordTokenUsage(ctx context.Context, in *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error)
RecordPromptUsage(ctx context.Context, in *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error)
RecordToolUsage(ctx context.Context, in *RecordToolUsageRequest) (*RecordToolUsageResponse, error)
RecordModelThought(ctx context.Context, in *RecordModelThoughtRequest) (*RecordModelThoughtResponse, error)
}
type drpcRecorderClient struct {
cc drpc.Conn
}
func NewDRPCRecorderClient(cc drpc.Conn) DRPCRecorderClient {
return &drpcRecorderClient{cc}
}
func (c *drpcRecorderClient) DRPCConn() drpc.Conn { return c.cc }
func (c *drpcRecorderClient) RecordInterception(ctx context.Context, in *RecordInterceptionRequest) (*RecordInterceptionResponse, error) {
out := new(RecordInterceptionResponse)
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordInterception", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
func (c *drpcRecorderClient) RecordInterceptionEnded(ctx context.Context, in *RecordInterceptionEndedRequest) (*RecordInterceptionEndedResponse, error) {
out := new(RecordInterceptionEndedResponse)
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordInterceptionEnded", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
func (c *drpcRecorderClient) RecordTokenUsage(ctx context.Context, in *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error) {
out := new(RecordTokenUsageResponse)
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordTokenUsage", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
func (c *drpcRecorderClient) RecordPromptUsage(ctx context.Context, in *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error) {
out := new(RecordPromptUsageResponse)
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordPromptUsage", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
func (c *drpcRecorderClient) RecordToolUsage(ctx context.Context, in *RecordToolUsageRequest) (*RecordToolUsageResponse, error) {
out := new(RecordToolUsageResponse)
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordToolUsage", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
func (c *drpcRecorderClient) RecordModelThought(ctx context.Context, in *RecordModelThoughtRequest) (*RecordModelThoughtResponse, error) {
out := new(RecordModelThoughtResponse)
err := c.cc.Invoke(ctx, "/proto.Recorder/RecordModelThought", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
type DRPCRecorderServer interface {
RecordInterception(context.Context, *RecordInterceptionRequest) (*RecordInterceptionResponse, error)
RecordInterceptionEnded(context.Context, *RecordInterceptionEndedRequest) (*RecordInterceptionEndedResponse, error)
RecordTokenUsage(context.Context, *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error)
RecordPromptUsage(context.Context, *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error)
RecordToolUsage(context.Context, *RecordToolUsageRequest) (*RecordToolUsageResponse, error)
RecordModelThought(context.Context, *RecordModelThoughtRequest) (*RecordModelThoughtResponse, error)
}
type DRPCRecorderUnimplementedServer struct{}
func (s *DRPCRecorderUnimplementedServer) RecordInterception(context.Context, *RecordInterceptionRequest) (*RecordInterceptionResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
func (s *DRPCRecorderUnimplementedServer) RecordInterceptionEnded(context.Context, *RecordInterceptionEndedRequest) (*RecordInterceptionEndedResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
func (s *DRPCRecorderUnimplementedServer) RecordTokenUsage(context.Context, *RecordTokenUsageRequest) (*RecordTokenUsageResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
func (s *DRPCRecorderUnimplementedServer) RecordPromptUsage(context.Context, *RecordPromptUsageRequest) (*RecordPromptUsageResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
func (s *DRPCRecorderUnimplementedServer) RecordToolUsage(context.Context, *RecordToolUsageRequest) (*RecordToolUsageResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
func (s *DRPCRecorderUnimplementedServer) RecordModelThought(context.Context, *RecordModelThoughtRequest) (*RecordModelThoughtResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
type DRPCRecorderDescription struct{}
func (DRPCRecorderDescription) NumMethods() int { return 6 }
func (DRPCRecorderDescription) Method(n int) (string, drpc.Encoding, drpc.Receiver, interface{}, bool) {
switch n {
case 0:
return "/proto.Recorder/RecordInterception", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCRecorderServer).
RecordInterception(
ctx,
in1.(*RecordInterceptionRequest),
)
}, DRPCRecorderServer.RecordInterception, true
case 1:
return "/proto.Recorder/RecordInterceptionEnded", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCRecorderServer).
RecordInterceptionEnded(
ctx,
in1.(*RecordInterceptionEndedRequest),
)
}, DRPCRecorderServer.RecordInterceptionEnded, true
case 2:
return "/proto.Recorder/RecordTokenUsage", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCRecorderServer).
RecordTokenUsage(
ctx,
in1.(*RecordTokenUsageRequest),
)
}, DRPCRecorderServer.RecordTokenUsage, true
case 3:
return "/proto.Recorder/RecordPromptUsage", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCRecorderServer).
RecordPromptUsage(
ctx,
in1.(*RecordPromptUsageRequest),
)
}, DRPCRecorderServer.RecordPromptUsage, true
case 4:
return "/proto.Recorder/RecordToolUsage", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCRecorderServer).
RecordToolUsage(
ctx,
in1.(*RecordToolUsageRequest),
)
}, DRPCRecorderServer.RecordToolUsage, true
case 5:
return "/proto.Recorder/RecordModelThought", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCRecorderServer).
RecordModelThought(
ctx,
in1.(*RecordModelThoughtRequest),
)
}, DRPCRecorderServer.RecordModelThought, true
default:
return "", nil, nil, nil, false
}
}
func DRPCRegisterRecorder(mux drpc.Mux, impl DRPCRecorderServer) error {
return mux.Register(impl, DRPCRecorderDescription{})
}
type DRPCRecorder_RecordInterceptionStream interface {
drpc.Stream
SendAndClose(*RecordInterceptionResponse) error
}
type drpcRecorder_RecordInterceptionStream struct {
drpc.Stream
}
func (x *drpcRecorder_RecordInterceptionStream) SendAndClose(m *RecordInterceptionResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCRecorder_RecordInterceptionEndedStream interface {
drpc.Stream
SendAndClose(*RecordInterceptionEndedResponse) error
}
type drpcRecorder_RecordInterceptionEndedStream struct {
drpc.Stream
}
func (x *drpcRecorder_RecordInterceptionEndedStream) SendAndClose(m *RecordInterceptionEndedResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCRecorder_RecordTokenUsageStream interface {
drpc.Stream
SendAndClose(*RecordTokenUsageResponse) error
}
type drpcRecorder_RecordTokenUsageStream struct {
drpc.Stream
}
func (x *drpcRecorder_RecordTokenUsageStream) SendAndClose(m *RecordTokenUsageResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCRecorder_RecordPromptUsageStream interface {
drpc.Stream
SendAndClose(*RecordPromptUsageResponse) error
}
type drpcRecorder_RecordPromptUsageStream struct {
drpc.Stream
}
func (x *drpcRecorder_RecordPromptUsageStream) SendAndClose(m *RecordPromptUsageResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCRecorder_RecordToolUsageStream interface {
drpc.Stream
SendAndClose(*RecordToolUsageResponse) error
}
type drpcRecorder_RecordToolUsageStream struct {
drpc.Stream
}
func (x *drpcRecorder_RecordToolUsageStream) SendAndClose(m *RecordToolUsageResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCRecorder_RecordModelThoughtStream interface {
drpc.Stream
SendAndClose(*RecordModelThoughtResponse) error
}
type drpcRecorder_RecordModelThoughtStream struct {
drpc.Stream
}
func (x *drpcRecorder_RecordModelThoughtStream) SendAndClose(m *RecordModelThoughtResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCMCPConfiguratorClient interface {
DRPCConn() drpc.Conn
GetMCPServerConfigs(ctx context.Context, in *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error)
GetMCPServerAccessTokensBatch(ctx context.Context, in *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error)
}
type drpcMCPConfiguratorClient struct {
cc drpc.Conn
}
func NewDRPCMCPConfiguratorClient(cc drpc.Conn) DRPCMCPConfiguratorClient {
return &drpcMCPConfiguratorClient{cc}
}
func (c *drpcMCPConfiguratorClient) DRPCConn() drpc.Conn { return c.cc }
func (c *drpcMCPConfiguratorClient) GetMCPServerConfigs(ctx context.Context, in *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error) {
out := new(GetMCPServerConfigsResponse)
err := c.cc.Invoke(ctx, "/proto.MCPConfigurator/GetMCPServerConfigs", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
func (c *drpcMCPConfiguratorClient) GetMCPServerAccessTokensBatch(ctx context.Context, in *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error) {
out := new(GetMCPServerAccessTokensBatchResponse)
err := c.cc.Invoke(ctx, "/proto.MCPConfigurator/GetMCPServerAccessTokensBatch", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
type DRPCMCPConfiguratorServer interface {
GetMCPServerConfigs(context.Context, *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error)
GetMCPServerAccessTokensBatch(context.Context, *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error)
}
type DRPCMCPConfiguratorUnimplementedServer struct{}
func (s *DRPCMCPConfiguratorUnimplementedServer) GetMCPServerConfigs(context.Context, *GetMCPServerConfigsRequest) (*GetMCPServerConfigsResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
func (s *DRPCMCPConfiguratorUnimplementedServer) GetMCPServerAccessTokensBatch(context.Context, *GetMCPServerAccessTokensBatchRequest) (*GetMCPServerAccessTokensBatchResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
type DRPCMCPConfiguratorDescription struct{}
func (DRPCMCPConfiguratorDescription) NumMethods() int { return 2 }
func (DRPCMCPConfiguratorDescription) Method(n int) (string, drpc.Encoding, drpc.Receiver, interface{}, bool) {
switch n {
case 0:
return "/proto.MCPConfigurator/GetMCPServerConfigs", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCMCPConfiguratorServer).
GetMCPServerConfigs(
ctx,
in1.(*GetMCPServerConfigsRequest),
)
}, DRPCMCPConfiguratorServer.GetMCPServerConfigs, true
case 1:
return "/proto.MCPConfigurator/GetMCPServerAccessTokensBatch", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCMCPConfiguratorServer).
GetMCPServerAccessTokensBatch(
ctx,
in1.(*GetMCPServerAccessTokensBatchRequest),
)
}, DRPCMCPConfiguratorServer.GetMCPServerAccessTokensBatch, true
default:
return "", nil, nil, nil, false
}
}
func DRPCRegisterMCPConfigurator(mux drpc.Mux, impl DRPCMCPConfiguratorServer) error {
return mux.Register(impl, DRPCMCPConfiguratorDescription{})
}
type DRPCMCPConfigurator_GetMCPServerConfigsStream interface {
drpc.Stream
SendAndClose(*GetMCPServerConfigsResponse) error
}
type drpcMCPConfigurator_GetMCPServerConfigsStream struct {
drpc.Stream
}
func (x *drpcMCPConfigurator_GetMCPServerConfigsStream) SendAndClose(m *GetMCPServerConfigsResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCMCPConfigurator_GetMCPServerAccessTokensBatchStream interface {
drpc.Stream
SendAndClose(*GetMCPServerAccessTokensBatchResponse) error
}
type drpcMCPConfigurator_GetMCPServerAccessTokensBatchStream struct {
drpc.Stream
}
func (x *drpcMCPConfigurator_GetMCPServerAccessTokensBatchStream) SendAndClose(m *GetMCPServerAccessTokensBatchResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
type DRPCAuthorizerClient interface {
DRPCConn() drpc.Conn
IsAuthorized(ctx context.Context, in *IsAuthorizedRequest) (*IsAuthorizedResponse, error)
}
type drpcAuthorizerClient struct {
cc drpc.Conn
}
func NewDRPCAuthorizerClient(cc drpc.Conn) DRPCAuthorizerClient {
return &drpcAuthorizerClient{cc}
}
func (c *drpcAuthorizerClient) DRPCConn() drpc.Conn { return c.cc }
func (c *drpcAuthorizerClient) IsAuthorized(ctx context.Context, in *IsAuthorizedRequest) (*IsAuthorizedResponse, error) {
out := new(IsAuthorizedResponse)
err := c.cc.Invoke(ctx, "/proto.Authorizer/IsAuthorized", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}, in, out)
if err != nil {
return nil, err
}
return out, nil
}
type DRPCAuthorizerServer interface {
IsAuthorized(context.Context, *IsAuthorizedRequest) (*IsAuthorizedResponse, error)
}
type DRPCAuthorizerUnimplementedServer struct{}
func (s *DRPCAuthorizerUnimplementedServer) IsAuthorized(context.Context, *IsAuthorizedRequest) (*IsAuthorizedResponse, error) {
return nil, drpcerr.WithCode(errors.New("Unimplemented"), drpcerr.Unimplemented)
}
type DRPCAuthorizerDescription struct{}
func (DRPCAuthorizerDescription) NumMethods() int { return 1 }
func (DRPCAuthorizerDescription) Method(n int) (string, drpc.Encoding, drpc.Receiver, interface{}, bool) {
switch n {
case 0:
return "/proto.Authorizer/IsAuthorized", drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{},
func(srv interface{}, ctx context.Context, in1, in2 interface{}) (drpc.Message, error) {
return srv.(DRPCAuthorizerServer).
IsAuthorized(
ctx,
in1.(*IsAuthorizedRequest),
)
}, DRPCAuthorizerServer.IsAuthorized, true
default:
return "", nil, nil, nil, false
}
}
func DRPCRegisterAuthorizer(mux drpc.Mux, impl DRPCAuthorizerServer) error {
return mux.Register(impl, DRPCAuthorizerDescription{})
}
type DRPCAuthorizer_IsAuthorizedStream interface {
drpc.Stream
SendAndClose(*IsAuthorizedResponse) error
}
type drpcAuthorizer_IsAuthorizedStream struct {
drpc.Stream
}
func (x *drpcAuthorizer_IsAuthorizedStream) SendAndClose(m *IsAuthorizedResponse) error {
if err := x.MsgSend(m, drpcEncoding_File_coderd_aibridged_proto_aibridged_proto{}); err != nil {
return err
}
return x.CloseSend()
}
+9
View File
@@ -0,0 +1,9 @@
package aibridged
import "github.com/google/uuid"
type Request struct {
SessionKey string
APIKeyID string
InitiatorID uuid.UUID
}
+9
View File
@@ -0,0 +1,9 @@
package aibridged
import "github.com/coder/coder/v2/coderd/aibridged/proto"
type DRPCServer interface {
proto.DRPCRecorderServer
proto.DRPCMCPConfiguratorServer
proto.DRPCAuthorizerServer
}
+158
View File
@@ -0,0 +1,158 @@
package aibridged
import (
"context"
"encoding/json"
"fmt"
"golang.org/x/xerrors"
"google.golang.org/protobuf/types/known/anypb"
"google.golang.org/protobuf/types/known/structpb"
"google.golang.org/protobuf/types/known/timestamppb"
"github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/coderd/util/ptr"
)
var _ aibridge.Recorder = &recorderTranslation{}
// recorderTranslation satisfies the aibridge.Recorder interface and translates calls into dRPC calls to aibridgedserver.
type recorderTranslation struct {
apiKeyID string
client proto.DRPCRecorderClient
}
func (t *recorderTranslation) RecordInterception(ctx context.Context, req *aibridge.InterceptionRecord) error {
_, err := t.client.RecordInterception(ctx, &proto.RecordInterceptionRequest{
Id: req.ID,
ApiKeyId: t.apiKeyID,
InitiatorId: req.InitiatorID,
Provider: req.Provider,
ProviderName: req.ProviderName,
Model: req.Model,
UserAgent: req.UserAgent,
Client: req.Client,
ClientSessionId: req.ClientSessionID,
Metadata: marshalForProto(req.Metadata),
StartedAt: timestamppb.New(req.StartedAt),
CorrelatingToolCallId: req.CorrelatingToolCallID,
CredentialKind: req.CredentialKind,
CredentialHint: req.CredentialHint,
})
return err
}
func (t *recorderTranslation) RecordInterceptionEnded(ctx context.Context, req *aibridge.InterceptionRecordEnded) error {
_, err := t.client.RecordInterceptionEnded(ctx, &proto.RecordInterceptionEndedRequest{
Id: req.ID,
EndedAt: timestamppb.New(req.EndedAt),
})
return err
}
func (t *recorderTranslation) RecordPromptUsage(ctx context.Context, req *aibridge.PromptUsageRecord) error {
_, err := t.client.RecordPromptUsage(ctx, &proto.RecordPromptUsageRequest{
InterceptionId: req.InterceptionID,
MsgId: req.MsgID,
Prompt: req.Prompt,
Metadata: marshalForProto(req.Metadata),
CreatedAt: timestamppb.New(req.CreatedAt),
})
return err
}
func (t *recorderTranslation) RecordTokenUsage(ctx context.Context, req *aibridge.TokenUsageRecord) error {
merged := req.Metadata
if merged == nil {
merged = aibridge.Metadata{}
}
// Merge remaining extra token types into metadata.
for k, v := range req.ExtraTokenTypes {
merged[k] = v
}
_, err := t.client.RecordTokenUsage(ctx, &proto.RecordTokenUsageRequest{
InterceptionId: req.InterceptionID,
MsgId: req.MsgID,
InputTokens: req.Input,
OutputTokens: req.Output,
CacheReadInputTokens: req.CacheReadInputTokens,
CacheWriteInputTokens: req.CacheWriteInputTokens,
Metadata: marshalForProto(merged),
CreatedAt: timestamppb.New(req.CreatedAt),
})
return err
}
func (t *recorderTranslation) RecordToolUsage(ctx context.Context, req *aibridge.ToolUsageRecord) error {
serialized, err := json.Marshal(req.Args)
if err != nil {
return xerrors.Errorf("serialize tool %q args: %w", req.Tool, err)
}
var invErr *string
if req.InvocationError != nil {
invErr = ptr.Ref(req.InvocationError.Error())
}
_, err = t.client.RecordToolUsage(ctx, &proto.RecordToolUsageRequest{
InterceptionId: req.InterceptionID,
MsgId: req.MsgID,
ToolCallId: req.ToolCallID,
ServerUrl: req.ServerURL,
Tool: req.Tool,
Input: string(serialized),
Injected: req.Injected,
InvocationError: invErr,
Metadata: marshalForProto(req.Metadata),
CreatedAt: timestamppb.New(req.CreatedAt),
})
return err
}
func (t *recorderTranslation) RecordModelThought(ctx context.Context, req *aibridge.ModelThoughtRecord) error {
_, err := t.client.RecordModelThought(ctx, &proto.RecordModelThoughtRequest{
InterceptionId: req.InterceptionID,
Content: req.Content,
Metadata: marshalForProto(req.Metadata),
CreatedAt: timestamppb.New(req.CreatedAt),
})
return err
}
// marshalForProto will attempt to convert from aibridge.Metadata into a proto-friendly map[string]*anypb.Any.
// If any marshaling fails, rather return a map with the error details since we don't want to fail Record* funcs if metadata can't encode,
// since it's, well, metadata.
func marshalForProto(in aibridge.Metadata) map[string]*anypb.Any {
out := make(map[string]*anypb.Any, len(in))
if len(in) == 0 {
return out
}
// Instead of returning error, just encode error into metadata.
encodeErr := func(err error) map[string]*anypb.Any {
errVal, _ := anypb.New(structpb.NewStringValue(err.Error()))
mdVal, _ := anypb.New(structpb.NewStringValue(fmt.Sprintf("%+v", in)))
return map[string]*anypb.Any{
"error": errVal,
"metadata": mdVal,
}
}
for k, v := range in {
sv, err := structpb.NewValue(v)
if err != nil {
return encodeErr(err)
}
av, err := anypb.New(sv)
if err != nil {
return encodeErr(err)
}
out[k] = av
}
return out
}
+27
View File
@@ -0,0 +1,27 @@
package aibridged_test
import (
"net/http"
"sync/atomic"
"go.opentelemetry.io/otel"
)
var testTracer = otel.Tracer("aibridged_test")
var _ http.Handler = &mockAIUpstreamServer{}
type mockAIUpstreamServer struct {
hitCounter atomic.Int32
}
func (m *mockAIUpstreamServer) ServeHTTP(rw http.ResponseWriter, _ *http.Request) {
m.hitCounter.Add(1)
rw.WriteHeader(http.StatusTeapot)
_, _ = rw.Write([]byte(`i am a teapot`))
}
func (m *mockAIUpstreamServer) Hits() int32 {
return m.hitCounter.Load()
}
+659
View File
@@ -0,0 +1,659 @@
package aibridgedserver
import (
"context"
"database/sql"
"encoding/json"
"net/url"
"slices"
"strings"
"sync"
"github.com/google/uuid"
"github.com/hashicorp/go-multierror"
"golang.org/x/xerrors"
"google.golang.org/protobuf/types/known/anypb"
"google.golang.org/protobuf/types/known/structpb"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/aibridged"
"github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/coderd/aiseats"
"github.com/coder/coder/v2/coderd/apikey"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
"github.com/coder/coder/v2/coderd/database/dbtime"
"github.com/coder/coder/v2/coderd/externalauth"
"github.com/coder/coder/v2/coderd/httpmw"
codermcp "github.com/coder/coder/v2/coderd/mcp"
"github.com/coder/coder/v2/codersdk"
)
var (
ErrExpiredOrInvalidOAuthToken = xerrors.New("expired or invalid OAuth2 token")
ErrNoMCPConfigFound = xerrors.New("no MCP config found")
// These errors are returned by IsAuthorized. Since they're just returned as
// a generic dRPC error, it's difficult to tell them apart without string
// matching.
// TODO: return these errors to the client in a more structured/comparable
// way.
ErrInvalidKey = xerrors.New("invalid key")
ErrUnknownKey = xerrors.New("unknown key")
ErrExpired = xerrors.New("expired")
ErrUnknownUser = xerrors.New("unknown user")
ErrDeletedUser = xerrors.New("deleted user")
ErrSystemUser = xerrors.New("system user")
ErrNoExternalAuthLinkFound = xerrors.New("no external auth link found")
)
const (
InterceptionLogMarker = "interception log"
MetadataUserAgentKey = "request_user_agent"
)
var _ aibridged.DRPCServer = &Server{}
type store interface {
// Recorder-related queries.
InsertAIBridgeInterception(ctx context.Context, arg database.InsertAIBridgeInterceptionParams) (database.AIBridgeInterception, error)
InsertAIBridgeTokenUsage(ctx context.Context, arg database.InsertAIBridgeTokenUsageParams) (database.AIBridgeTokenUsage, error)
InsertAIBridgeUserPrompt(ctx context.Context, arg database.InsertAIBridgeUserPromptParams) (database.AIBridgeUserPrompt, error)
InsertAIBridgeToolUsage(ctx context.Context, arg database.InsertAIBridgeToolUsageParams) (database.AIBridgeToolUsage, error)
InsertAIBridgeModelThought(ctx context.Context, arg database.InsertAIBridgeModelThoughtParams) (database.AIBridgeModelThought, error)
UpdateAIBridgeInterceptionEnded(ctx context.Context, intcID database.UpdateAIBridgeInterceptionEndedParams) (database.AIBridgeInterception, error)
GetAIBridgeInterceptionLineageByToolCallID(ctx context.Context, toolCallID string) (database.GetAIBridgeInterceptionLineageByToolCallIDRow, error)
// MCPConfigurator-related queries.
GetExternalAuthLinksByUserID(ctx context.Context, userID uuid.UUID) ([]database.ExternalAuthLink, error)
// Authorizer-related queries.
GetAPIKeyByID(ctx context.Context, id string) (database.APIKey, error)
GetUserByID(ctx context.Context, id uuid.UUID) (database.User, error)
}
type Server struct {
// lifecycleCtx must be tied to the API server's lifecycle
// as when the API server shuts down, we want to cancel any
// long-running operations.
lifecycleCtx context.Context
store store
logger slog.Logger
externalAuthConfigs map[string]*externalauth.Config
coderMCPConfig *proto.MCPServerConfig // may be nil if not available
structuredLogging bool
aiSeatTracker aiseats.SeatTracker
}
func NewServer(lifecycleCtx context.Context, store store, logger slog.Logger, accessURL string,
bridgeCfg codersdk.AIBridgeConfig, externalAuthConfigs []*externalauth.Config, experiments codersdk.Experiments,
aiSeatTracker aiseats.SeatTracker,
) (*Server, error) {
eac := make(map[string]*externalauth.Config, len(externalAuthConfigs))
for _, cfg := range externalAuthConfigs {
// Only External Auth configs which are configured with an MCP URL are relevant to aibridged.
if cfg.MCPURL == "" {
continue
}
eac[cfg.ID] = cfg
}
srv := &Server{
lifecycleCtx: lifecycleCtx,
store: store,
logger: logger,
externalAuthConfigs: eac,
structuredLogging: bridgeCfg.StructuredLogging.Value(),
aiSeatTracker: aiSeatTracker,
}
if bridgeCfg.InjectCoderMCPTools {
logger.Warn(lifecycleCtx, "inject MCP tools option is deprecated and will be removed in a future release")
coderMCPConfig, err := getCoderMCPServerConfig(experiments, accessURL)
if err != nil {
logger.Warn(lifecycleCtx, "failed to retrieve coder MCP server config, Coder MCP will not be available", slog.Error(err))
}
srv.coderMCPConfig = coderMCPConfig
}
return srv, nil
}
func (s *Server) RecordInterception(ctx context.Context, in *proto.RecordInterceptionRequest) (*proto.RecordInterceptionResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
intcID, err := uuid.Parse(in.GetId())
if err != nil {
return nil, xerrors.Errorf("invalid interception ID %q: %w", in.GetId(), err)
}
initID, err := uuid.Parse(in.GetInitiatorId())
if err != nil {
return nil, xerrors.Errorf("invalid initiator ID %q: %w", in.GetInitiatorId(), err)
}
if in.ApiKeyId == "" {
return nil, xerrors.Errorf("empty API key ID")
}
metadata := metadataToMap(in.GetMetadata())
if in.UserAgent != "" {
if _, ok := metadata[MetadataUserAgentKey]; ok {
s.logger.Warn(ctx, "interception metadata contains user agent key, will be overwritten")
}
metadata[MetadataUserAgentKey] = in.UserAgent
}
// Look up the interception lineage using the correlating tool call ID.
parentID, rootID := s.findInterceptionLineage(ctx, in.GetCorrelatingToolCallId())
if s.structuredLogging {
s.logger.Info(ctx, InterceptionLogMarker,
slog.F("record_type", "interception_start"),
slog.F("interception_id", intcID.String()),
slog.F("initiator_id", initID.String()),
slog.F("api_key_id", in.ApiKeyId),
slog.F("provider", in.Provider),
slog.F("model", in.Model),
slog.F("client", in.Client),
slog.F("client_session_id", in.GetClientSessionId()),
slog.F("started_at", in.StartedAt.AsTime()),
slog.F("metadata", metadata),
slog.F("correlating_tool_call_id", in.GetCorrelatingToolCallId()),
slog.F("thread_parent_id", parentID),
slog.F("thread_root_id", rootID),
)
}
out, err := json.Marshal(metadata)
if err != nil {
s.logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err))
}
providerName := strings.TrimSpace(in.ProviderName)
if providerName == "" {
providerName = in.Provider
}
_, err = s.store.InsertAIBridgeInterception(ctx, database.InsertAIBridgeInterceptionParams{
ID: intcID,
APIKeyID: sql.NullString{String: in.ApiKeyId, Valid: true},
Client: sql.NullString{String: in.Client, Valid: in.Client != ""},
ClientSessionID: sql.NullString{String: in.GetClientSessionId(), Valid: in.GetClientSessionId() != ""},
InitiatorID: initID,
Provider: in.Provider,
ProviderName: providerName,
Model: in.Model,
Metadata: out,
StartedAt: in.StartedAt.AsTime(),
ThreadParentInterceptionID: uuid.NullUUID{UUID: parentID, Valid: parentID != uuid.Nil},
ThreadRootInterceptionID: uuid.NullUUID{UUID: rootID, Valid: rootID != uuid.Nil},
CredentialKind: credentialKindOrDefault(in.CredentialKind),
CredentialHint: in.CredentialHint,
})
if err != nil {
return nil, xerrors.Errorf("start interception: %w", err)
}
reason := aiseats.ReasonAIBridge("provider=" + in.Provider + ", model=" + in.Model)
s.aiSeatTracker.RecordUsage(ctx, initID, reason)
return &proto.RecordInterceptionResponse{}, nil
}
func (s *Server) RecordInterceptionEnded(ctx context.Context, in *proto.RecordInterceptionEndedRequest) (*proto.RecordInterceptionEndedResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
intcID, err := uuid.Parse(in.GetId())
if err != nil {
return nil, xerrors.Errorf("invalid interception ID %q: %w", in.GetId(), err)
}
if s.structuredLogging {
s.logger.Info(ctx, InterceptionLogMarker,
slog.F("record_type", "interception_end"),
slog.F("interception_id", intcID.String()),
slog.F("ended_at", in.EndedAt.AsTime()),
)
}
_, err = s.store.UpdateAIBridgeInterceptionEnded(ctx, database.UpdateAIBridgeInterceptionEndedParams{
ID: intcID,
EndedAt: in.EndedAt.AsTime(),
})
if err != nil {
return nil, xerrors.Errorf("end interception: %w", err)
}
return &proto.RecordInterceptionEndedResponse{}, nil
}
func (s *Server) RecordTokenUsage(ctx context.Context, in *proto.RecordTokenUsageRequest) (*proto.RecordTokenUsageResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
intcID, err := uuid.Parse(in.GetInterceptionId())
if err != nil {
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
}
metadata := metadataToMap(in.GetMetadata())
if s.structuredLogging {
s.logger.Info(ctx, InterceptionLogMarker,
slog.F("record_type", "token_usage"),
slog.F("interception_id", intcID.String()),
slog.F("msg_id", in.GetMsgId()),
slog.F("input_tokens", in.GetInputTokens()),
slog.F("output_tokens", in.GetOutputTokens()),
slog.F("cache_read_input_tokens", in.GetCacheReadInputTokens()),
slog.F("cache_write_input_tokens", in.GetCacheWriteInputTokens()),
slog.F("created_at", in.GetCreatedAt().AsTime()),
slog.F("metadata", metadata),
)
}
out, err := json.Marshal(metadata)
if err != nil {
s.logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err))
}
_, err = s.store.InsertAIBridgeTokenUsage(ctx, database.InsertAIBridgeTokenUsageParams{
ID: uuid.New(),
InterceptionID: intcID,
ProviderResponseID: in.GetMsgId(),
InputTokens: in.GetInputTokens(),
OutputTokens: in.GetOutputTokens(),
CacheReadInputTokens: in.GetCacheReadInputTokens(),
CacheWriteInputTokens: in.GetCacheWriteInputTokens(),
Metadata: out,
CreatedAt: in.GetCreatedAt().AsTime(),
})
if err != nil {
return nil, xerrors.Errorf("insert token usage: %w", err)
}
return &proto.RecordTokenUsageResponse{}, nil
}
func (s *Server) RecordPromptUsage(ctx context.Context, in *proto.RecordPromptUsageRequest) (*proto.RecordPromptUsageResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
intcID, err := uuid.Parse(in.GetInterceptionId())
if err != nil {
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
}
metadata := metadataToMap(in.GetMetadata())
if s.structuredLogging {
s.logger.Info(ctx, InterceptionLogMarker,
slog.F("record_type", "prompt_usage"),
slog.F("interception_id", intcID.String()),
slog.F("msg_id", in.GetMsgId()),
slog.F("prompt", in.GetPrompt()),
slog.F("created_at", in.GetCreatedAt().AsTime()),
slog.F("metadata", metadata),
)
}
out, err := json.Marshal(metadata)
if err != nil {
s.logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err))
}
_, err = s.store.InsertAIBridgeUserPrompt(ctx, database.InsertAIBridgeUserPromptParams{
ID: uuid.New(),
InterceptionID: intcID,
ProviderResponseID: in.GetMsgId(),
Prompt: in.GetPrompt(),
Metadata: out,
CreatedAt: in.GetCreatedAt().AsTime(),
})
if err != nil {
return nil, xerrors.Errorf("insert user prompt: %w", err)
}
return &proto.RecordPromptUsageResponse{}, nil
}
func (s *Server) RecordToolUsage(ctx context.Context, in *proto.RecordToolUsageRequest) (*proto.RecordToolUsageResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
intcID, err := uuid.Parse(in.GetInterceptionId())
if err != nil {
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
}
metadata := metadataToMap(in.GetMetadata())
if s.structuredLogging {
s.logger.Info(ctx, InterceptionLogMarker,
slog.F("record_type", "tool_usage"),
slog.F("interception_id", intcID.String()),
slog.F("msg_id", in.GetMsgId()),
slog.F("tool_call_id", in.GetToolCallId()),
slog.F("tool", in.GetTool()),
slog.F("input", in.GetInput()),
slog.F("server_url", in.GetServerUrl()),
slog.F("injected", in.GetInjected()),
slog.F("invocation_error", in.GetInvocationError()),
slog.F("created_at", in.GetCreatedAt().AsTime()),
slog.F("metadata", metadata),
)
}
out, err := json.Marshal(metadata)
if err != nil {
s.logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err))
}
_, err = s.store.InsertAIBridgeToolUsage(ctx, database.InsertAIBridgeToolUsageParams{
ID: uuid.New(),
InterceptionID: intcID,
ProviderResponseID: in.GetMsgId(),
ProviderToolCallID: sql.NullString{String: in.GetToolCallId(), Valid: in.GetToolCallId() != ""},
ServerUrl: sql.NullString{String: in.GetServerUrl(), Valid: in.ServerUrl != nil},
Tool: in.GetTool(),
Input: in.GetInput(),
Injected: in.GetInjected(),
InvocationError: sql.NullString{String: in.GetInvocationError(), Valid: in.InvocationError != nil},
Metadata: out,
CreatedAt: in.GetCreatedAt().AsTime(),
})
if err != nil {
return nil, xerrors.Errorf("insert tool usage: %w", err)
}
return &proto.RecordToolUsageResponse{}, nil
}
func (s *Server) RecordModelThought(ctx context.Context, in *proto.RecordModelThoughtRequest) (*proto.RecordModelThoughtResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
intcID, err := uuid.Parse(in.GetInterceptionId())
if err != nil {
return nil, xerrors.Errorf("failed to parse interception_id %q: %w", in.GetInterceptionId(), err)
}
metadata := metadataToMap(in.GetMetadata())
if s.structuredLogging {
s.logger.Info(ctx, InterceptionLogMarker,
slog.F("record_type", "model_thought"),
slog.F("interception_id", intcID.String()),
slog.F("content", in.GetContent()),
slog.F("created_at", in.GetCreatedAt().AsTime()),
slog.F("metadata", metadata),
)
}
out, err := json.Marshal(metadata)
if err != nil {
s.logger.Warn(ctx, "failed to marshal aibridge metadata from proto to JSON", slog.F("metadata", in), slog.Error(err))
}
_, err = s.store.InsertAIBridgeModelThought(ctx, database.InsertAIBridgeModelThoughtParams{
InterceptionID: intcID,
Content: in.GetContent(),
Metadata: out,
CreatedAt: in.GetCreatedAt().AsTime(),
})
if err != nil {
return nil, xerrors.Errorf("insert model thought: %w", err)
}
return &proto.RecordModelThoughtResponse{}, nil
}
// findInterceptionLineage looks up the parent interception and the root
// of the thread by finding which interception recorded a tool usage with
// the given tool call ID. Returns (parentID, rootID); both will be
// uuid.Nil if no match is found or the tool call ID is empty.
func (s *Server) findInterceptionLineage(ctx context.Context, toolCallID string) (parent uuid.UUID, root uuid.UUID) {
if toolCallID == "" {
return uuid.Nil, uuid.Nil
}
lineage, err := s.store.GetAIBridgeInterceptionLineageByToolCallID(ctx, toolCallID)
if err != nil {
s.logger.Warn(ctx, "failed to retrieve interception lineage",
slog.Error(err), slog.F("tool_call_id", toolCallID))
return uuid.Nil, uuid.Nil
}
return lineage.ThreadParentID, lineage.ThreadRootID
}
func (s *Server) GetMCPServerConfigs(_ context.Context, _ *proto.GetMCPServerConfigsRequest) (*proto.GetMCPServerConfigsResponse, error) {
cfgs := make([]*proto.MCPServerConfig, 0, len(s.externalAuthConfigs))
for _, eac := range s.externalAuthConfigs {
var allowlist, denylist string
if eac.MCPToolAllowRegex != nil {
allowlist = eac.MCPToolAllowRegex.String()
}
if eac.MCPToolDenyRegex != nil {
denylist = eac.MCPToolDenyRegex.String()
}
cfgs = append(cfgs, &proto.MCPServerConfig{
Id: eac.ID,
Url: eac.MCPURL,
ToolAllowRegex: allowlist,
ToolDenyRegex: denylist,
})
}
return &proto.GetMCPServerConfigsResponse{
CoderMcpConfig: s.coderMCPConfig, // it's fine if this is nil
ExternalAuthMcpConfigs: cfgs,
}, nil
}
func (s *Server) GetMCPServerAccessTokensBatch(ctx context.Context, in *proto.GetMCPServerAccessTokensBatchRequest) (*proto.GetMCPServerAccessTokensBatchResponse, error) {
if len(in.GetMcpServerConfigIds()) == 0 {
return &proto.GetMCPServerAccessTokensBatchResponse{}, nil
}
userID, err := uuid.Parse(in.GetUserId())
if err != nil {
return nil, xerrors.Errorf("parse user_id: %w", err)
}
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
links, err := s.store.GetExternalAuthLinksByUserID(ctx, userID)
if err != nil {
return nil, xerrors.Errorf("fetch external auth links: %w", err)
}
if len(links) == 0 {
return &proto.GetMCPServerAccessTokensBatchResponse{}, nil
}
// Ensure unique to prevent unnecessary effort.
ids := in.GetMcpServerConfigIds()
slices.Sort(ids)
ids = slices.Compact(ids)
var (
wg sync.WaitGroup
errs error
mu sync.Mutex
tokens = make(map[string]string, len(ids))
tokenErrs = make(map[string]string)
)
externalAuthLoop:
for _, id := range ids {
eac, ok := s.externalAuthConfigs[id]
if !ok {
mu.Lock()
s.logger.Warn(ctx, "no MCP server config found by given ID", slog.F("id", id))
tokenErrs[id] = ErrNoMCPConfigFound.Error()
mu.Unlock()
continue
}
for _, link := range links {
if link.ProviderID != eac.ID {
continue
}
// Validate all configured External Auth links concurrently.
wg.Add(1)
go func() {
defer wg.Done()
// TODO: timeout.
valid, _, validateErr := eac.ValidateToken(ctx, link.OAuthToken())
mu.Lock()
defer mu.Unlock()
if !valid {
// TODO: attempt refresh.
s.logger.Warn(ctx, "invalid/expired access token, cannot auto-configure MCP", slog.F("provider", link.ProviderID), slog.Error(validateErr))
tokenErrs[id] = ErrExpiredOrInvalidOAuthToken.Error()
return
}
if validateErr != nil {
errs = multierror.Append(errs, validateErr)
tokenErrs[id] = validateErr.Error()
} else {
tokens[id] = link.OAuthAccessToken
}
}()
continue externalAuthLoop
}
// No link found for this external auth config, so include a generic
// error.
mu.Lock()
tokenErrs[id] = ErrNoExternalAuthLinkFound.Error()
mu.Unlock()
}
wg.Wait()
return &proto.GetMCPServerAccessTokensBatchResponse{
AccessTokens: tokens,
Errors: tokenErrs,
}, errs
}
// IsAuthorized validates a given Coder API key and returns the user ID to which it belongs (if valid).
//
// NOTE: this should really be using the code from [httpmw.ExtractAPIKey]. That function not only validates the key
// but handles many other cases like updating last used, expiry, etc. This code does not currently use it for
// a few reasons:
//
// 1. [httpmw.ExtractAPIKey] relies on keys being given in specific headers [httpmw.APITokenFromRequest] which AI
// bridge requests will not conform to.
// 2. The code mixes many different concerns, and handles HTTP responses too, which is undesirable here.
// 3. The core logic would need to be extracted, but that will surely be a complex & time-consuming distraction right now.
// 4. Once we have an Early Access release of AI Bridge, we need to return to this.
//
// TODO: replace with logic from [httpmw.ExtractAPIKey].
func (s *Server) IsAuthorized(ctx context.Context, in *proto.IsAuthorizedRequest) (*proto.IsAuthorizedResponse, error) {
//nolint:gocritic // AIBridged has specific authz rules.
ctx = dbauthz.AsAIBridged(ctx)
// Key matches expected format.
keyID, keySecret, err := httpmw.SplitAPIToken(in.GetKey())
if err != nil {
return nil, ErrInvalidKey
}
// Key exists.
key, err := s.store.GetAPIKeyByID(ctx, keyID)
if err != nil {
s.logger.Warn(ctx, "failed to retrieve API key by id", slog.F("key_id", keyID), slog.Error(err))
return nil, ErrUnknownKey
}
// Key has not expired.
now := dbtime.Now()
if key.ExpiresAt.Before(now) {
return nil, ErrExpired
}
// Key secret matches.
if !apikey.ValidateHash(key.HashedSecret, keySecret) {
return nil, ErrInvalidKey
}
// User exists.
user, err := s.store.GetUserByID(ctx, key.UserID)
if err != nil {
s.logger.Warn(ctx, "failed to retrieve API key user", slog.F("key_id", keyID), slog.F("user_id", key.UserID), slog.Error(err))
return nil, ErrUnknownUser
}
// User is not deleted or a system user.
if user.Deleted {
return nil, ErrDeletedUser
}
if user.IsSystem {
return nil, ErrSystemUser
}
return &proto.IsAuthorizedResponse{
OwnerId: key.UserID.String(),
ApiKeyId: key.ID,
Username: user.Username,
}, nil
}
// Deprecated: Injected MCP in AI Bridge is deprecated and will be removed in a future release.
func getCoderMCPServerConfig(experiments codersdk.Experiments, accessURL string) (*proto.MCPServerConfig, error) {
// Both the MCP & OAuth2 experiments are currently required in order to use our
// internal MCP server.
if !experiments.Enabled(codersdk.ExperimentMCPServerHTTP) {
return nil, xerrors.Errorf("%q experiment not enabled", codersdk.ExperimentMCPServerHTTP)
}
if !experiments.Enabled(codersdk.ExperimentOAuth2) {
return nil, xerrors.Errorf("%q experiment not enabled", codersdk.ExperimentOAuth2)
}
u, err := url.JoinPath(accessURL, codermcp.MCPEndpoint)
if err != nil {
return nil, xerrors.Errorf("build MCP URL with %q: %w", accessURL, err)
}
return &proto.MCPServerConfig{
Id: aibridged.InternalMCPServerID,
Url: u,
}, nil
}
// credentialKindOrDefault converts the proto credential kind string to
// the database enum, defaulting to "centralized" when the value is
// empty or not a valid enum member.
func credentialKindOrDefault(kind string) database.CredentialKind {
ck := database.CredentialKind(kind)
if !ck.Valid() {
return database.CredentialKindCentralized
}
return ck
}
func metadataToMap(in map[string]*anypb.Any) map[string]any {
meta := make(map[string]any, len(in))
for k, v := range in {
if v == nil {
continue
}
var sv structpb.Value
if err := v.UnmarshalTo(&sv); err == nil {
meta[k] = sv.AsInterface()
}
}
return meta
}
File diff suppressed because it is too large Load Diff
+4
View File
@@ -2198,6 +2198,10 @@ type API struct {
// UsageInserter is a pointer to an atomic pointer because it is passed to
// multiple components.
UsageInserter *atomic.Pointer[usage.Inserter]
// aibridgedHandler is the in-memory aibridge HTTP handler. Set by
// RegisterInMemoryAIBridgedHTTPHandler; read by the enterprise
// /api/v2/aibridge route (license-gated).
aibridgedHandler http.Handler
UpdatesProvider tailnet.WorkspaceUpdatesProvider