mirror of
https://github.com/coder/coder.git
synced 2026-09-21 20:51:01 +08:00
feat: initialize aibridged & mount API handler (#19798)
Addresses https://github.com/coder/internal/issues/987
This commit is contained in:
+7
-1
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
})
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}()
|
||||
},
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user