mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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;
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
package aibridged
|
||||
|
||||
import "github.com/google/uuid"
|
||||
|
||||
type Request struct {
|
||||
SessionKey string
|
||||
APIKeyID string
|
||||
InitiatorID uuid.UUID
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
package aibridged
|
||||
|
||||
import "github.com/coder/coder/v2/coderd/aibridged/proto"
|
||||
|
||||
type DRPCServer interface {
|
||||
proto.DRPCRecorderServer
|
||||
proto.DRPCMCPConfiguratorServer
|
||||
proto.DRPCAuthorizerServer
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user