mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: session list API (#23202)
<!-- If you have used AI to produce some or all of this PR, please ensure you have read our [AI Contribution guidelines](https://coder.com/docs/about/contributing/AI_CONTRIBUTING) before submitting. --> _Disclaimer:_ _initially_ _produced_ _by_ _Claude_ _Opus_ _4\.6,_ _heavily_ _modified_ _and_ _reviewed_ _by_ _me._ Closes https://github.com/coder/internal/issues/1360 Adds a new `/api/v2/aibridge/sessions` API which returns "sessions". Sessions, as defined in the [RFC](https://www.notion.so/coderhq/AI-Bridge-Sessions-Threads-2ccd579be59280f28021d3baf7472fbe?source=copy_link), are a set of interceptions logically grouped by a session key issued by the client. The API design for this endpoint was done in [this doc](https://github.com/coder/internal/issues/1360). If the client has not provided a session ID, we will revert to the thread root ID, and if that's not present we use the interception's own ID (i.e. a session of a single interception - which is effectively what we show currently in our `/api/v2/aibridge/interceptions` API). The SQL query looks gnarly but it's relatively simple, and seems to perform well (~200ms) even when I import dogfood's `aibridge_*` tables into my workspace. If we need to improve performance on this later we can investigate materialized views, perhaps, but for now I don't think it's warranted. --- _The PR looks large but it's got a lot of generated code; the actual changes aren't huge._
This commit is contained in:
@@ -2,6 +2,7 @@ package coderd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"time"
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
@@ -22,8 +24,10 @@ import (
|
||||
|
||||
const (
|
||||
maxListInterceptionsLimit = 1000
|
||||
maxListSessionsLimit = 1000
|
||||
maxListModelsLimit = 1000
|
||||
defaultListInterceptionsLimit = 100
|
||||
defaultListSessionsLimit = 100
|
||||
defaultListModelsLimit = 100
|
||||
// aiBridgeRateLimitWindow is the fixed duration for rate limiting AI Bridge
|
||||
// requests. This is hardcoded to keep configuration simple.
|
||||
@@ -43,6 +47,7 @@ func aibridgeHandler(api *API, middlewares ...func(http.Handler) http.Handler) f
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(middlewares...)
|
||||
r.Get("/interceptions", api.aiBridgeListInterceptions)
|
||||
r.Get("/sessions", api.aiBridgeListSessions)
|
||||
r.Get("/models", api.aiBridgeListModels)
|
||||
})
|
||||
|
||||
@@ -176,6 +181,130 @@ func (api *API) aiBridgeListInterceptions(rw http.ResponseWriter, r *http.Reques
|
||||
})
|
||||
}
|
||||
|
||||
// aiBridgeListSessions returns AI Bridge sessions (aggregated interceptions).
|
||||
//
|
||||
// @Summary List AI Bridge sessions
|
||||
// @ID list-ai-bridge-sessions
|
||||
// @Security CoderSessionToken
|
||||
// @Produce json
|
||||
// @Tags AI Bridge
|
||||
// @Param q query string false "Search query in the format `key:value`. Available keys are: initiator, provider, model, client, session_id, started_after, started_before."
|
||||
// @Param limit query int false "Page limit"
|
||||
// @Param after_session_id query string false "Cursor pagination after session ID (cannot be used with offset)"
|
||||
// @Param offset query int false "Offset pagination (cannot be used with after_session_id)"
|
||||
// @Success 200 {object} codersdk.AIBridgeListSessionsResponse
|
||||
// @Router /aibridge/sessions [get]
|
||||
func (api *API) aiBridgeListSessions(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
apiKey := httpmw.APIKey(r)
|
||||
|
||||
page, ok := coderd.ParsePagination(rw, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
afterSessionID := r.URL.Query().Get("after_session_id")
|
||||
if afterSessionID != "" && page.Offset != 0 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Query parameters have invalid values.",
|
||||
Detail: "Cannot use both after_session_id and offset pagination in the same request.",
|
||||
})
|
||||
return
|
||||
}
|
||||
if page.Limit == 0 {
|
||||
page.Limit = defaultListSessionsLimit
|
||||
}
|
||||
if page.Limit > maxListSessionsLimit || page.Limit < 1 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid pagination limit value.",
|
||||
Detail: fmt.Sprintf("Pagination limit must be in range (0, %d]", maxListSessionsLimit),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
queryStr := r.URL.Query().Get("q")
|
||||
filter, errs := searchquery.AIBridgeSessions(ctx, api.Database, queryStr, page, apiKey.UserID, afterSessionID)
|
||||
if len(errs) > 0 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid session search query.",
|
||||
Validations: errs,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Validate the cursor session exists before running the main query.
|
||||
if afterSessionID != "" {
|
||||
//nolint:exhaustruct // Only need session_id filter and limit.
|
||||
cursor, err := api.Database.ListAIBridgeSessions(ctx, database.ListAIBridgeSessionsParams{
|
||||
SessionID: afterSessionID,
|
||||
Limit: 1,
|
||||
})
|
||||
if err != nil {
|
||||
api.Logger.Error(ctx, "error validating after_session_id cursor", slog.Error(err))
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error validating after_session_id cursor.",
|
||||
Detail: "", // Don't leak database issue to client.
|
||||
})
|
||||
return
|
||||
}
|
||||
if len(cursor) == 0 {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Query parameter has invalid value.",
|
||||
Detail: fmt.Sprintf("after_session_id: session %q not found", afterSessionID),
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
var (
|
||||
count int64
|
||||
rows []database.ListAIBridgeSessionsRow
|
||||
)
|
||||
err := api.Database.InTx(func(db database.Store) error {
|
||||
var err error
|
||||
count, err = db.CountAIBridgeSessions(ctx, database.CountAIBridgeSessionsParams{
|
||||
StartedAfter: filter.StartedAfter,
|
||||
StartedBefore: filter.StartedBefore,
|
||||
InitiatorID: filter.InitiatorID,
|
||||
Provider: filter.Provider,
|
||||
Model: filter.Model,
|
||||
Client: filter.Client,
|
||||
SessionID: filter.SessionID,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("count authorized aibridge sessions: %w", err)
|
||||
}
|
||||
|
||||
rows, err = db.ListAIBridgeSessions(ctx, filter)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("list aibridge sessions: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}, &database.TxOptions{
|
||||
Isolation: sql.LevelRepeatableRead, // Consistency across queries tables while writes may be occurring.
|
||||
ReadOnly: true,
|
||||
TxIdentifier: "aibridge_list_sessions",
|
||||
})
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error getting AI Bridge sessions.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
sessions := make([]codersdk.AIBridgeSession, len(rows))
|
||||
for i, row := range rows {
|
||||
sessions[i] = db2sdk.AIBridgeSession(row)
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, codersdk.AIBridgeListSessionsResponse{
|
||||
Count: count,
|
||||
Sessions: sessions,
|
||||
})
|
||||
}
|
||||
|
||||
// aiBridgeListModels returns all AI Bridge models a user can see.
|
||||
//
|
||||
// @Summary List AI Bridge models
|
||||
|
||||
@@ -665,6 +665,560 @@ func TestAIBridgeListInterceptions(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func aibridgeOpts(t *testing.T) *coderdenttest.Options {
|
||||
t.Helper()
|
||||
dv := coderdtest.DeploymentValues(t)
|
||||
dv.AI.BridgeConfig.Enabled = serpent.Bool(true)
|
||||
return &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
DeploymentValues: dv,
|
||||
},
|
||||
LicenseOptions: &coderdenttest.LicenseOptions{
|
||||
Features: license.Features{
|
||||
codersdk.FeatureAIBridge: 1,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestAIBridgeListSessions(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("EmptyDB", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, _ := coderdenttest.New(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
//nolint:gocritic // Owner role is irrelevant here.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, res.Sessions)
|
||||
require.EqualValues(t, 0, res.Count)
|
||||
})
|
||||
|
||||
t.Run("OK", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, db, firstUser := coderdenttest.NewWithDatabase(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
now := dbtime.Now()
|
||||
|
||||
// Session 1: Two interceptions sharing client_session_id "session-A".
|
||||
s1i1EndedAt := now.Add(time.Minute)
|
||||
s1i1 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4",
|
||||
StartedAt: now,
|
||||
Client: sql.NullString{String: "claude-code", Valid: true},
|
||||
ClientSessionID: sql.NullString{String: "session-A", Valid: true},
|
||||
}, &s1i1EndedAt)
|
||||
s1i2EndedAt := now.Add(2 * time.Minute)
|
||||
dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4-haiku",
|
||||
StartedAt: now.Add(time.Minute),
|
||||
Client: sql.NullString{String: "claude-code", Valid: true},
|
||||
ClientSessionID: sql.NullString{String: "session-A", Valid: true},
|
||||
ThreadRootInterceptionID: uuid.NullUUID{UUID: s1i1.ID, Valid: true},
|
||||
ThreadParentInterceptionID: uuid.NullUUID{UUID: s1i1.ID, Valid: true},
|
||||
}, &s1i2EndedAt)
|
||||
|
||||
// Add token usages to session 1 interceptions.
|
||||
dbgen.AIBridgeTokenUsage(t, db, database.InsertAIBridgeTokenUsageParams{
|
||||
InterceptionID: s1i1.ID,
|
||||
InputTokens: 100,
|
||||
OutputTokens: 50,
|
||||
CreatedAt: now,
|
||||
})
|
||||
dbgen.AIBridgeTokenUsage(t, db, database.InsertAIBridgeTokenUsageParams{
|
||||
InterceptionID: s1i1.ID,
|
||||
InputTokens: 200,
|
||||
OutputTokens: 75,
|
||||
CreatedAt: now.Add(time.Second),
|
||||
})
|
||||
|
||||
// Add user prompts to session 1.
|
||||
dbgen.AIBridgeUserPrompt(t, db, database.InsertAIBridgeUserPromptParams{
|
||||
InterceptionID: s1i1.ID,
|
||||
Prompt: "first prompt",
|
||||
CreatedAt: now,
|
||||
})
|
||||
dbgen.AIBridgeUserPrompt(t, db, database.InsertAIBridgeUserPromptParams{
|
||||
InterceptionID: s1i1.ID,
|
||||
Prompt: "last prompt in session",
|
||||
CreatedAt: now.Add(time.Minute),
|
||||
})
|
||||
|
||||
// Session 2: Thread-based session (no client_session_id, shared thread_root_id).
|
||||
s2i1EndedAt := now.Add(-time.Hour + time.Minute)
|
||||
s2i1 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4",
|
||||
StartedAt: now.Add(-time.Hour),
|
||||
}, &s2i1EndedAt)
|
||||
s2i2EndedAt := now.Add(-time.Hour + 2*time.Minute)
|
||||
dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4",
|
||||
StartedAt: now.Add(-time.Hour + time.Minute),
|
||||
ThreadRootInterceptionID: uuid.NullUUID{UUID: s2i1.ID, Valid: true},
|
||||
ThreadParentInterceptionID: uuid.NullUUID{UUID: s2i1.ID, Valid: true},
|
||||
}, &s2i2EndedAt)
|
||||
|
||||
// Session 3: Standalone interception (no client_session_id, no thread_root_id).
|
||||
s3EndedAt := now.Add(-2*time.Hour + time.Minute)
|
||||
s3i1 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4",
|
||||
StartedAt: now.Add(-2 * time.Hour),
|
||||
}, &s3EndedAt)
|
||||
|
||||
// Session 4: Two distinct thread roots in one client_session_id.
|
||||
s4i1EndedAt := now.Add(-3*time.Hour + time.Minute)
|
||||
dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4",
|
||||
StartedAt: now.Add(-3 * time.Hour),
|
||||
ClientSessionID: sql.NullString{String: "session-multi", Valid: true},
|
||||
}, &s4i1EndedAt)
|
||||
s4i2EndedAt := now.Add(-3*time.Hour + 2*time.Minute)
|
||||
dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4",
|
||||
StartedAt: now.Add(-3*time.Hour + time.Minute),
|
||||
ClientSessionID: sql.NullString{String: "session-multi", Valid: true},
|
||||
}, &s4i2EndedAt)
|
||||
|
||||
//nolint:gocritic // Owner role is irrelevant here.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 4, res.Count)
|
||||
require.Len(t, res.Sessions, 4)
|
||||
|
||||
// Sessions ordered by started_at DESC: session-A (now), then
|
||||
// thread-based (now-1h), then standalone (now-2h), then
|
||||
// multi-thread (now-3h).
|
||||
require.Equal(t, "session-A", res.Sessions[0].ID)
|
||||
require.Equal(t, s2i1.ID.String(), res.Sessions[1].ID)
|
||||
require.Equal(t, s3i1.ID.String(), res.Sessions[2].ID)
|
||||
require.Equal(t, "session-multi", res.Sessions[3].ID)
|
||||
|
||||
// Verify session 1 aggregations.
|
||||
s1 := res.Sessions[0]
|
||||
require.ElementsMatch(t, []string{"anthropic"}, s1.Providers)
|
||||
require.ElementsMatch(t, []string{"claude-4", "claude-4-haiku"}, s1.Models)
|
||||
require.NotNil(t, s1.Client)
|
||||
require.Equal(t, "claude-code", *s1.Client)
|
||||
require.EqualValues(t, 300, s1.TokenUsageSummary.InputTokens)
|
||||
require.EqualValues(t, 125, s1.TokenUsageSummary.OutputTokens)
|
||||
require.NotNil(t, s1.LastPrompt)
|
||||
require.Equal(t, "last prompt in session", *s1.LastPrompt)
|
||||
// Two interceptions in session-A, but they share a thread root,
|
||||
// so thread count is 1.
|
||||
require.EqualValues(t, 1, s1.Threads)
|
||||
|
||||
// Verify session 2 (thread-based).
|
||||
s2 := res.Sessions[1]
|
||||
require.ElementsMatch(t, []string{"openai"}, s2.Providers)
|
||||
// Thread count: the root interception and its child share the same
|
||||
// thread root, so count is 1.
|
||||
require.EqualValues(t, 1, s2.Threads)
|
||||
|
||||
// Verify session 3 (standalone).
|
||||
s3 := res.Sessions[2]
|
||||
require.EqualValues(t, 1, s3.Threads)
|
||||
require.Nil(t, s3.LastPrompt)
|
||||
|
||||
// Verify session 4 (multiple threads). Thread A has a root +
|
||||
// child (1 thread), thread B is a standalone root (1 thread),
|
||||
// so total is 2.
|
||||
s4 := res.Sessions[3]
|
||||
require.EqualValues(t, 2, s4.Threads)
|
||||
require.ElementsMatch(t, []string{"anthropic", "openai"}, s4.Providers)
|
||||
require.ElementsMatch(t, []string{"claude-4", "gpt-4"}, s4.Models)
|
||||
})
|
||||
|
||||
t.Run("Pagination", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, db, firstUser := coderdenttest.NewWithDatabase(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
now := dbtime.Now()
|
||||
// Create 5 standalone sessions with different start times.
|
||||
allSessionIDs := make([]string, 5)
|
||||
for i := range 5 {
|
||||
endedAt := now.Add(-time.Duration(i)*time.Hour + time.Minute)
|
||||
intc := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
StartedAt: now.Add(-time.Duration(i) * time.Hour),
|
||||
}, &endedAt)
|
||||
// Standalone session: ID = interception UUID string.
|
||||
allSessionIDs[i] = intc.ID.String()
|
||||
}
|
||||
|
||||
// Test offset pagination.
|
||||
//nolint:gocritic // Owner role is irrelevant here.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Pagination: codersdk.Pagination{Limit: 2},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, res.Sessions, 2)
|
||||
require.EqualValues(t, 5, res.Count)
|
||||
require.Equal(t, allSessionIDs[0], res.Sessions[0].ID)
|
||||
require.Equal(t, allSessionIDs[1], res.Sessions[1].ID)
|
||||
|
||||
// Second page with offset.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Pagination: codersdk.Pagination{Limit: 2, Offset: 2},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, res.Sessions, 2)
|
||||
require.Equal(t, allSessionIDs[2], res.Sessions[0].ID)
|
||||
require.Equal(t, allSessionIDs[3], res.Sessions[1].ID)
|
||||
|
||||
// Test cursor pagination.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Pagination: codersdk.Pagination{Limit: 2},
|
||||
AfterSessionID: allSessionIDs[1],
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, res.Sessions, 2)
|
||||
require.Equal(t, allSessionIDs[2], res.Sessions[0].ID)
|
||||
require.Equal(t, allSessionIDs[3], res.Sessions[1].ID)
|
||||
|
||||
// Test mutual exclusion of cursor and offset.
|
||||
_, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Pagination: codersdk.Pagination{Limit: 2, Offset: 1},
|
||||
AfterSessionID: allSessionIDs[0],
|
||||
})
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Contains(t, sdkErr.Detail, "Cannot use both after_session_id and offset pagination")
|
||||
})
|
||||
|
||||
t.Run("AfterSessionIDNotFound", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, _ := coderdenttest.New(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
//nolint:gocritic // Owner role is irrelevant here.
|
||||
_, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Pagination: codersdk.Pagination{Limit: 10},
|
||||
AfterSessionID: "nonexistent-session-id",
|
||||
})
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||
require.Equal(t, `after_session_id: session "nonexistent-session-id" not found`, sdkErr.Detail)
|
||||
})
|
||||
|
||||
t.Run("Filters", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, db, firstUser := coderdenttest.NewWithDatabase(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
_, user2 := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID)
|
||||
|
||||
now := dbtime.Now()
|
||||
|
||||
// Session from user1 with provider "anthropic" and client "claude-code".
|
||||
s1EndedAt := now.Add(time.Minute)
|
||||
s1 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4",
|
||||
StartedAt: now,
|
||||
Client: sql.NullString{String: "claude-code", Valid: true},
|
||||
}, &s1EndedAt)
|
||||
|
||||
// Session from user2 with provider "openai".
|
||||
s2EndedAt := now.Add(-time.Hour + time.Minute)
|
||||
s2 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: user2.ID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4",
|
||||
StartedAt: now.Add(-time.Hour),
|
||||
}, &s2EndedAt)
|
||||
|
||||
// Filter by initiator.
|
||||
//nolint:gocritic // Owner role is irrelevant; testing filter behavior.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Initiator: user2.Username,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Equal(t, s2.ID.String(), res.Sessions[0].ID)
|
||||
|
||||
// Filter by provider.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Provider: "anthropic",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Equal(t, s1.ID.String(), res.Sessions[0].ID)
|
||||
|
||||
// Filter by model.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Model: "gpt-4",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Equal(t, s2.ID.String(), res.Sessions[0].ID)
|
||||
|
||||
// Filter by client.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Client: "claude-code",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Equal(t, s1.ID.String(), res.Sessions[0].ID)
|
||||
|
||||
// Filter by time range.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
StartedAfter: now.Add(-30 * time.Minute),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Equal(t, s1.ID.String(), res.Sessions[0].ID)
|
||||
|
||||
// Filter by session_id.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
SessionID: s2.ID.String(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Len(t, res.Sessions, 1)
|
||||
require.Equal(t, s2.ID.String(), res.Sessions[0].ID)
|
||||
|
||||
// Filter by session_id with no match.
|
||||
res, err = client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
SessionID: "nonexistent-session-id",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 0, res.Count)
|
||||
require.Empty(t, res.Sessions)
|
||||
})
|
||||
|
||||
t.Run("Authorized", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
adminClient, db, firstUser := coderdenttest.NewWithDatabase(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
secondUserClient, secondUser := coderdtest.CreateAnotherUser(t, adminClient, firstUser.OrganizationID)
|
||||
|
||||
now := dbtime.Now()
|
||||
i1EndedAt := now.Add(time.Minute)
|
||||
i1 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
StartedAt: now,
|
||||
}, &i1EndedAt)
|
||||
i2 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: secondUser.ID,
|
||||
StartedAt: now.Add(-time.Hour),
|
||||
}, &now)
|
||||
|
||||
// Admin can see all sessions.
|
||||
//nolint:gocritic // Intentionally testing admin/owner visibility.
|
||||
res, err := adminClient.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 2, res.Count)
|
||||
require.Len(t, res.Sessions, 2)
|
||||
require.Equal(t, i1.ID.String(), res.Sessions[0].ID)
|
||||
require.Equal(t, i2.ID.String(), res.Sessions[1].ID)
|
||||
|
||||
// Second user can only see their own sessions.
|
||||
res, err = secondUserClient.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Len(t, res.Sessions, 1)
|
||||
require.Equal(t, i2.ID.String(), res.Sessions[0].ID)
|
||||
})
|
||||
|
||||
t.Run("SessionIDCollisionAcrossUsers", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, db, firstUser := coderdenttest.NewWithDatabase(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
_, user2 := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID)
|
||||
|
||||
now := dbtime.Now()
|
||||
|
||||
// Two users share the same client_session_id. They must be
|
||||
// treated as distinct sessions.
|
||||
sharedSessionID := "shared-session-id"
|
||||
u1EndedAt := now.Add(time.Minute)
|
||||
u1Interception := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
Provider: "anthropic",
|
||||
Model: "claude-4",
|
||||
StartedAt: now,
|
||||
Client: sql.NullString{String: "claude-code", Valid: true},
|
||||
ClientSessionID: sql.NullString{String: sharedSessionID, Valid: true},
|
||||
}, &u1EndedAt)
|
||||
dbgen.AIBridgeTokenUsage(t, db, database.InsertAIBridgeTokenUsageParams{
|
||||
InterceptionID: u1Interception.ID,
|
||||
InputTokens: 100,
|
||||
OutputTokens: 50,
|
||||
CreatedAt: now,
|
||||
})
|
||||
|
||||
u2EndedAt := now.Add(-time.Hour + time.Minute)
|
||||
u2Interception := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: user2.ID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4",
|
||||
StartedAt: now.Add(-time.Hour),
|
||||
Client: sql.NullString{String: "cursor", Valid: true},
|
||||
ClientSessionID: sql.NullString{String: sharedSessionID, Valid: true},
|
||||
}, &u2EndedAt)
|
||||
dbgen.AIBridgeTokenUsage(t, db, database.InsertAIBridgeTokenUsageParams{
|
||||
InterceptionID: u2Interception.ID,
|
||||
InputTokens: 200,
|
||||
OutputTokens: 75,
|
||||
CreatedAt: now.Add(-time.Hour),
|
||||
})
|
||||
|
||||
// Admin should see two distinct sessions despite the shared
|
||||
// session_id, each with the correct user and token counts.
|
||||
//nolint:gocritic // Owner role is irrelevant; testing collision behavior.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 2, res.Count)
|
||||
require.Len(t, res.Sessions, 2)
|
||||
|
||||
// Both sessions share the same ID string but belong to
|
||||
// different users.
|
||||
require.Equal(t, sharedSessionID, res.Sessions[0].ID)
|
||||
require.Equal(t, sharedSessionID, res.Sessions[1].ID)
|
||||
require.NotEqual(t, res.Sessions[0].Initiator.ID, res.Sessions[1].Initiator.ID)
|
||||
|
||||
// Verify token counts are not merged across users.
|
||||
for _, s := range res.Sessions {
|
||||
if s.Initiator.ID == firstUser.UserID {
|
||||
require.EqualValues(t, 100, s.TokenUsageSummary.InputTokens)
|
||||
require.EqualValues(t, 50, s.TokenUsageSummary.OutputTokens)
|
||||
} else {
|
||||
require.EqualValues(t, 200, s.TokenUsageSummary.InputTokens)
|
||||
require.EqualValues(t, 75, s.TokenUsageSummary.OutputTokens)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("InflightSessions", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, db, firstUser := coderdenttest.NewWithDatabase(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
now := dbtime.Now()
|
||||
i1EndedAt := now.Add(time.Minute)
|
||||
i1 := dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
StartedAt: now,
|
||||
}, &i1EndedAt)
|
||||
// Inflight interception (no ended_at) should not appear as a session.
|
||||
dbgen.AIBridgeInterception(t, db, database.InsertAIBridgeInterceptionParams{
|
||||
InitiatorID: firstUser.UserID,
|
||||
StartedAt: now.Add(-time.Hour),
|
||||
}, nil)
|
||||
|
||||
//nolint:gocritic // Owner role is irrelevant; testing inflight filtering.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, res.Count)
|
||||
require.Len(t, res.Sessions, 1)
|
||||
require.Equal(t, i1.ID.String(), res.Sessions[0].ID)
|
||||
})
|
||||
|
||||
t.Run("FilterErrors", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, _ := coderdenttest.New(t, aibridgeOpts(t))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
q string
|
||||
want []codersdk.ValidationError
|
||||
}{
|
||||
{
|
||||
name: "UnknownUsername",
|
||||
q: "initiator:unknown",
|
||||
want: []codersdk.ValidationError{
|
||||
{
|
||||
Field: "initiator",
|
||||
Detail: `Query param "initiator" has invalid value: user "unknown" either does not exist, or you are unauthorized to view them`,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "InvalidStartedAfter",
|
||||
q: "started_after:invalid",
|
||||
want: []codersdk.ValidationError{
|
||||
{
|
||||
Field: "started_after",
|
||||
Detail: `Query param "started_after" must be a valid date format (2006-01-02T15:04:05.999999999Z07:00): parsing time "INVALID" as "2006-01-02T15:04:05.999999999Z07:00": cannot parse "INVALID" as "2006"`,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "InvalidStartedBefore",
|
||||
q: "started_before:invalid",
|
||||
want: []codersdk.ValidationError{
|
||||
{
|
||||
Field: "started_before",
|
||||
Detail: `Query param "started_before" must be a valid date format (2006-01-02T15:04:05.999999999Z07:00): parsing time "INVALID" as "2006-01-02T15:04:05.999999999Z07:00": cannot parse "INVALID" as "2006"`,
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "InvalidBeforeAfterRange",
|
||||
q: `started_after:"2025-01-01T00:00:00Z" started_before:"2024-01-01T00:00:00Z"`,
|
||||
want: []codersdk.ValidationError{
|
||||
{
|
||||
Field: "started_before",
|
||||
Detail: `Query param "started_before" has invalid value: "started_before" must be after "started_after" if set`,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
FilterQuery: tc.q,
|
||||
})
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Equal(t, tc.want, sdkErr.Validations)
|
||||
require.Empty(t, res.Sessions)
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("PaginationLimitValidation", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
client, _ := coderdenttest.New(t, aibridgeOpts(t))
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
//nolint:gocritic // Owner role is irrelevant; testing pagination validation.
|
||||
res, err := client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
|
||||
Pagination: codersdk.Pagination{
|
||||
Limit: 1001,
|
||||
},
|
||||
})
|
||||
var sdkErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &sdkErr)
|
||||
require.Contains(t, sdkErr.Message, "Invalid pagination limit value.")
|
||||
require.Empty(t, res.Sessions)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAIBridgeRouting(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user