mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
refactor(coderd): stop storing chat gateway key IDs and drop the columns (#27171)
> Mux is working on behalf of Mike. ## Summary Stop reading and writing the legacy `api_key_id` columns on chat messages and queued messages, and drop the columns in the same PR. Runtime AI Gateway attribution continues to use the per-user synthetic key introduced by #27170. With the columns gone, `sqlc` generates `database.ChatMessage` and `database.ChatQueuedMessage` without `api_key_id`, so no transitional query scaffolding is needed. Migration `000548` drops the `api_key_id` columns. #27170 already removed their foreign keys, so the down migration re-adds nullable text columns without constraints. Previous column values cannot be restored. Also moves the model config validation in `CreateChat` above the message-building work so a disabled or invalid model fails fast. On main this mattered more: the old ordering minted a synthetic API key before rejecting the request. Deploy note: replicas still running the previous release write `api_key_id` on insert, so chat message inserts on old replicas fail during the rolling window after the column drop. This was previously split across two PRs to avoid that window; per review feedback the split added more churn than it was worth for an experimental surface. Depends on #27170 (merged).
This commit is contained in:
@@ -13,7 +13,7 @@ Chatd attributes AI Gateway requests with a synthetic API key owned by the chat
|
||||
|
||||
Synthetic keys expire after 30 days. When less than 24 hours remain, chatd extends the expiry of the existing row in place instead of replacing it, because an in-flight generation may have already delegated the current key ID to the gateway. The key ID is therefore stable for the lifetime of the user. Mints and extensions are serialized with a per-user advisory lock, since the partial unique index on token names only covers `login_type = 'token'` rows. The generated token is discarded, so the stored key cannot be used as a bearer credential, and it carries a minimal scope as defense in depth.
|
||||
|
||||
The legacy `api_key_id` columns on messages and queued messages are still stamped with the synthetic key for rolling deployment compatibility, but they no longer have foreign keys to `api_keys`. They are not the source of gateway routing. Stale IDs are harmless because chatd resolves attribution from `chats.owner_id`.
|
||||
Messages and queued messages no longer carry `api_key_id` columns; attribution is resolved solely from `chats.owner_id`. The drop migration discards any IDs stamped by older replicas, and its rollback restores the columns as nullable without backfilling them.
|
||||
|
||||
Deleting a synthetic key (password reset, explicit key deletion, dbpurge of long-expired keys) does not touch chat messages, queued messages, or their version fields. Chatd mints a replacement on the next request without mutating history. User suspension and deletion still block delegated gateway authorization.
|
||||
|
||||
|
||||
+9
-102
@@ -1286,9 +1286,10 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
return database.Chat{}, limitErr
|
||||
}
|
||||
|
||||
apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, opts.OwnerID)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("ensure synthetic API key: %w", err)
|
||||
if opts.ModelConfigID != uuid.Nil {
|
||||
if err := requireEnabledChatModelConfig(ctx, p.db, opts.ModelConfigID); err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
}
|
||||
|
||||
labelsJSON, err := json.Marshal(opts.Labels)
|
||||
@@ -1332,13 +1333,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
initialMessages = append(initialMessages, systemMessage(userPromptContent, opts.ModelConfigID))
|
||||
}
|
||||
initialMessages = append(initialMessages, systemMessage(workspaceAwarenessContent, opts.ModelConfigID))
|
||||
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, opts.ModelConfigID, opts.OwnerID, apiKeyID, opts.ReasoningEffort))
|
||||
|
||||
if opts.ModelConfigID != uuid.Nil {
|
||||
if err := requireEnabledChatModelConfig(ctx, p.db, opts.ModelConfigID); err != nil {
|
||||
return database.Chat{}, err
|
||||
}
|
||||
}
|
||||
initialMessages = append(initialMessages, userMessage(userContent, opts.ModelConfigID, opts.OwnerID, opts.ReasoningEffort))
|
||||
|
||||
result, err := chatstate.CreateChat(ctx, p.db, p.pubsub, chatstate.CreateChatInput{
|
||||
OrganizationID: opts.OrganizationID,
|
||||
@@ -1418,15 +1413,6 @@ func (p *Server) SendMessage(
|
||||
requestedPlanMode := opts.PlanMode
|
||||
requestedMCPServerIDs := opts.MCPServerIDs
|
||||
|
||||
chat, err := p.db.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, xerrors.Errorf("load chat: %w", err)
|
||||
}
|
||||
apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, chat.OwnerID)
|
||||
if err != nil {
|
||||
return SendMessageResult{}, xerrors.Errorf("ensure synthetic API key: %w", err)
|
||||
}
|
||||
|
||||
var result SendMessageResult
|
||||
machine := p.newChatMachine(opts.ChatID)
|
||||
updateErr := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
|
||||
@@ -1491,7 +1477,7 @@ func (p *Server) SendMessage(
|
||||
// Queue capacity is enforced inside tx.SendMessage; this
|
||||
// wrapper only propagates the typed error.
|
||||
sendResult, err := tx.SendMessage(chatstate.SendMessageInput{
|
||||
Message: userMessageWithAPIKeyID(content, modelConfigID, messageCreatedBy, apiKeyID, opts.ReasoningEffort),
|
||||
Message: userMessage(content, modelConfigID, messageCreatedBy, opts.ReasoningEffort),
|
||||
BusyBehavior: busyBehaviorToChatState(busyBehavior),
|
||||
})
|
||||
if err != nil {
|
||||
@@ -1665,15 +1651,6 @@ func (p *Server) EditMessage(
|
||||
if err != nil {
|
||||
return EditMessageResult{}, xerrors.Errorf("marshal message content: %w", err)
|
||||
}
|
||||
chat, err := p.db.GetChatByID(ctx, opts.ChatID)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, xerrors.Errorf("load chat: %w", err)
|
||||
}
|
||||
apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, chat.OwnerID)
|
||||
if err != nil {
|
||||
return EditMessageResult{}, xerrors.Errorf("ensure synthetic API key: %w", err)
|
||||
}
|
||||
|
||||
var (
|
||||
result EditMessageResult
|
||||
editedMsg database.ChatMessage
|
||||
@@ -1744,7 +1721,6 @@ func (p *Server) EditMessage(
|
||||
Content: content,
|
||||
ModelConfigIDOverride: modelOverride,
|
||||
ReasoningEffortOverride: reasoningEffortOverride,
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: true},
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, chatstate.ErrEditedMessageNotUser) {
|
||||
@@ -2759,16 +2735,6 @@ type chatMessage struct {
|
||||
runtimeMs int64
|
||||
}
|
||||
|
||||
type userChatMessage struct {
|
||||
chatMessage
|
||||
apiKeyID string
|
||||
}
|
||||
|
||||
func (m userChatMessage) withCreatedBy(id uuid.UUID) userChatMessage {
|
||||
m.chatMessage = m.chatMessage.withCreatedBy(id)
|
||||
return m
|
||||
}
|
||||
|
||||
func newChatMessage(
|
||||
role database.ChatMessageRole,
|
||||
content pqtype.NullRawMessage,
|
||||
@@ -2785,25 +2751,6 @@ func newChatMessage(
|
||||
}
|
||||
}
|
||||
|
||||
func newUserChatMessage(
|
||||
apiKeyID string,
|
||||
content pqtype.NullRawMessage,
|
||||
visibility database.ChatMessageVisibility,
|
||||
modelConfigID uuid.UUID,
|
||||
contentVersion int16,
|
||||
) userChatMessage {
|
||||
return userChatMessage{
|
||||
chatMessage: newChatMessage(
|
||||
database.ChatMessageRoleUser,
|
||||
content,
|
||||
visibility,
|
||||
modelConfigID,
|
||||
contentVersion,
|
||||
),
|
||||
apiKeyID: apiKeyID,
|
||||
}
|
||||
}
|
||||
|
||||
func (m chatMessage) withCreatedBy(id uuid.UUID) chatMessage {
|
||||
m.createdBy = id
|
||||
return m
|
||||
@@ -2812,10 +2759,8 @@ func (m chatMessage) withCreatedBy(id uuid.UUID) chatMessage {
|
||||
func appendMessageFields(
|
||||
params *database.InsertChatMessagesParams,
|
||||
msg chatMessage,
|
||||
apiKeyID string,
|
||||
) {
|
||||
params.CreatedBy = append(params.CreatedBy, msg.createdBy)
|
||||
params.APIKeyID = append(params.APIKeyID, apiKeyID)
|
||||
params.ModelConfigID = append(params.ModelConfigID, msg.modelConfigID)
|
||||
params.ReasoningEffort = append(params.ReasoningEffort, "")
|
||||
params.Role = append(params.Role, msg.role)
|
||||
@@ -2834,21 +2779,7 @@ func appendMessageFields(
|
||||
params.RuntimeMs = append(params.RuntimeMs, msg.runtimeMs)
|
||||
}
|
||||
|
||||
func appendChatMessage(params *database.InsertChatMessagesParams, msg chatMessage) {
|
||||
if msg.role == database.ChatMessageRoleUser {
|
||||
panic("developer error: use appendUserChatMessage for user-role messages")
|
||||
}
|
||||
appendMessageFields(params, msg, "")
|
||||
}
|
||||
|
||||
func appendUserChatMessage(params *database.InsertChatMessagesParams, msg userChatMessage) {
|
||||
appendMessageFields(params, msg.chatMessage, msg.apiKeyID)
|
||||
}
|
||||
|
||||
// BuildSingleUserChatMessageInsertParams creates batch insert params for
|
||||
// one user message, requiring an apiKeyID for AI Gateway attribution.
|
||||
// BuildSingleChatMessageInsertParams creates batch insert params for one
|
||||
// non-user message using the shared chat message builder.
|
||||
// BuildSingleChatMessageInsertParams builds insert parameters for one chat message.
|
||||
func BuildSingleChatMessageInsertParams(
|
||||
chatID uuid.UUID,
|
||||
role database.ChatMessageRole,
|
||||
@@ -2858,38 +2789,14 @@ func BuildSingleChatMessageInsertParams(
|
||||
contentVersion int16,
|
||||
createdBy uuid.UUID,
|
||||
) database.InsertChatMessagesParams {
|
||||
params := database.InsertChatMessagesParams{ //nolint:exhaustruct // Fields populated by appendChatMessage.
|
||||
params := database.InsertChatMessagesParams{ //nolint:exhaustruct // Fields populated by appendMessageFields.
|
||||
ChatID: chatID,
|
||||
}
|
||||
msg := newChatMessage(role, content, visibility, modelConfigID, contentVersion)
|
||||
if createdBy != uuid.Nil {
|
||||
msg = msg.withCreatedBy(createdBy)
|
||||
}
|
||||
if role == database.ChatMessageRoleUser {
|
||||
appendMessageFields(¶ms, msg, "")
|
||||
} else {
|
||||
appendChatMessage(¶ms, msg)
|
||||
}
|
||||
return params
|
||||
}
|
||||
|
||||
func BuildSingleUserChatMessageInsertParams(
|
||||
chatID uuid.UUID,
|
||||
apiKeyID string,
|
||||
content pqtype.NullRawMessage,
|
||||
visibility database.ChatMessageVisibility,
|
||||
modelConfigID uuid.UUID,
|
||||
contentVersion int16,
|
||||
createdBy uuid.UUID,
|
||||
) database.InsertChatMessagesParams {
|
||||
params := database.InsertChatMessagesParams{ //nolint:exhaustruct // Fields populated by appendUserChatMessage.
|
||||
ChatID: chatID,
|
||||
}
|
||||
msg := newUserChatMessage(apiKeyID, content, visibility, modelConfigID, contentVersion)
|
||||
if createdBy != uuid.Nil {
|
||||
msg = msg.withCreatedBy(createdBy)
|
||||
}
|
||||
appendUserChatMessage(¶ms, msg)
|
||||
appendMessageFields(¶ms, msg)
|
||||
return params
|
||||
}
|
||||
|
||||
|
||||
@@ -743,11 +743,6 @@ func TestRenameChatTitle(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func withChatMessageAPIKeyID(message database.ChatMessage, apiKeyID string) database.ChatMessage {
|
||||
message.APIKeyID = sqlNullString(apiKeyID)
|
||||
return message
|
||||
}
|
||||
|
||||
// requireOutgoingRequestModel asserts that the outgoing request body
|
||||
// requests wantModel. This is so that mock transports can still
|
||||
// verify the outgoing request asked for the expected model.
|
||||
@@ -867,12 +862,12 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
|
||||
LimitVal: manualTitleMessageWindowLimit,
|
||||
},
|
||||
).Return([]database.ChatMessage{
|
||||
withChatMessageAPIKeyID(mustChatMessage(
|
||||
mustChatMessage(
|
||||
t,
|
||||
database.ChatMessageRoleUser,
|
||||
database.ChatMessageVisibilityBoth,
|
||||
codersdk.ChatMessageText(userPrompt),
|
||||
), activeAPIKeyID),
|
||||
),
|
||||
mustChatMessage(
|
||||
t,
|
||||
database.ChatMessageRoleAssistant,
|
||||
@@ -1019,12 +1014,12 @@ func TestRegenerateChatTitle_SkipsPersistWhenTitleChangedConcurrently(t *testing
|
||||
LimitVal: manualTitleMessageWindowLimit,
|
||||
},
|
||||
).Return([]database.ChatMessage{
|
||||
withChatMessageAPIKeyID(mustChatMessage(
|
||||
mustChatMessage(
|
||||
t,
|
||||
database.ChatMessageRoleUser,
|
||||
database.ChatMessageVisibilityBoth,
|
||||
codersdk.ChatMessageText(userPrompt),
|
||||
), activeAPIKeyID),
|
||||
),
|
||||
}, nil)
|
||||
db.EXPECT().GetChatMessagesByChatIDDescPaginated(
|
||||
gomock.Any(),
|
||||
|
||||
+30
-230
@@ -74,12 +74,6 @@ type recordedOpenAIRequest struct {
|
||||
ContentLength int64
|
||||
}
|
||||
|
||||
func testAPIKeyID(t testing.TB, db database.Store, userID uuid.UUID) string {
|
||||
t.Helper()
|
||||
key, _ := dbgen.APIKey(t, db, database.APIKey{ID: uuid.NewString(), UserID: userID})
|
||||
return key.ID
|
||||
}
|
||||
|
||||
func chatAIGatewayTransportFactoryPointer(factory aibridge.TransportFactory) *atomic.Pointer[aibridge.TransportFactory] {
|
||||
var factoryPtr atomic.Pointer[aibridge.TransportFactory]
|
||||
factoryPtr.Store(&factory)
|
||||
@@ -770,7 +764,6 @@ func TestExploreChatUsesPersistedMCPSnapshot(t *testing.T) {
|
||||
codersdk.ChatMessageText("inspect the codebase"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
createdExplore, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
@@ -794,7 +787,6 @@ func TestExploreChatUsesPersistedMCPSnapshot(t *testing.T) {
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: webSearchModel.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
@@ -1603,175 +1595,6 @@ func TestUpdateChatHeartbeatsRequiresOwnership(t *testing.T) {
|
||||
require.Equal(t, chat.ID, ids[0])
|
||||
}
|
||||
|
||||
func TestCreateChatPersistsSyntheticAPIKeyIDOnInitialUserMessage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "create-chat-synthetic-api-key-id",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
require.Equal(t, database.ChatMessageRoleUser, messages[0].Role)
|
||||
require.True(t, messages[0].APIKeyID.Valid)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, gatewayKey.ID, messages[0].APIKeyID.String)
|
||||
}
|
||||
|
||||
func TestSendMessagePersistsSyntheticAPIKeyIDOnUserMessage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
Title: "send-message-synthetic-api-key-id",
|
||||
})
|
||||
|
||||
result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
Content: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("message with synthetic api key id"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, result.Queued)
|
||||
require.True(t, result.Message.APIKeyID.Valid)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, gatewayKey.ID, result.Message.APIKeyID.String)
|
||||
|
||||
stored, err := db.GetChatMessageByID(ctx, result.Message.ID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, stored.APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, stored.APIKeyID.String)
|
||||
}
|
||||
|
||||
func TestSendMessagePersistsSyntheticAPIKeyIDOnQueuedUserMessage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "queue-synthetic-api-key-id",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
|
||||
ID: chat.ID,
|
||||
Status: database.ChatStatusRunning,
|
||||
WorkerID: uuid.NullUUID{UUID: uuid.New(), Valid: true},
|
||||
StartedAt: sql.NullTime{Time: time.Now(), Valid: true},
|
||||
HeartbeatAt: sql.NullTime{Time: time.Now(), Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")},
|
||||
BusyBehavior: chatd.SendMessageBusyBehaviorQueue,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.Queued)
|
||||
require.NotNil(t, result.QueuedMessage)
|
||||
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.QueuedMessage.APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, result.QueuedMessage.APIKeyID.String)
|
||||
|
||||
queued, err := db.GetChatQueuedMessages(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, queued, 1)
|
||||
require.True(t, queued[0].APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, queued[0].APIKeyID.String)
|
||||
}
|
||||
|
||||
func TestEditMessagePersistsSyntheticAPIKeyIDOnReplacement(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: "edit-synthetic-api-key-id",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("original")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
|
||||
result, err := replica.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
EditedMessageID: messages[0].ID,
|
||||
CreatedBy: user.ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.Message.APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, result.Message.APIKeyID.String)
|
||||
|
||||
stored, err := db.GetChatMessageByID(ctx, result.Message.ID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, stored.APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, stored.APIKeyID.String)
|
||||
}
|
||||
|
||||
func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -2773,7 +2596,6 @@ func TestRecoverStaleRequiresActionChat(t *testing.T) {
|
||||
codersdk.ChatMessageText("hello"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
@@ -2789,7 +2611,6 @@ func TestRecoverStaleRequiresActionChat(t *testing.T) {
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
@@ -2867,7 +2688,6 @@ func TestNewReplicaRecoversStaleChatFromDeadReplica(t *testing.T) {
|
||||
codersdk.ChatMessageText("hello"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
@@ -2882,7 +2702,6 @@ func TestNewReplicaRecoversStaleChatFromDeadReplica(t *testing.T) {
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
@@ -5459,11 +5278,6 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) {
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
contextContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
|
||||
Type: codersdk.ChatMessagePartTypeContextFile,
|
||||
ContextFileAgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true},
|
||||
@@ -5473,9 +5287,9 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) {
|
||||
ContextFileDirectory: "/home/coder/project",
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
_, err = db.InsertChatMessages(ctx, chatd.BuildSingleUserChatMessageInsertParams(
|
||||
_, err = db.InsertChatMessages(ctx, chatd.BuildSingleChatMessageInsertParams(
|
||||
chat.ID,
|
||||
gatewayKey.ID,
|
||||
database.ChatMessageRoleUser,
|
||||
contextContent,
|
||||
database.ChatMessageVisibilityBoth,
|
||||
model.ID,
|
||||
@@ -5497,14 +5311,17 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) {
|
||||
chatResult := waitForTerminalChat(ctx, t, db, chat.ID)
|
||||
require.Equal(t, database.ChatStatusWaiting, chatResult.Status)
|
||||
require.False(t, chatResult.LastError.Valid)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
messages := chatMessages(ctx, t, db, chat.ID)
|
||||
promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
compressed := compressedChatSummarizedMessages(t, append(promptMessages, messages...))
|
||||
require.Len(t, compressed.summaries, 1)
|
||||
require.True(t, compressed.summaries[0].APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, compressed.summaries[0].APIKeyID.String)
|
||||
|
||||
requests := factory.RequestsSnapshot()
|
||||
require.NotEmpty(t, requests)
|
||||
@@ -6779,7 +6596,6 @@ func userMessageForTest(
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelID, Valid: true},
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: apiKeyID != ""},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8357,29 +8173,15 @@ func insertChatMessageParts(
|
||||
t.Helper()
|
||||
content, err := chatprompt.MarshalParts(parts)
|
||||
require.NoError(t, err)
|
||||
var params database.InsertChatMessagesParams
|
||||
if role == database.ChatMessageRoleUser {
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: createdBy})
|
||||
params = chatd.BuildSingleUserChatMessageInsertParams(
|
||||
chatID,
|
||||
apiKey.ID,
|
||||
content,
|
||||
database.ChatMessageVisibilityBoth,
|
||||
modelID,
|
||||
chatprompt.CurrentContentVersion,
|
||||
createdBy,
|
||||
)
|
||||
} else {
|
||||
params = chatd.BuildSingleChatMessageInsertParams(
|
||||
chatID,
|
||||
role,
|
||||
content,
|
||||
database.ChatMessageVisibilityBoth,
|
||||
modelID,
|
||||
chatprompt.CurrentContentVersion,
|
||||
createdBy,
|
||||
)
|
||||
}
|
||||
params := chatd.BuildSingleChatMessageInsertParams(
|
||||
chatID,
|
||||
role,
|
||||
content,
|
||||
database.ChatMessageVisibilityBoth,
|
||||
modelID,
|
||||
chatprompt.CurrentContentVersion,
|
||||
createdBy,
|
||||
)
|
||||
messages, err := db.InsertChatMessages(ctx, params)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
@@ -10211,12 +10013,11 @@ func seedAIGatewayOpenAITestDependencies(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
openAIURL string,
|
||||
) (database.User, database.Organization, database.AIProvider, database.ChatModelConfig, database.APIKey) {
|
||||
) (database.User, database.Organization, database.AIProvider, database.ChatModelConfig) {
|
||||
t.Helper()
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
@@ -10239,7 +10040,7 @@ func seedAIGatewayOpenAITestDependencies(
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
return user, org, provider, model, apiKey
|
||||
return user, org, provider, model
|
||||
}
|
||||
|
||||
func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) {
|
||||
@@ -10258,7 +10059,7 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) {
|
||||
})
|
||||
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
|
||||
|
||||
user, org, provider, model, _ := seedAIGatewayOpenAITestDependencies(t, db, openAIURL)
|
||||
user, org, provider, model := seedAIGatewayOpenAITestDependencies(t, db, openAIURL)
|
||||
|
||||
creator := newTestServer(t, db, ps, uuid.New())
|
||||
chat, err := creator.CreateChat(ctx, chatd.CreateOptions{
|
||||
@@ -10271,11 +10072,6 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) {
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, events, cancel, ok := creator.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
@@ -10292,6 +10088,11 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) {
|
||||
chatResult := waitForTerminalChat(ctx, t, db, chat.ID)
|
||||
require.Equal(t, database.ChatStatusWaiting, chatResult.Status)
|
||||
require.False(t, chatResult.LastError.Valid)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
requests := factory.RequestsSnapshot()
|
||||
require.NotEmpty(t, requests)
|
||||
@@ -10324,7 +10125,7 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) {
|
||||
return chattest.OpenAINonStreamingResponse(`{"title":"AI Gateway Workspace"}`)
|
||||
})
|
||||
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
|
||||
user, org, provider, model, _ := seedAIGatewayOpenAITestDependencies(t, db, openAIURL)
|
||||
user, org, provider, model := seedAIGatewayOpenAITestDependencies(t, db, openAIURL)
|
||||
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
|
||||
|
||||
creator := newTestServer(t, db, ps, uuid.New())
|
||||
@@ -10339,11 +10140,6 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) {
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
const contextText = "# Project instructions\nAlways keep routing metadata."
|
||||
// Workspace context is sourced from the agent's pinned snapshot. Seed it so
|
||||
@@ -10374,6 +10170,11 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) {
|
||||
pinned, err := db.ListChatContextResourcesByChatID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, pinned, "workspace context should be pinned to the chat")
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: chatd.GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
requests := factory.RequestsSnapshot()
|
||||
require.NotEmpty(t, requests)
|
||||
@@ -12189,7 +11990,6 @@ func TestPromoteQueuedPreservesReasoningEffort(t *testing.T) {
|
||||
Content: content.RawMessage,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
ReasoningEffort: database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffortHigh, Valid: true},
|
||||
APIKeyID: sql.NullString{String: testAPIKeyID(t, db, user.ID), Valid: true},
|
||||
CreatedBy: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -2,7 +2,6 @@ package chatstate_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"slices"
|
||||
"sync"
|
||||
@@ -34,13 +33,6 @@ type testFixture struct {
|
||||
User database.User
|
||||
Org database.Organization
|
||||
Model database.ChatModelConfig
|
||||
APIKey database.APIKey
|
||||
}
|
||||
|
||||
// apiKeyID returns the fixture API key wrapped for the chatstate
|
||||
// inputs that require a non-null api_key_id (for example EditMessage).
|
||||
func (f *testFixture) apiKeyID() sql.NullString {
|
||||
return sql.NullString{String: f.APIKey.ID, Valid: true}
|
||||
}
|
||||
|
||||
func newTestFixture(t *testing.T) *testFixture {
|
||||
@@ -60,7 +52,6 @@ func newTestFixture(t *testing.T) *testFixture {
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
IsDefault: true,
|
||||
})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
pub := newRecordingPubsub()
|
||||
return &testFixture{
|
||||
DB: db,
|
||||
@@ -69,7 +60,6 @@ func newTestFixture(t *testing.T) *testFixture {
|
||||
User: user,
|
||||
Org: org,
|
||||
Model: model,
|
||||
APIKey: apiKey,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -35,7 +35,6 @@ type Message struct {
|
||||
ContextLimit sql.NullInt64
|
||||
TotalCostMicros sql.NullInt64
|
||||
RuntimeMs sql.NullInt64
|
||||
APIKeyID sql.NullString
|
||||
}
|
||||
|
||||
// toInsertParams converts a batch of Messages into the parallel-array
|
||||
@@ -51,7 +50,6 @@ func toInsertParams(chatID uuid.UUID, messages []Message) database.InsertChatMes
|
||||
CreatedBy: make([]uuid.UUID, n),
|
||||
ModelConfigID: make([]uuid.UUID, n),
|
||||
ReasoningEffort: make([]string, n),
|
||||
APIKeyID: make([]string, n),
|
||||
Role: make([]database.ChatMessageRole, n),
|
||||
Content: make([]string, n),
|
||||
ContentVersion: make([]int16, n),
|
||||
@@ -73,9 +71,6 @@ func toInsertParams(chatID uuid.UUID, messages []Message) database.InsertChatMes
|
||||
if m.ReasoningEffort.Valid {
|
||||
params.ReasoningEffort[i] = string(m.ReasoningEffort.ChatReasoningEffort)
|
||||
}
|
||||
if m.APIKeyID.Valid {
|
||||
params.APIKeyID[i] = m.APIKeyID.String
|
||||
}
|
||||
params.Role[i] = m.Role
|
||||
if m.Content.Valid {
|
||||
params.Content[i] = string(m.Content.RawMessage)
|
||||
|
||||
@@ -249,7 +249,6 @@ func testEditMessageSynthesizesToolCancellationsBeforeReplacement(t *testing.T)
|
||||
MessageID: secondUserID,
|
||||
CreatedBy: f.User.ID,
|
||||
Content: editedContent,
|
||||
APIKeyID: f.apiKeyID(),
|
||||
})
|
||||
return err
|
||||
}))
|
||||
|
||||
@@ -236,7 +236,6 @@ func (tx *Tx) insertQueuedMessage(ownerFallback uuid.UUID, m Message) (database.
|
||||
ModelConfigID: m.ModelConfigID,
|
||||
ReasoningEffort: m.ReasoningEffort,
|
||||
CreatedBy: createdBy,
|
||||
APIKeyID: m.APIKeyID,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -251,7 +250,6 @@ func messageFromQueuedRow(q database.ChatQueuedMessage) Message {
|
||||
ReasoningEffort: q.ReasoningEffort,
|
||||
CreatedBy: uuid.NullUUID{UUID: q.CreatedBy, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: q.APIKeyID,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -491,7 +489,6 @@ type EditMessageInput struct {
|
||||
Content pqtype.NullRawMessage
|
||||
ModelConfigIDOverride uuid.NullUUID
|
||||
ReasoningEffortOverride database.NullChatReasoningEffort
|
||||
APIKeyID sql.NullString
|
||||
}
|
||||
|
||||
// EditMessageResult is returned by [Tx.EditMessage].
|
||||
@@ -571,10 +568,6 @@ func (tx *Tx) EditMessage(input EditMessageInput) (EditMessageResult, error) {
|
||||
if input.ReasoningEffortOverride.Valid {
|
||||
reasoningEffort = input.ReasoningEffortOverride
|
||||
}
|
||||
apiKeyID := input.APIKeyID
|
||||
if !apiKeyID.Valid {
|
||||
return EditMessageResult{}, xerrors.Errorf("api_key_id is required")
|
||||
}
|
||||
replacement := Message{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: input.Content,
|
||||
@@ -583,7 +576,6 @@ func (tx *Tx) EditMessage(input EditMessageInput) (EditMessageResult, error) {
|
||||
ReasoningEffort: reasoningEffort,
|
||||
CreatedBy: uuid.NullUUID{UUID: input.CreatedBy, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: apiKeyID,
|
||||
}
|
||||
insertedReplacement, err := tx.insertMessages([]Message{replacement})
|
||||
if err != nil {
|
||||
|
||||
@@ -794,9 +794,14 @@ func assertChatMessageText(t *testing.T, msg database.ChatMessage, want string)
|
||||
// matrix cases that need to verify the body inserted into
|
||||
// chat_queued_messages via SendMessage.
|
||||
func assertQueuedMessageText(t *testing.T, queued database.ChatQueuedMessage, want string) {
|
||||
t.Helper()
|
||||
assertQueuedMessageContent(t, queued.Content, want)
|
||||
}
|
||||
|
||||
func assertQueuedMessageContent(t *testing.T, content json.RawMessage, want string) {
|
||||
t.Helper()
|
||||
var parts []codersdk.ChatMessagePart
|
||||
require.NoError(t, json.Unmarshal(queued.Content, &parts), "unmarshal queued content")
|
||||
require.NoError(t, json.Unmarshal(content, &parts), "unmarshal queued content")
|
||||
require.Len(t, parts, 1, "expected exactly one queued content part")
|
||||
require.Equal(t, codersdk.ChatMessagePartTypeText, parts[0].Type,
|
||||
"expected a text content part")
|
||||
@@ -814,7 +819,7 @@ func assertQueueBodiesInOrder(ctx context.Context, t *testing.T, f *testFixture,
|
||||
require.NoError(t, err)
|
||||
require.Len(t, rows, len(want), "queue length must match expected bodies")
|
||||
for i, r := range rows {
|
||||
assertQueuedMessageText(t, r, want[i])
|
||||
assertQueuedMessageContent(t, r.Content, want[i])
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -136,7 +136,6 @@ func applyEditMessage(t *testing.T, f *testFixture, tx *chatstate.Tx, seeded see
|
||||
MessageID: seeded.initialUserMessageID,
|
||||
CreatedBy: f.User.ID,
|
||||
Content: content,
|
||||
APIKeyID: f.apiKeyID(),
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
|
||||
@@ -29,7 +27,7 @@ func systemMessage(rawContent pqtype.NullRawMessage, modelConfigID uuid.UUID) ch
|
||||
}
|
||||
}
|
||||
|
||||
func userMessageWithAPIKeyID(rawContent pqtype.NullRawMessage, modelConfigID, createdBy uuid.UUID, apiKeyID string, reasoningEffort *string) chatstate.Message {
|
||||
func userMessage(rawContent pqtype.NullRawMessage, modelConfigID, createdBy uuid.UUID, reasoningEffort *string) chatstate.Message {
|
||||
var effort database.NullChatReasoningEffort
|
||||
if reasoningEffort != nil && *reasoningEffort != "" {
|
||||
effort = database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffort(*reasoningEffort), Valid: true}
|
||||
@@ -42,7 +40,6 @@ func userMessageWithAPIKeyID(rawContent pqtype.NullRawMessage, modelConfigID, cr
|
||||
ReasoningEffort: effort,
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: createdBy != uuid.Nil},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: apiKeyID != ""},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -740,7 +740,6 @@ func (s *taskStarter) generateCompaction(
|
||||
}
|
||||
messages, err := buildCompactionMessages(buildCompactionMessagesInput{
|
||||
modelConfigID: prepared.ModelConfigID,
|
||||
activeAPIKeyID: prepared.ModelBuildOptions.ActiveAPIKeyID,
|
||||
toolCallID: compactionOpts.ToolCallID,
|
||||
toolName: compactionOpts.ToolName,
|
||||
compaction: compactionOutcome(outcome),
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package chatd //nolint:testpackage // Exercises unexported re-derivation helpers.
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
@@ -93,7 +92,6 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) {
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := chatdTestContext(t)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
@@ -135,7 +133,6 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) {
|
||||
},
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
@@ -161,6 +158,74 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) {
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortMedium, *providerOptions.ReasoningEffort)
|
||||
}
|
||||
|
||||
func TestPrepareGenerationSubagentUsesOwnerSyntheticAPIKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := chatdTestContext(t)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
provider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
}, "test-key")
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Model: "gpt-4o-mini",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
}, func(p *database.InsertChatModelConfigParams) {
|
||||
p.Enabled = true
|
||||
})
|
||||
parent := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
})
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true},
|
||||
RootChatID: uuid.NullUUID{UUID: parent.ID, Valid: true},
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "subagent attribution",
|
||||
ClientType: database.ChatClientTypeApi,
|
||||
InitialMessages: []chatstate.Message{
|
||||
{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: mustMarshalText(t, "inspect the workspace"),
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
server := newInternalTestServer(
|
||||
t,
|
||||
db,
|
||||
ps,
|
||||
chatprovider.ProviderAPIKeys{},
|
||||
withInternalTestServerTransportFactory(&aibridgeTestFactory{}),
|
||||
)
|
||||
prepared, err := server.prepareGeneration(ctx, generationPrepareInput{
|
||||
Chat: created.Chat,
|
||||
Messages: created.InitialMessages,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(prepared.Cleanup)
|
||||
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, gatewayKey.ID, prepared.ModelBuildOptions.ActiveAPIKeyID)
|
||||
}
|
||||
|
||||
// TestDeriveFinalTurnRunResult exercises the re-derivation path that replaces
|
||||
// the old in-memory generationSideEffects stash. The server here never ran
|
||||
// prepareGeneration, so a passing test proves the finish-turn inputs are
|
||||
@@ -196,7 +261,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
|
||||
p.Enabled = true
|
||||
p.IsDefault = true
|
||||
})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
@@ -212,7 +276,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
@@ -314,7 +377,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
|
||||
DisplayName: "gpt-4o-mini",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
@@ -330,7 +392,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
@@ -202,7 +202,6 @@ func userTextMessage(t *testing.T, text string, createdBy uuid.UUID, modelConfig
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: apiKeyID != ""},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -255,7 +255,6 @@ func textFromParts(parts []codersdk.ChatMessagePart) string {
|
||||
|
||||
type buildCompactionMessagesInput struct {
|
||||
modelConfigID uuid.UUID
|
||||
activeAPIKeyID string
|
||||
toolCallID string
|
||||
toolName string
|
||||
compaction compactionOutcome
|
||||
@@ -319,7 +318,6 @@ func buildCompactionMessages(input buildCompactionMessagesInput) (compactionMess
|
||||
Visibility: database.ChatMessageVisibilityModel,
|
||||
ModelConfigID: uuid.NullUUID{UUID: input.modelConfigID, Valid: input.modelConfigID != uuid.Nil},
|
||||
ContentVersion: contentVersion,
|
||||
APIKeyID: sql.NullString{String: input.activeAPIKeyID, Valid: input.activeAPIKeyID != ""},
|
||||
},
|
||||
baseMessage(database.ChatMessageRoleAssistant, database.ChatMessageVisibilityUser, input.modelConfigID, contentVersion, assistantContent),
|
||||
baseMessage(database.ChatMessageRoleTool, database.ChatMessageVisibilityBoth, input.modelConfigID, contentVersion, toolContent),
|
||||
|
||||
@@ -365,10 +365,6 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func sqlNullString(value string) sql.NullString {
|
||||
return sql.NullString{String: value, Valid: value != ""}
|
||||
}
|
||||
|
||||
func TestAIBridgeRoutingFailClosed(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -1056,11 +1056,6 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
if modelConfigID == uuid.Nil {
|
||||
return database.Chat{}, xerrors.New("model config is required")
|
||||
}
|
||||
childAPIKeyID, err := p.ensureSyntheticAPIKeyID(ctx, parent.OwnerID)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("ensure synthetic API key: %w", err)
|
||||
}
|
||||
|
||||
childPlanMode := parent.PlanMode
|
||||
if opts.planModeOverride != nil {
|
||||
childPlanMode = *opts.planModeOverride
|
||||
@@ -1131,7 +1126,7 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
// workspace context the same way a top-level chat does: pinned from the
|
||||
// agent's latest snapshot (see hydrateChatContextOnCreate below). The
|
||||
// parent's context is not copied into child history.
|
||||
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, modelConfigID, parent.OwnerID, childAPIKeyID, opts.reasoningEffortOverride))
|
||||
initialMessages = append(initialMessages, userMessage(userContent, modelConfigID, parent.OwnerID, opts.reasoningEffortOverride))
|
||||
|
||||
publisher := p.pubsub
|
||||
if publisher == nil {
|
||||
|
||||
@@ -651,89 +651,6 @@ func upsertInternalUserChatPersonalModelOverride(
|
||||
)
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChatPersistsOwnerSyntheticAPIKeyID(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, model := seedInternalChatDeps(t, db)
|
||||
parent := createInternalParentChat(
|
||||
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-child-key",
|
||||
)
|
||||
|
||||
child, err := server.createChildSubagentChatWithOptions(
|
||||
ctx,
|
||||
parent,
|
||||
"inspect the workspace",
|
||||
"",
|
||||
childSubagentChatOptions{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: child.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
for _, message := range messages {
|
||||
if message.Role != database.ChatMessageRoleUser {
|
||||
continue
|
||||
}
|
||||
require.True(t, message.APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, message.APIKeyID.String)
|
||||
return
|
||||
}
|
||||
require.Fail(t, "child user message not found")
|
||||
}
|
||||
|
||||
func TestSendSubagentMessagePersistsOwnerSyntheticAPIKeyID(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, model := seedInternalChatDeps(t, db)
|
||||
parent, child := createParentChildChats(ctx, t, server, user, org, model)
|
||||
setChatStatus(ctx, t, db, child.ID, database.ChatStatusWaiting, "")
|
||||
|
||||
_, err := server.sendSubagentMessage(
|
||||
ctx,
|
||||
parent.ID,
|
||||
child.ID,
|
||||
"follow up",
|
||||
SendMessageBusyBehaviorInterrupt,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{
|
||||
UserID: user.ID,
|
||||
TokenName: GatewayTokenName(user.ID),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: child.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
var latestUserMessage database.ChatMessage
|
||||
for _, message := range messages {
|
||||
if message.Role == database.ChatMessageRoleUser && message.ID > latestUserMessage.ID {
|
||||
latestUserMessage = message
|
||||
}
|
||||
}
|
||||
require.NotZero(t, latestUserMessage.ID)
|
||||
require.True(t, latestUserMessage.APIKeyID.Valid)
|
||||
require.Equal(t, gatewayKey.ID, latestUserMessage.APIKeyID.String)
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -217,18 +217,16 @@ func TestSyntheticAPIKeyDeletionDoesNotMutateChatState(t *testing.T) {
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
})
|
||||
message := dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
dbgen.ChatMessage(t, db, database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
Role: database.ChatMessageRoleUser,
|
||||
APIKeyID: sql.NullString{String: syntheticID, Valid: true},
|
||||
})
|
||||
queued, err := db.InsertChatQueuedMessage(t.Context(), database.InsertChatQueuedMessageParams{
|
||||
_, err = db.InsertChatQueuedMessage(t.Context(), database.InsertChatQueuedMessageParams{
|
||||
ChatID: chat.ID,
|
||||
Content: json.RawMessage(`[]`),
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: syntheticID, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -247,16 +245,6 @@ func TestSyntheticAPIKeyDeletionDoesNotMutateChatState(t *testing.T) {
|
||||
require.Equal(t, before.QueueVersion, after.QueueVersion)
|
||||
require.Equal(t, before.GenerationAttempt, after.GenerationAttempt)
|
||||
|
||||
stored, err := db.GetChatMessageByID(t.Context(), message.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, sql.NullString{String: syntheticID, Valid: true}, stored.APIKeyID)
|
||||
storedQueued, err := db.GetChatQueuedMessageByID(t.Context(), database.GetChatQueuedMessageByIDParams{
|
||||
ID: queued.ID,
|
||||
ChatID: chat.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, sql.NullString{String: syntheticID, Valid: true}, storedQueued.APIKeyID)
|
||||
|
||||
remintedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID)
|
||||
require.NoError(t, err)
|
||||
require.NotEqual(t, syntheticID, remintedID)
|
||||
|
||||
@@ -1044,7 +1044,6 @@ func taskUserTextMessage(t *testing.T, text string, createdBy uuid.UUID, modelCo
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: apiKeyID != ""},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -57,7 +57,6 @@ func TestUpdateLastTurnSummaryRejectsStaleWrites(t *testing.T) {
|
||||
codersdk.ChatMessageText("hello"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: owner.ID})
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: owner.ID,
|
||||
@@ -72,7 +71,6 @@ func TestUpdateLastTurnSummaryRejectsStaleWrites(t *testing.T) {
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelCfg.ID, Valid: true},
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user