diff --git a/coderd/coderd.go b/coderd/coderd.go index d8a52c6f6e..b59ad85a48 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -44,6 +44,8 @@ import ( "tailscale.com/types/key" "tailscale.com/util/singleflight" + "github.com/coder/coder/v2/provisionerd/proto" + "cdr.dev/slog" "github.com/coder/quartz" "github.com/coder/serpent" @@ -95,7 +97,6 @@ import ( "github.com/coder/coder/v2/coderd/workspacestats" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/codersdk/healthsdk" - "github.com/coder/coder/v2/provisionerd/proto" "github.com/coder/coder/v2/provisionersdk" "github.com/coder/coder/v2/site" "github.com/coder/coder/v2/tailnet" @@ -999,6 +1000,11 @@ func New(options *Options) *API { // Experimental routes are not guaranteed to be stable and may change at any time. r.Route("/api/experimental", func(r chi.Router) { + r.NotFound(func(rw http.ResponseWriter, _ *http.Request) { httpapi.RouteNotFound(rw) }) + + // Only this group should be subject to apiKeyMiddleware; aibridged will mount its own + // router and handles key validation in a different fashion. + // See enterprise/x/aibridged/http.go. r.Group(func(r chi.Router) { r.Use(apiKeyMiddleware) r.Route("/aitasks", func(r chi.Router) { diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index e0da9c5ac1..6704c82118 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -175,15 +175,15 @@ func (q *querier) authorizePrebuiltWorkspace(ctx context.Context, action policy. return xerrors.Errorf("authorize context: %w", workspaceErr) } -// authorizeAIBridgeInterceptionUpdate validates that the context's actor matches the initiator of the AIBridgeInterception. +// authorizeAIBridgeInterceptionAction validates that the context's actor matches the initiator of the AIBridgeInterception. // This is used by all of the sub-resources which fall under the [ResourceAibridgeInterception] umbrella. -func (q *querier) authorizeAIBridgeInterceptionUpdate(ctx context.Context, interceptionID uuid.UUID) error { +func (q *querier) authorizeAIBridgeInterceptionAction(ctx context.Context, action policy.Action, interceptionID uuid.UUID) error { inter, err := q.db.GetAIBridgeInterceptionByID(ctx, interceptionID) if err != nil { return xerrors.Errorf("fetch aibridge interception %q: %w", interceptionID, err) } - err = q.authorizeContext(ctx, policy.ActionUpdate, inter.RBACObject()) + err = q.authorizeContext(ctx, action, inter.RBACObject()) if err != nil { return err } @@ -1928,6 +1928,37 @@ func (q *querier) GetAIBridgeInterceptionByID(ctx context.Context, id uuid.UUID) return fetch(q.log, q.auth, q.db.GetAIBridgeInterceptionByID)(ctx, id) } +func (q *querier) GetAIBridgeInterceptions(ctx context.Context) ([]database.AIBridgeInterception, error) { + fetch := func(ctx context.Context, _ any) ([]database.AIBridgeInterception, error) { + return q.db.GetAIBridgeInterceptions(ctx) + } + return fetchWithPostFilter(q.auth, policy.ActionRead, fetch)(ctx, nil) +} + +func (q *querier) GetAIBridgeTokenUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeTokenUsage, error) { + // All aibridge_token_usages records belong to the initiator of their associated interception. + if err := q.authorizeAIBridgeInterceptionAction(ctx, policy.ActionRead, interceptionID); err != nil { + return nil, err + } + return q.db.GetAIBridgeTokenUsagesByInterceptionID(ctx, interceptionID) +} + +func (q *querier) GetAIBridgeToolUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeToolUsage, error) { + // All aibridge_token_usages records belong to the initiator of their associated interception. + if err := q.authorizeAIBridgeInterceptionAction(ctx, policy.ActionRead, interceptionID); err != nil { + return nil, err + } + return q.db.GetAIBridgeToolUsagesByInterceptionID(ctx, interceptionID) +} + +func (q *querier) GetAIBridgeUserPromptsByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeUserPrompt, error) { + // All aibridge_token_usages records belong to the initiator of their associated interception. + if err := q.authorizeAIBridgeInterceptionAction(ctx, policy.ActionRead, interceptionID); err != nil { + return nil, err + } + return q.db.GetAIBridgeUserPromptsByInterceptionID(ctx, interceptionID) +} + func (q *querier) GetAPIKeyByID(ctx context.Context, id string) (database.APIKey, error) { return fetch(q.log, q.auth, q.db.GetAPIKeyByID)(ctx, id) } @@ -3813,7 +3844,7 @@ func (q *querier) InsertAIBridgeInterception(ctx context.Context, arg database.I func (q *querier) InsertAIBridgeTokenUsage(ctx context.Context, arg database.InsertAIBridgeTokenUsageParams) error { // All aibridge_token_usages records belong to the initiator of their associated interception. - if err := q.authorizeAIBridgeInterceptionUpdate(ctx, arg.InterceptionID); err != nil { + if err := q.authorizeAIBridgeInterceptionAction(ctx, policy.ActionUpdate, arg.InterceptionID); err != nil { return err } return q.db.InsertAIBridgeTokenUsage(ctx, arg) @@ -3821,7 +3852,7 @@ func (q *querier) InsertAIBridgeTokenUsage(ctx context.Context, arg database.Ins func (q *querier) InsertAIBridgeToolUsage(ctx context.Context, arg database.InsertAIBridgeToolUsageParams) error { // All aibridge_tool_usages records belong to the initiator of their associated interception. - if err := q.authorizeAIBridgeInterceptionUpdate(ctx, arg.InterceptionID); err != nil { + if err := q.authorizeAIBridgeInterceptionAction(ctx, policy.ActionUpdate, arg.InterceptionID); err != nil { return err } return q.db.InsertAIBridgeToolUsage(ctx, arg) @@ -3829,7 +3860,7 @@ func (q *querier) InsertAIBridgeToolUsage(ctx context.Context, arg database.Inse func (q *querier) InsertAIBridgeUserPrompt(ctx context.Context, arg database.InsertAIBridgeUserPromptParams) error { // All aibridge_user_prompts records belong to the initiator of their associated interception. - if err := q.authorizeAIBridgeInterceptionUpdate(ctx, arg.InterceptionID); err != nil { + if err := q.authorizeAIBridgeInterceptionAction(ctx, policy.ActionUpdate, arg.InterceptionID); err != nil { return err } return q.db.InsertAIBridgeUserPrompt(ctx, arg) diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 1eb92e3680..a20efb8be4 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -4335,10 +4335,10 @@ func TestInsertAPIKey_AsPrebuildsUser(t *testing.T) { func (s *MethodTestSuite) TestAIBridge() { s.Run("GetAIBridgeInterceptionByID", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - sessID := uuid.UUID{2} - sess := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: sessID}) - db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), sessID).Return(sess, nil).AnyTimes() - check.Args(sessID).Asserts(sess, policy.ActionRead).Returns(sess) + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() + check.Args(intID).Asserts(intc, policy.ActionRead).Returns(intc) })) s.Run("InsertAIBridgeInterception", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { @@ -4348,39 +4348,76 @@ func (s *MethodTestSuite) TestAIBridge() { user.IsSystem = false user.Deleted = false - sessID := uuid.UUID{2} - sess := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: sessID, InitiatorID: initID}) + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID, InitiatorID: initID}) - params := database.InsertAIBridgeInterceptionParams{ID: sess.ID, InitiatorID: sess.InitiatorID, Provider: sess.Provider, Model: sess.Model} + params := database.InsertAIBridgeInterceptionParams{ID: intc.ID, InitiatorID: intc.InitiatorID, Provider: intc.Provider, Model: intc.Model} db.EXPECT().GetUserByID(gomock.Any(), initID).Return(user, nil).AnyTimes() // Validation. - db.EXPECT().InsertAIBridgeInterception(gomock.Any(), params).Return(sess, nil).AnyTimes() - check.Args(params).Asserts(sess, policy.ActionCreate) + db.EXPECT().InsertAIBridgeInterception(gomock.Any(), params).Return(intc, nil).AnyTimes() + check.Args(params).Asserts(intc, policy.ActionCreate) })) s.Run("InsertAIBridgeTokenUsage", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - sessID := uuid.UUID{2} - sess := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: sessID}) - params := database.InsertAIBridgeTokenUsageParams{InterceptionID: sess.ID} - db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), sessID).Return(sess, nil).AnyTimes() // Validation. + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + params := database.InsertAIBridgeTokenUsageParams{InterceptionID: intc.ID} + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() // Validation. db.EXPECT().InsertAIBridgeTokenUsage(gomock.Any(), params).Return(nil).AnyTimes() - check.Args(params).Asserts(sess, policy.ActionUpdate) + check.Args(params).Asserts(intc, policy.ActionUpdate) })) s.Run("InsertAIBridgeToolUsage", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - sessID := uuid.UUID{2} - sess := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: sessID}) - params := database.InsertAIBridgeToolUsageParams{InterceptionID: sess.ID} - db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), sessID).Return(sess, nil).AnyTimes() // Validation. + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + params := database.InsertAIBridgeToolUsageParams{InterceptionID: intc.ID} + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() // Validation. db.EXPECT().InsertAIBridgeToolUsage(gomock.Any(), params).Return(nil).AnyTimes() - check.Args(params).Asserts(sess, policy.ActionUpdate) + check.Args(params).Asserts(intc, policy.ActionUpdate) })) s.Run("InsertAIBridgeUserPrompt", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - sessID := uuid.UUID{2} - sess := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: sessID}) - params := database.InsertAIBridgeUserPromptParams{InterceptionID: sess.ID} - db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), sessID).Return(sess, nil).AnyTimes() // Validation. + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + params := database.InsertAIBridgeUserPromptParams{InterceptionID: intc.ID} + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() // Validation. db.EXPECT().InsertAIBridgeUserPrompt(gomock.Any(), params).Return(nil).AnyTimes() - check.Args(params).Asserts(sess, policy.ActionUpdate) + check.Args(params).Asserts(intc, policy.ActionUpdate) + })) + + s.Run("GetAIBridgeInterceptions", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + a := testutil.Fake(s.T(), faker, database.AIBridgeInterception{}) + b := testutil.Fake(s.T(), faker, database.AIBridgeInterception{}) + db.EXPECT().GetAIBridgeInterceptions(gomock.Any()).Return([]database.AIBridgeInterception{a, b}, nil).AnyTimes() + check.Args().Asserts(a, policy.ActionRead, b, policy.ActionRead).Returns([]database.AIBridgeInterception{a, b}) + })) + + s.Run("GetAIBridgeTokenUsagesByInterceptionID", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + tok := testutil.Fake(s.T(), faker, database.AIBridgeTokenUsage{InterceptionID: intID}) + toks := []database.AIBridgeTokenUsage{tok} + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() // Validation. + db.EXPECT().GetAIBridgeTokenUsagesByInterceptionID(gomock.Any(), intID).Return(toks, nil).AnyTimes() + check.Args(intID).Asserts(intc, policy.ActionRead).Returns(toks) + })) + + s.Run("GetAIBridgeToolUsagesByInterceptionID", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + tool := testutil.Fake(s.T(), faker, database.AIBridgeToolUsage{InterceptionID: intID}) + tools := []database.AIBridgeToolUsage{tool} + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() // Validation. + db.EXPECT().GetAIBridgeToolUsagesByInterceptionID(gomock.Any(), intID).Return(tools, nil).AnyTimes() + check.Args(intID).Asserts(intc, policy.ActionRead).Returns(tools) + })) + + s.Run("GetAIBridgeUserPromptsByInterceptionID", s.Mocked(func(db *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + intID := uuid.UUID{2} + intc := testutil.Fake(s.T(), faker, database.AIBridgeInterception{ID: intID}) + pr := testutil.Fake(s.T(), faker, database.AIBridgeUserPrompt{InterceptionID: intID}) + prs := []database.AIBridgeUserPrompt{pr} + db.EXPECT().GetAIBridgeInterceptionByID(gomock.Any(), intID).Return(intc, nil).AnyTimes() // Validation. + db.EXPECT().GetAIBridgeUserPromptsByInterceptionID(gomock.Any(), intID).Return(prs, nil).AnyTimes() + check.Args(intID).Asserts(intc, policy.ActionRead).Returns(prs) })) } diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 6f520a904a..d7d20f2355 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -593,6 +593,34 @@ func (m queryMetricsStore) GetAIBridgeInterceptionByID(ctx context.Context, id u return r0, r1 } +func (m queryMetricsStore) GetAIBridgeInterceptions(ctx context.Context) ([]database.AIBridgeInterception, error) { + start := time.Now() + r0, r1 := m.s.GetAIBridgeInterceptions(ctx) + m.queryLatencies.WithLabelValues("GetAIBridgeInterceptions").Observe(time.Since(start).Seconds()) + return r0, r1 +} + +func (m queryMetricsStore) GetAIBridgeTokenUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeTokenUsage, error) { + start := time.Now() + r0, r1 := m.s.GetAIBridgeTokenUsagesByInterceptionID(ctx, interceptionID) + m.queryLatencies.WithLabelValues("GetAIBridgeTokenUsagesByInterceptionID").Observe(time.Since(start).Seconds()) + return r0, r1 +} + +func (m queryMetricsStore) GetAIBridgeToolUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeToolUsage, error) { + start := time.Now() + r0, r1 := m.s.GetAIBridgeToolUsagesByInterceptionID(ctx, interceptionID) + m.queryLatencies.WithLabelValues("GetAIBridgeToolUsagesByInterceptionID").Observe(time.Since(start).Seconds()) + return r0, r1 +} + +func (m queryMetricsStore) GetAIBridgeUserPromptsByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeUserPrompt, error) { + start := time.Now() + r0, r1 := m.s.GetAIBridgeUserPromptsByInterceptionID(ctx, interceptionID) + m.queryLatencies.WithLabelValues("GetAIBridgeUserPromptsByInterceptionID").Observe(time.Since(start).Seconds()) + return r0, r1 +} + func (m queryMetricsStore) GetAPIKeyByID(ctx context.Context, id string) (database.APIKey, error) { start := time.Now() apiKey, err := m.s.GetAPIKeyByID(ctx, id) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index e96759f828..66de8ec6e5 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -1109,6 +1109,66 @@ func (mr *MockStoreMockRecorder) GetAIBridgeInterceptionByID(ctx, id any) *gomoc return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIBridgeInterceptionByID", reflect.TypeOf((*MockStore)(nil).GetAIBridgeInterceptionByID), ctx, id) } +// GetAIBridgeInterceptions mocks base method. +func (m *MockStore) GetAIBridgeInterceptions(ctx context.Context) ([]database.AIBridgeInterception, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAIBridgeInterceptions", ctx) + ret0, _ := ret[0].([]database.AIBridgeInterception) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAIBridgeInterceptions indicates an expected call of GetAIBridgeInterceptions. +func (mr *MockStoreMockRecorder) GetAIBridgeInterceptions(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIBridgeInterceptions", reflect.TypeOf((*MockStore)(nil).GetAIBridgeInterceptions), ctx) +} + +// GetAIBridgeTokenUsagesByInterceptionID mocks base method. +func (m *MockStore) GetAIBridgeTokenUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeTokenUsage, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAIBridgeTokenUsagesByInterceptionID", ctx, interceptionID) + ret0, _ := ret[0].([]database.AIBridgeTokenUsage) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAIBridgeTokenUsagesByInterceptionID indicates an expected call of GetAIBridgeTokenUsagesByInterceptionID. +func (mr *MockStoreMockRecorder) GetAIBridgeTokenUsagesByInterceptionID(ctx, interceptionID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIBridgeTokenUsagesByInterceptionID", reflect.TypeOf((*MockStore)(nil).GetAIBridgeTokenUsagesByInterceptionID), ctx, interceptionID) +} + +// GetAIBridgeToolUsagesByInterceptionID mocks base method. +func (m *MockStore) GetAIBridgeToolUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeToolUsage, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAIBridgeToolUsagesByInterceptionID", ctx, interceptionID) + ret0, _ := ret[0].([]database.AIBridgeToolUsage) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAIBridgeToolUsagesByInterceptionID indicates an expected call of GetAIBridgeToolUsagesByInterceptionID. +func (mr *MockStoreMockRecorder) GetAIBridgeToolUsagesByInterceptionID(ctx, interceptionID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIBridgeToolUsagesByInterceptionID", reflect.TypeOf((*MockStore)(nil).GetAIBridgeToolUsagesByInterceptionID), ctx, interceptionID) +} + +// GetAIBridgeUserPromptsByInterceptionID mocks base method. +func (m *MockStore) GetAIBridgeUserPromptsByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]database.AIBridgeUserPrompt, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetAIBridgeUserPromptsByInterceptionID", ctx, interceptionID) + ret0, _ := ret[0].([]database.AIBridgeUserPrompt) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetAIBridgeUserPromptsByInterceptionID indicates an expected call of GetAIBridgeUserPromptsByInterceptionID. +func (mr *MockStoreMockRecorder) GetAIBridgeUserPromptsByInterceptionID(ctx, interceptionID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIBridgeUserPromptsByInterceptionID", reflect.TypeOf((*MockStore)(nil).GetAIBridgeUserPromptsByInterceptionID), ctx, interceptionID) +} + // GetAPIKeyByID mocks base method. func (m *MockStore) GetAPIKeyByID(ctx context.Context, id string) (database.APIKey, error) { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 2cd8bd3b25..7d77f9aad5 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -149,6 +149,10 @@ type sqlcQuerier interface { // and returns the preset with the most parameters (largest subset). FindMatchingPresetID(ctx context.Context, arg FindMatchingPresetIDParams) (uuid.UUID, error) GetAIBridgeInterceptionByID(ctx context.Context, id uuid.UUID) (AIBridgeInterception, error) + GetAIBridgeInterceptions(ctx context.Context) ([]AIBridgeInterception, error) + GetAIBridgeTokenUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeTokenUsage, error) + GetAIBridgeToolUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeToolUsage, error) + GetAIBridgeUserPromptsByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeUserPrompt, error) GetAPIKeyByID(ctx context.Context, id string) (APIKey, error) // there is no unique constraint on empty token names GetAPIKeyByName(ctx context.Context, arg GetAPIKeyByNameParams) (APIKey, error) diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index f2e4e33c33..edc5966f55 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -129,6 +129,147 @@ func (q *sqlQuerier) GetAIBridgeInterceptionByID(ctx context.Context, id uuid.UU return i, err } +const getAIBridgeInterceptions = `-- name: GetAIBridgeInterceptions :many +SELECT id, initiator_id, provider, model, started_at, metadata FROM aibridge_interceptions +` + +func (q *sqlQuerier) GetAIBridgeInterceptions(ctx context.Context) ([]AIBridgeInterception, error) { + rows, err := q.db.QueryContext(ctx, getAIBridgeInterceptions) + if err != nil { + return nil, err + } + defer rows.Close() + var items []AIBridgeInterception + for rows.Next() { + var i AIBridgeInterception + if err := rows.Scan( + &i.ID, + &i.InitiatorID, + &i.Provider, + &i.Model, + &i.StartedAt, + &i.Metadata, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const getAIBridgeTokenUsagesByInterceptionID = `-- name: GetAIBridgeTokenUsagesByInterceptionID :many +SELECT id, interception_id, provider_response_id, input_tokens, output_tokens, metadata, created_at FROM aibridge_token_usages WHERE interception_id = $1::uuid +` + +func (q *sqlQuerier) GetAIBridgeTokenUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeTokenUsage, error) { + rows, err := q.db.QueryContext(ctx, getAIBridgeTokenUsagesByInterceptionID, interceptionID) + if err != nil { + return nil, err + } + defer rows.Close() + var items []AIBridgeTokenUsage + for rows.Next() { + var i AIBridgeTokenUsage + if err := rows.Scan( + &i.ID, + &i.InterceptionID, + &i.ProviderResponseID, + &i.InputTokens, + &i.OutputTokens, + &i.Metadata, + &i.CreatedAt, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const getAIBridgeToolUsagesByInterceptionID = `-- name: GetAIBridgeToolUsagesByInterceptionID :many +SELECT id, interception_id, provider_response_id, server_url, tool, input, injected, invocation_error, metadata, created_at FROM aibridge_tool_usages WHERE interception_id = $1::uuid +` + +func (q *sqlQuerier) GetAIBridgeToolUsagesByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeToolUsage, error) { + rows, err := q.db.QueryContext(ctx, getAIBridgeToolUsagesByInterceptionID, interceptionID) + if err != nil { + return nil, err + } + defer rows.Close() + var items []AIBridgeToolUsage + for rows.Next() { + var i AIBridgeToolUsage + if err := rows.Scan( + &i.ID, + &i.InterceptionID, + &i.ProviderResponseID, + &i.ServerUrl, + &i.Tool, + &i.Input, + &i.Injected, + &i.InvocationError, + &i.Metadata, + &i.CreatedAt, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const getAIBridgeUserPromptsByInterceptionID = `-- name: GetAIBridgeUserPromptsByInterceptionID :many +SELECT id, interception_id, provider_response_id, prompt, metadata, created_at FROM aibridge_user_prompts WHERE interception_id = $1::uuid +` + +func (q *sqlQuerier) GetAIBridgeUserPromptsByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeUserPrompt, error) { + rows, err := q.db.QueryContext(ctx, getAIBridgeUserPromptsByInterceptionID, interceptionID) + if err != nil { + return nil, err + } + defer rows.Close() + var items []AIBridgeUserPrompt + for rows.Next() { + var i AIBridgeUserPrompt + if err := rows.Scan( + &i.ID, + &i.InterceptionID, + &i.ProviderResponseID, + &i.Prompt, + &i.Metadata, + &i.CreatedAt, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + const insertAIBridgeInterception = `-- name: InsertAIBridgeInterception :one INSERT INTO aibridge_interceptions (id, initiator_id, provider, model, metadata, started_at) VALUES ($1::uuid, $2::uuid, $3, $4, COALESCE($5::jsonb, '{}'::jsonb), $6) diff --git a/coderd/database/queries/aibridge.sql b/coderd/database/queries/aibridge.sql index 863a5f3051..dab2073cea 100644 --- a/coderd/database/queries/aibridge.sql +++ b/coderd/database/queries/aibridge.sql @@ -26,3 +26,15 @@ INSERT INTO aibridge_tool_usages ( -- name: GetAIBridgeInterceptionByID :one SELECT * FROM aibridge_interceptions WHERE id = @id::uuid; + +-- name: GetAIBridgeInterceptions :many +SELECT * FROM aibridge_interceptions; + +-- name: GetAIBridgeTokenUsagesByInterceptionID :many +SELECT * FROM aibridge_token_usages WHERE interception_id = @interception_id::uuid; + +-- name: GetAIBridgeUserPromptsByInterceptionID :many +SELECT * FROM aibridge_user_prompts WHERE interception_id = @interception_id::uuid; + +-- name: GetAIBridgeToolUsagesByInterceptionID :many +SELECT * FROM aibridge_tool_usages WHERE interception_id = @interception_id::uuid; diff --git a/enterprise/cli/aibridged.go b/enterprise/cli/aibridged.go new file mode 100644 index 0000000000..9e59327039 --- /dev/null +++ b/enterprise/cli/aibridged.go @@ -0,0 +1,47 @@ +//go:build !slim + +package cli + +import ( + "context" + + "golang.org/x/xerrors" + + "github.com/coder/aibridge" + "github.com/coder/coder/v2/enterprise/coderd" + "github.com/coder/coder/v2/enterprise/x/aibridged" +) + +func newAIBridgeDaemon(coderAPI *coderd.API) (*aibridged.Server, error) { + ctx := context.Background() + coderAPI.Logger.Debug(ctx, "starting in-memory aibridge daemon") + + logger := coderAPI.Logger.Named("aibridged") + + // Setup supported providers. + providers := []aibridge.Provider{ + aibridge.NewOpenAIProvider(aibridge.ProviderConfig{ + BaseURL: coderAPI.DeploymentValues.AI.BridgeConfig.OpenAI.BaseURL.String(), + Key: coderAPI.DeploymentValues.AI.BridgeConfig.OpenAI.Key.String(), + }), + aibridge.NewAnthropicProvider(aibridge.ProviderConfig{ + BaseURL: coderAPI.DeploymentValues.AI.BridgeConfig.Anthropic.BaseURL.String(), + Key: coderAPI.DeploymentValues.AI.BridgeConfig.Anthropic.Key.String(), + }), + } + + // Create pool for reusable stateful [aibridge.RequestBridge] instances (one per user). + pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger.Named("pool")) // TODO: configurable. + if err != nil { + return nil, xerrors.Errorf("create request pool: %w", err) + } + + // Create daemon. + srv, err := aibridged.New(ctx, pool, func(dialCtx context.Context) (aibridged.DRPCClient, error) { + return coderAPI.CreateInMemoryAIBridgeServer(dialCtx) + }, logger) + if err != nil { + return nil, xerrors.Errorf("start in-memory aibridge daemon: %w", err) + } + return srv, nil +} diff --git a/enterprise/cli/server.go b/enterprise/cli/server.go index f58ec86b58..1e3c71f6eb 100644 --- a/enterprise/cli/server.go +++ b/enterprise/cli/server.go @@ -7,6 +7,7 @@ import ( "database/sql" "encoding/base64" "errors" + "fmt" "io" "net/url" @@ -15,6 +16,7 @@ import ( "tailscale.com/types/key" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/cryptorand" "github.com/coder/coder/v2/enterprise/audit" "github.com/coder/coder/v2/enterprise/audit/backends" @@ -23,6 +25,7 @@ import ( "github.com/coder/coder/v2/enterprise/coderd/usage" "github.com/coder/coder/v2/enterprise/dbcrypt" "github.com/coder/coder/v2/enterprise/trialer" + "github.com/coder/coder/v2/enterprise/x/aibridged" "github.com/coder/coder/v2/tailnet" "github.com/coder/quartz" "github.com/coder/serpent" @@ -143,6 +146,33 @@ func (r *RootCmd) Server(_ func()) *serpent.Command { } closers.Add(publisher) + experiments := agplcoderd.ReadExperiments(options.Logger, options.DeploymentValues.Experiments.Value()) + + var aibridgeDaemon *aibridged.Server + // In-memory aibridge daemon. + if options.DeploymentValues.AI.BridgeConfig.Enabled { + if experiments.Enabled(codersdk.ExperimentAIBridge) { + aibridgeDaemon, err = newAIBridgeDaemon(api) + if err != nil { + return nil, nil, xerrors.Errorf("create aibridged: %w", err) + } + + api.RegisterInMemoryAIBridgedHTTPHandler(aibridgeDaemon) + + // When running as an in-memory daemon, the HTTP handler is wired into the + // coderd API and therefore is subject to its context. Calling Close() on + // aibridged will NOT affect in-flight requests but those will be closed once + // the API server is itself shutdown. + closers.Add(aibridgeDaemon) + } else { + api.Logger.Warn(ctx, fmt.Sprintf("CODER_AIBRIDGE_ENABLED=true but experiment %q not enabled", codersdk.ExperimentAIBridge)) + } + } else { + if experiments.Enabled(codersdk.ExperimentAIBridge) { + api.Logger.Warn(ctx, "aibridge experiment enabled but CODER_AIBRIDGE_ENABLED=false") + } + } + return api.AGPL, closers, nil }) diff --git a/enterprise/coderd/aibridged.go b/enterprise/coderd/aibridged.go new file mode 100644 index 0000000000..51afb24e6a --- /dev/null +++ b/enterprise/coderd/aibridged.go @@ -0,0 +1,118 @@ +package coderd + +import ( + "context" + "errors" + "io" + "net/http" + + "github.com/go-chi/chi/v5" + "golang.org/x/xerrors" + "storj.io/drpc/drpcmux" + "storj.io/drpc/drpcserver" + + "cdr.dev/slog" + + "github.com/coder/coder/v2/coderd/httpmw" + "github.com/coder/coder/v2/coderd/tracing" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/codersdk/drpcsdk" + "github.com/coder/coder/v2/enterprise/x/aibridged" + aibridgedproto "github.com/coder/coder/v2/enterprise/x/aibridged/proto" + "github.com/coder/coder/v2/enterprise/x/aibridgedserver" +) + +// 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 *aibridged.Server) { + if srv == nil { + panic("aibridged cannot be nil") + } + + if api.AGPL.RootHandler == nil { + panic("api.RootHandler cannot be nil") + } + + aibridgeEndpoint := "/api/experimental/aibridge" + + r := chi.NewRouter() + r.Group(func(r chi.Router) { + r.Use(httpmw.RequireExperiment(api.AGPL.Experiments, codersdk.ExperimentAIBridge)) + r.HandleFunc("/*", http.StripPrefix(aibridgeEndpoint, srv).ServeHTTP) + }) + + api.AGPL.RootHandler.Mount(aibridgeEndpoint, r) +} + +// 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.ExternalAuthConfigs, api.AGPL.Experiments) + 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.AGPL.WebsocketWaitMutex.Lock() + api.AGPL.WebsocketWaitGroup.Add(1) + api.AGPL.WebsocketWaitMutex.Unlock() + go func() { + defer api.AGPL.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 +} diff --git a/enterprise/x/aibridged/aibridged.go b/enterprise/x/aibridged/aibridged.go index 04ae617c20..a1fa4022ff 100644 --- a/enterprise/x/aibridged/aibridged.go +++ b/enterprise/x/aibridged/aibridged.go @@ -3,6 +3,7 @@ package aibridged import ( "context" "errors" + "io" "net/http" "sync" "time" @@ -14,6 +15,8 @@ import ( "github.com/coder/retry" ) +var _ io.Closer = &Server{} + // Server provides the AI Bridge functionality. // It is responsible for: // - receiving requests on /api/experimental/aibridged/* // TODO: update endpoint once out of experimental @@ -183,3 +186,10 @@ func (s *Server) Shutdown(ctx context.Context) error { }) 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) +} diff --git a/enterprise/x/aibridged/aibridged_integration_test.go b/enterprise/x/aibridged/aibridged_integration_test.go new file mode 100644 index 0000000000..69d7627e04 --- /dev/null +++ b/enterprise/x/aibridged/aibridged_integration_test.go @@ -0,0 +1,242 @@ +package aibridged_test + +import ( + "bytes" + "context" + "fmt" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/coder/aibridge" + "github.com/coder/coder/v2/coderd/coderdtest" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/coderd/database/dbtestutil" + "github.com/coder/coder/v2/coderd/database/dbtime" + "github.com/coder/coder/v2/coderd/externalauth" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" + "github.com/coder/coder/v2/enterprise/x/aibridged" + "github.com/coder/coder/v2/testutil" +) + +// TestIntegration is not an exhaustive test against the upstream AI providers' SDKs (see coder/aibridge for those). +// This test validates that: +// - intercepted requests can be authenticated/authorized +// - requests can be routed to an appropriate handler +// - responses can be returned as expected +// - interceptions are logged, as well as their related prompt, token, and tool calls +// - MCP server configurations are returned as expected +func TestIntegration(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + + // Create mock MCP server. + var mcpTokenReceived string + mockMCPServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + t.Logf("Mock MCP server received request: %s %s", r.Method, r.URL.Path) + + if r.Method == http.MethodPost && r.URL.Path == "/" { + // Mark that init was called. + mcpTokenReceived = r.Header.Get("Authorization") + t.Log("MCP init request received") + + // Return a basic MCP init response. + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Mcp-Session-Id", "test-session-123") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "jsonrpc": "2.0", + "id": 1, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": { + "name": "test-mcp-server", + "version": "1.0.0" + } + } + }`)) + } + })) + t.Cleanup(mockMCPServer.Close) + t.Logf("Mock MCP server running at: %s", mockMCPServer.URL) + + // Set up mock OpenAI server that returns a tool call response. + mockOpenAI := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "id": "chatcmpl-BwkyFElDIr1egmFyfQ9z4vPBto7m2", + "object": "chat.completion", + "created": 1753343279, + "model": "gpt-4.1-2025-04-14", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": null, + "tool_calls": [ + { + "id": "call_KjzAbhiZC6nk81tQzL7pwlpc", + "type": "function", + "function": { + "name": "read_file", + "arguments": "{\"path\":\"README.md\"}" + } + } + ], + "refusal": null, + "annotations": [] + }, + "logprobs": null, + "finish_reason": "tool_calls" + } + ], + "usage": { + "prompt_tokens": 60, + "completion_tokens": 15, + "total_tokens": 75, + "prompt_tokens_details": { + "cached_tokens": 0, + "audio_tokens": 0 + }, + "completion_tokens_details": { + "reasoning_tokens": 0, + "audio_tokens": 0, + "accepted_prediction_tokens": 0, + "rejected_prediction_tokens": 0 + } + }, + "service_tier": "default", + "system_fingerprint": "fp_b3f1157249" +}`)) + })) + t.Cleanup(mockOpenAI.Close) + + db, ps := dbtestutil.NewDB(t) + client, _, api, firstUser := coderdenttest.NewWithAPI(t, &coderdenttest.Options{ + Options: &coderdtest.Options{ + Database: db, + Pubsub: ps, + ExternalAuthConfigs: []*externalauth.Config{ + { + InstrumentedOAuth2Config: &testutil.OAuth2Config{}, + ID: "mock", + Type: "mock", + DisplayName: "Mock", + MCPURL: mockMCPServer.URL, + }, + }, + }, + }) + + userClient, user := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID) + + // Create an API token for the user. + apiKey, err := userClient.CreateToken(ctx, "me", codersdk.CreateTokenRequest{ + TokenName: fmt.Sprintf("test-key-%d", time.Now().UnixNano()), + Lifetime: time.Hour, + Scope: codersdk.APIKeyScopeAll, + }) + require.NoError(t, err) + + // Create external auth link for the user. + authLink, err := db.InsertExternalAuthLink(dbauthz.AsSystemRestricted(ctx), database.InsertExternalAuthLinkParams{ + ProviderID: "mock", + UserID: user.ID, + CreatedAt: dbtime.Now(), + UpdatedAt: dbtime.Now(), + OAuthAccessToken: "test-mock-token", + OAuthRefreshToken: "test-refresh-token", + OAuthExpiry: dbtime.Now().Add(time.Hour), + }) + require.NoError(t, err) + + // Create aibridge server & client. + aiBridgeClient, err := api.CreateInMemoryAIBridgeServer(ctx) + require.NoError(t, err) + + logger := testutil.Logger(t) + providers := []aibridge.Provider{aibridge.NewOpenAIProvider(aibridge.ProviderConfig{BaseURL: mockOpenAI.URL})} + pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger) + require.NoError(t, err) + + // Given: aibridged is started. + srv, err := aibridged.New(t.Context(), pool, func(ctx context.Context) (aibridged.DRPCClient, error) { + return aiBridgeClient, nil + }, logger) + require.NoError(t, err, "create new aibridged") + t.Cleanup(func() { + _ = srv.Shutdown(ctx) + }) + + // When: a request is made to aibridged. + req, err := http.NewRequestWithContext(ctx, http.MethodPost, "/openai/v1/chat/completions", bytes.NewBufferString(`{ + "messages": [ + { + "role": "user", + "content": "how large is the README.md file in my current path" + } + ], + "model": "gpt-4.1", + "tools": [ + { + "type": "function", + "function": { + "name": "read_file", + "description": "Read the contents of a file at the given path.", + "parameters": { + "properties": { + "path": { + "type": "string" + } + }, + "required": [ + "path" + ], + "type": "object" + } + } + } + ] +}`)) + require.NoError(t, err, "make request to test server") + req.Header.Add("Authorization", "Bearer "+apiKey.Key) + req.Header.Add("Accept", "application/json") + + // When: aibridged handles the request. + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + // Then: the interception & related records are stored. + interceptions, err := db.GetAIBridgeInterceptions(ctx) + require.NoError(t, err) + require.Len(t, interceptions, 1) + + prompts, err := db.GetAIBridgeUserPromptsByInterceptionID(ctx, interceptions[0].ID) + require.NoError(t, err) + require.Len(t, prompts, 1) + require.Equal(t, prompts[0].Prompt, "how large is the README.md file in my current path") + + tokens, err := db.GetAIBridgeTokenUsagesByInterceptionID(ctx, interceptions[0].ID) + require.NoError(t, err) + require.Len(t, tokens, 1) + require.EqualValues(t, tokens[0].InputTokens, 60) + require.EqualValues(t, tokens[0].OutputTokens, 15) + + tools, err := db.GetAIBridgeToolUsagesByInterceptionID(ctx, interceptions[0].ID) + require.NoError(t, err) + require.Len(t, tools, 1) + require.False(t, tools[0].Injected) + + // Then: the MCP server was initialized. + require.Contains(t, mcpTokenReceived, authLink.OAuthAccessToken, "mock MCP server not requested") +} diff --git a/enterprise/x/aibridged/pool.go b/enterprise/x/aibridged/pool.go index 97c08703c7..309f8fc61f 100644 --- a/enterprise/x/aibridged/pool.go +++ b/enterprise/x/aibridged/pool.go @@ -2,7 +2,6 @@ package aibridged import ( "context" - "errors" "net/http" "sync" "time" @@ -73,13 +72,7 @@ func NewCachedBridgePool(options PoolOptions, providers []aibridge.Provider, log // Run the eviction in the background since ristretto blocks sets until a free slot is available. go func() { - if err := item.Value.Shutdown(shutdownCtx); err != nil { - if errors.Is(err, context.DeadlineExceeded) { - logger.Debug(shutdownCtx, "bridge shutdown timed out") - } else { - logger.Debug(shutdownCtx, "bridge shutdown failed", slog.Error(err)) - } - } + _ = item.Value.Shutdown(shutdownCtx) }() }, })