mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: provider key policies and user provider settings (#23751)
This commit is contained in:
+65
-19
@@ -1585,17 +1585,17 @@ func (p *Server) acquireManualTitleLock(ctx context.Context, chatID uuid.UUID) e
|
||||
if err != nil {
|
||||
return xerrors.Errorf("lock chat for manual title regeneration: %w", err)
|
||||
}
|
||||
if isFreshManualTitleLock(lockedChat, now) {
|
||||
// Only a fresh manual lock or a chat without a real worker should
|
||||
// block title regeneration. Running chats with a real worker may
|
||||
// regenerate their title concurrently, and last write wins.
|
||||
hasRealWorker := lockedChat.Status == database.ChatStatusRunning &&
|
||||
lockedChat.WorkerID.Valid &&
|
||||
lockedChat.WorkerID.UUID != manualTitleLockWorkerID
|
||||
if lockedChat.Status == database.ChatStatusPending ||
|
||||
(lockedChat.Status == database.ChatStatusRunning && !hasRealWorker) ||
|
||||
isFreshManualTitleLock(lockedChat, now) {
|
||||
return ErrManualTitleRegenerationInProgress
|
||||
}
|
||||
|
||||
// Only write the lock marker when no real worker owns WorkerID.
|
||||
// When a real worker is running, we skip the DB lock but still
|
||||
// allow regeneration. The frontend prevents same-browser
|
||||
// double-clicks, and concurrent regeneration from different
|
||||
// replicas is harmless, last write wins.
|
||||
hasRealWorker := lockedChat.WorkerID.Valid &&
|
||||
lockedChat.WorkerID.UUID != manualTitleLockWorkerID
|
||||
if hasRealWorker {
|
||||
return nil
|
||||
}
|
||||
@@ -1658,7 +1658,7 @@ func (p *Server) RegenerateChatTitle(
|
||||
// keeping chat ownership authorization at the HTTP layer.
|
||||
//nolint:gocritic // Non-admin users need chatd-scoped config reads here.
|
||||
chatdCtx := dbauthz.AsChatd(ctx)
|
||||
keys, err := p.resolveProviderAPIKeys(chatdCtx)
|
||||
keys, err := p.resolveUserProviderAPIKeys(chatdCtx, chat.OwnerID)
|
||||
if err != nil {
|
||||
return database.Chat{}, xerrors.Errorf("resolve chat providers: %w", err)
|
||||
}
|
||||
@@ -4808,7 +4808,7 @@ func (p *Server) resolveChatModel(
|
||||
})
|
||||
g.Go(func() error {
|
||||
var err error
|
||||
keys, err = p.resolveProviderAPIKeys(ctx)
|
||||
keys, err = p.resolveUserProviderAPIKeys(ctx, chat.OwnerID)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("resolve provider API keys: %w", err)
|
||||
}
|
||||
@@ -4830,8 +4830,9 @@ func (p *Server) resolveChatModel(
|
||||
return model, dbConfig, keys, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolveProviderAPIKeys(
|
||||
func (p *Server) resolveUserProviderAPIKeys(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
) (chatprovider.ProviderAPIKeys, error) {
|
||||
providers, err := p.configCache.EnabledProviders(ctx)
|
||||
if err != nil {
|
||||
@@ -4840,17 +4841,62 @@ func (p *Server) resolveProviderAPIKeys(
|
||||
err,
|
||||
)
|
||||
}
|
||||
dbProviders := make(
|
||||
configuredProviders := make(
|
||||
[]chatprovider.ConfiguredProvider, 0, len(providers),
|
||||
)
|
||||
for _, provider := range providers {
|
||||
dbProviders = append(dbProviders, chatprovider.ConfiguredProvider{
|
||||
Provider: provider.Provider,
|
||||
APIKey: provider.APIKey,
|
||||
BaseURL: provider.BaseUrl,
|
||||
})
|
||||
configuredProviders = append(
|
||||
configuredProviders, chatprovider.ConfiguredProvider{
|
||||
ProviderID: provider.ID,
|
||||
Provider: provider.Provider,
|
||||
APIKey: provider.APIKey,
|
||||
BaseURL: provider.BaseUrl,
|
||||
CentralAPIKeyEnabled: provider.CentralApiKeyEnabled,
|
||||
AllowUserAPIKey: provider.AllowUserApiKey,
|
||||
AllowCentralAPIKeyFallback: provider.AllowCentralApiKeyFallback,
|
||||
},
|
||||
)
|
||||
}
|
||||
return chatprovider.MergeProviderAPIKeys(p.providerAPIKeys, dbProviders), nil
|
||||
allowAnyUserAPIKey := false
|
||||
for _, provider := range configuredProviders {
|
||||
if provider.AllowUserAPIKey {
|
||||
allowAnyUserAPIKey = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
userKeys := []chatprovider.UserProviderKey{}
|
||||
if allowAnyUserAPIKey {
|
||||
userKeyRows, err := p.db.GetUserChatProviderKeys(ctx, ownerID)
|
||||
if err != nil {
|
||||
return chatprovider.ProviderAPIKeys{}, xerrors.Errorf(
|
||||
"get user chat provider keys: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
userKeys = make([]chatprovider.UserProviderKey, 0, len(userKeyRows))
|
||||
for _, userKey := range userKeyRows {
|
||||
userKeys = append(userKeys, chatprovider.UserProviderKey{
|
||||
ChatProviderID: userKey.ChatProviderID,
|
||||
APIKey: userKey.APIKey,
|
||||
})
|
||||
}
|
||||
}
|
||||
keys, _ := chatprovider.ResolveUserProviderKeys(
|
||||
p.providerAPIKeys,
|
||||
configuredProviders,
|
||||
userKeys,
|
||||
)
|
||||
enabledProviders := make(map[string]struct{}, len(configuredProviders))
|
||||
for _, provider := range configuredProviders {
|
||||
normalizedProvider := chatprovider.NormalizeProvider(provider.Provider)
|
||||
if normalizedProvider == "" {
|
||||
continue
|
||||
}
|
||||
enabledProviders[normalizedProvider] = struct{}{}
|
||||
}
|
||||
chatprovider.PruneDisabledProviderKeys(&keys, enabledProviders)
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// resolveModelConfig looks up the chat's model config by its
|
||||
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
@@ -99,9 +100,10 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
|
||||
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil)
|
||||
db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return([]database.ChatProvider{{
|
||||
Provider: "openai",
|
||||
APIKey: "test-key",
|
||||
BaseUrl: serverURL,
|
||||
Provider: "openai",
|
||||
CentralApiKeyEnabled: true,
|
||||
APIKey: "test-key",
|
||||
BaseUrl: serverURL,
|
||||
}}, nil)
|
||||
db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows)
|
||||
db.EXPECT().GetChatMessagesByChatIDAscPaginated(
|
||||
@@ -261,9 +263,10 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
|
||||
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil)
|
||||
db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return([]database.ChatProvider{{
|
||||
Provider: "openai",
|
||||
APIKey: "test-key",
|
||||
BaseUrl: serverURL,
|
||||
Provider: "openai",
|
||||
CentralApiKeyEnabled: true,
|
||||
APIKey: "test-key",
|
||||
BaseUrl: serverURL,
|
||||
}}, nil)
|
||||
db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows)
|
||||
db.EXPECT().GetChatMessagesByChatIDAscPaginated(
|
||||
@@ -378,6 +381,87 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveUserProviderAPIKeys_StripsDisabledFallbackKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
ownerID := uuid.New()
|
||||
|
||||
server := &Server{
|
||||
db: db,
|
||||
configCache: newChatConfigCache(
|
||||
context.Background(),
|
||||
db,
|
||||
quartz.NewReal(),
|
||||
),
|
||||
providerAPIKeys: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "openai-deployment-key",
|
||||
Anthropic: "anthropic-deployment-key",
|
||||
ByProvider: map[string]string{
|
||||
"openai": "openai-deployment-key",
|
||||
"anthropic": "anthropic-deployment-key",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
"openai": "https://openai.example.com",
|
||||
"anthropic": "https://anthropic.example.com",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return([]database.ChatProvider{{
|
||||
Provider: "anthropic",
|
||||
CentralApiKeyEnabled: true,
|
||||
AllowCentralApiKeyFallback: true,
|
||||
}}, nil)
|
||||
|
||||
keys, err := server.resolveUserProviderAPIKeys(ctx, ownerID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, keys.OpenAI)
|
||||
require.Empty(t, keys.APIKey("openai"))
|
||||
require.Empty(t, keys.BaseURL("openai"))
|
||||
require.Equal(t, "anthropic-deployment-key", keys.Anthropic)
|
||||
require.Equal(t, "anthropic-deployment-key", keys.APIKey("anthropic"))
|
||||
require.Equal(t, "https://anthropic.example.com", keys.BaseURL("anthropic"))
|
||||
require.Equal(t, map[string]string{"anthropic": "anthropic-deployment-key"}, keys.ByProvider)
|
||||
require.Equal(t, map[string]string{"anthropic": "https://anthropic.example.com"}, keys.BaseURLByProvider)
|
||||
}
|
||||
|
||||
func TestResolveUserProviderAPIKeys_SkipsUserKeyLookupWhenNoProviderAllowsUserKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
ownerID := uuid.New()
|
||||
|
||||
server := &Server{
|
||||
db: db,
|
||||
configCache: newChatConfigCache(
|
||||
context.Background(),
|
||||
db,
|
||||
quartz.NewReal(),
|
||||
),
|
||||
providerAPIKeys: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "openai-deployment-key",
|
||||
ByProvider: map[string]string{
|
||||
"openai": "openai-deployment-key",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
db.EXPECT().GetEnabledChatProviders(gomock.Any()).Return([]database.ChatProvider{{
|
||||
Provider: "openai",
|
||||
CentralApiKeyEnabled: true,
|
||||
}}, nil)
|
||||
|
||||
keys, err := server.resolveUserProviderAPIKeys(ctx, ownerID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "openai-deployment-key", keys.OpenAI)
|
||||
require.Equal(t, "openai-deployment-key", keys.APIKey("openai"))
|
||||
}
|
||||
|
||||
func TestRefreshChatWorkspaceSnapshot_NoReloadWhenWorkspacePresent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -523,7 +607,8 @@ func TestPersistInstructionFilesIncludesAgentMetadata(t *testing.T) {
|
||||
workspacesdk.LSResponse{},
|
||||
codersdk.NewTestError(404, "POST", "/api/v0/list-directory"),
|
||||
).AnyTimes()
|
||||
conn.EXPECT().ReadFile(gomock.Any(),
|
||||
conn.EXPECT().ReadFile(
|
||||
gomock.Any(),
|
||||
"/home/coder/project/AGENTS.md",
|
||||
int64(0),
|
||||
int64(maxInstructionFileBytes+1)).Return(
|
||||
|
||||
+248
-18
@@ -2893,12 +2893,13 @@ func seedChatDependenciesWithProvider(
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: provider,
|
||||
DisplayName: provider,
|
||||
APIKey: "test-key",
|
||||
BaseUrl: baseURL,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
Provider: provider,
|
||||
DisplayName: provider,
|
||||
APIKey: "test-key",
|
||||
BaseUrl: baseURL,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
model, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
@@ -2917,6 +2918,102 @@ func seedChatDependenciesWithProvider(
|
||||
return user, model
|
||||
}
|
||||
|
||||
func seedChatDependenciesWithProviderPolicy(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
provider string,
|
||||
baseURL string,
|
||||
apiKey string,
|
||||
centralAPIKeyEnabled bool,
|
||||
allowUserAPIKey bool,
|
||||
allowCentralAPIKeyFallback bool,
|
||||
) (database.User, database.ChatProvider, database.ChatModelConfig) {
|
||||
t.Helper()
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
providerConfig, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: provider,
|
||||
DisplayName: provider,
|
||||
APIKey: apiKey,
|
||||
BaseUrl: baseURL,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: centralAPIKeyEnabled,
|
||||
AllowUserApiKey: allowUserAPIKey,
|
||||
AllowCentralApiKeyFallback: allowCentralAPIKeyFallback,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
model, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
Provider: provider,
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
IsDefault: true,
|
||||
ContextLimit: 128000,
|
||||
CompressionThreshold: 70,
|
||||
Options: json.RawMessage(`{}`),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
return user, providerConfig, model
|
||||
}
|
||||
|
||||
func waitForTerminalChatStatusEvent(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
events <-chan codersdk.ChatStreamEvent,
|
||||
) codersdk.ChatStatus {
|
||||
t.Helper()
|
||||
|
||||
var terminalStatus codersdk.ChatStatus
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
for {
|
||||
select {
|
||||
case event, ok := <-events:
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if event.Type != codersdk.ChatStreamEventTypeStatus || event.Status == nil {
|
||||
continue
|
||||
}
|
||||
if event.Status.Status == codersdk.ChatStatusWaiting || event.Status.Status == codersdk.ChatStatusError {
|
||||
terminalStatus = event.Status.Status
|
||||
return true
|
||||
}
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
return terminalStatus
|
||||
}
|
||||
|
||||
func waitForTerminalChat(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
chatID uuid.UUID,
|
||||
) database.Chat {
|
||||
t.Helper()
|
||||
|
||||
var chatResult database.Chat
|
||||
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
|
||||
got, err := db.GetChatByID(ctx, chatID)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
chatResult = got
|
||||
return got.Status == database.ChatStatusWaiting || got.Status == database.ChatStatusError
|
||||
}, testutil.IntervalFast)
|
||||
|
||||
return chatResult
|
||||
}
|
||||
|
||||
// seedWorkspaceWithAgent creates a full workspace chain with a connected
|
||||
// agent. This is the common setup needed by tests that exercise tool
|
||||
// execution against a workspace.
|
||||
@@ -2973,12 +3070,15 @@ func setOpenAIProviderBaseURL(
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = db.UpdateChatProvider(ctx, database.UpdateChatProviderParams{
|
||||
ID: provider.ID,
|
||||
DisplayName: provider.DisplayName,
|
||||
APIKey: provider.APIKey,
|
||||
BaseUrl: baseURL,
|
||||
ApiKeyKeyID: provider.ApiKeyKeyID,
|
||||
Enabled: provider.Enabled,
|
||||
ID: provider.ID,
|
||||
DisplayName: provider.DisplayName,
|
||||
APIKey: provider.APIKey,
|
||||
BaseUrl: baseURL,
|
||||
ApiKeyKeyID: provider.ApiKeyKeyID,
|
||||
Enabled: provider.Enabled,
|
||||
CentralApiKeyEnabled: provider.CentralApiKeyEnabled,
|
||||
AllowUserApiKey: provider.AllowUserApiKey,
|
||||
AllowCentralApiKeyFallback: provider.AllowCentralApiKeyFallback,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
@@ -3552,12 +3652,13 @@ func TestComputerUseSubagentToolsAndModel(t *testing.T) {
|
||||
|
||||
// Add an Anthropic provider pointing to our mock server.
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: "anthropic",
|
||||
DisplayName: "Anthropic",
|
||||
APIKey: "test-anthropic-key",
|
||||
BaseUrl: anthropicSrv.URL,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
Provider: "anthropic",
|
||||
DisplayName: "Anthropic",
|
||||
APIKey: "test-anthropic-key",
|
||||
BaseUrl: anthropicSrv.URL,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -3841,6 +3942,135 @@ func TestInterruptChatPersistsPartialResponse(t *testing.T) {
|
||||
"partial assistant response should contain the streamed text")
|
||||
}
|
||||
|
||||
func TestProcessChat_UserProviderKey_Success(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
const userAPIKey = "user-test-key"
|
||||
|
||||
var authHeadersMu sync.Mutex
|
||||
authHeaders := make([]string, 0, 1)
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
authHeadersMu.Lock()
|
||||
authHeaders = append(authHeaders, req.Header.Get("Authorization"))
|
||||
authHeadersMu.Unlock()
|
||||
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("user provider key success")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("hello from the saved user key")...,
|
||||
)
|
||||
})
|
||||
|
||||
user, provider, model := seedChatDependenciesWithProviderPolicy(
|
||||
ctx,
|
||||
t,
|
||||
db,
|
||||
"openai-compat",
|
||||
openAIURL,
|
||||
"",
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
_, err := db.UpsertUserChatProviderKey(ctx, database.UpsertUserChatProviderKeyParams{
|
||||
UserID: user.ID,
|
||||
ChatProviderID: provider.ID,
|
||||
APIKey: userAPIKey,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
creator := newTestServer(t, db, ps, uuid.New())
|
||||
chat, err := creator.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "user-provider-key-success",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("say hello"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, events, cancel, ok := creator.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
t.Cleanup(cancel)
|
||||
|
||||
_ = newActiveTestServer(t, db, ps)
|
||||
|
||||
terminalStatus := waitForTerminalChatStatusEvent(ctx, t, events)
|
||||
require.Equal(t, codersdk.ChatStatusWaiting, terminalStatus)
|
||||
|
||||
chatResult := waitForTerminalChat(ctx, t, db, chat.ID)
|
||||
require.Equal(t, database.ChatStatusWaiting, chatResult.Status)
|
||||
require.False(t, chatResult.LastError.Valid)
|
||||
|
||||
authHeadersMu.Lock()
|
||||
recordedAuthHeaders := append([]string(nil), authHeaders...)
|
||||
authHeadersMu.Unlock()
|
||||
require.Contains(t, recordedAuthHeaders, "Bearer "+userAPIKey)
|
||||
}
|
||||
|
||||
func TestProcessChat_UserProviderKey_MissingKeyError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
var llmCalls atomic.Int32
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
llmCalls.Add(1)
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("unexpected non-streaming request")
|
||||
}
|
||||
return chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("unexpected streaming request")...,
|
||||
)
|
||||
})
|
||||
|
||||
user, _, model := seedChatDependenciesWithProviderPolicy(
|
||||
ctx,
|
||||
t,
|
||||
db,
|
||||
"openai-compat",
|
||||
openAIURL,
|
||||
"",
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
|
||||
creator := newTestServer(t, db, ps, uuid.New())
|
||||
chat, err := creator.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
Title: "user-provider-key-missing",
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("say hello"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, events, cancel, ok := creator.Subscribe(ctx, chat.ID, nil, 0)
|
||||
require.True(t, ok)
|
||||
t.Cleanup(cancel)
|
||||
|
||||
_ = newActiveTestServer(t, db, ps)
|
||||
|
||||
terminalStatus := waitForTerminalChatStatusEvent(ctx, t, events)
|
||||
require.Equal(t, codersdk.ChatStatusError, terminalStatus)
|
||||
|
||||
chatResult := waitForTerminalChat(ctx, t, db, chat.ID)
|
||||
require.Equal(t, database.ChatStatusError, chatResult.Status)
|
||||
require.True(t, chatResult.LastError.Valid, "LastError should be set")
|
||||
require.NotEmpty(t, chatResult.LastError.String)
|
||||
require.NotContains(t, chatResult.LastError.String, "panicked")
|
||||
require.NotEqual(t, database.ChatStatusRunning, chatResult.Status)
|
||||
require.Zero(t, llmCalls.Load(), "missing user key should fail before any LLM request")
|
||||
}
|
||||
|
||||
func TestProcessChatPanicRecovery(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -1459,11 +1459,12 @@ func TestNulEscapeRoundTrip(t *testing.T) {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: "openai",
|
||||
DisplayName: "openai",
|
||||
APIKey: "test-key",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
Provider: "openai",
|
||||
DisplayName: "openai",
|
||||
APIKey: "test-key",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -1943,11 +1944,12 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: "anthropic",
|
||||
DisplayName: "anthropic",
|
||||
APIKey: "test-key",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
Provider: "anthropic",
|
||||
DisplayName: "anthropic",
|
||||
APIKey: "test-key",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
@@ -81,11 +81,28 @@ type ProviderAPIKeys struct {
|
||||
BaseURLByProvider map[string]string
|
||||
}
|
||||
|
||||
// UserProviderKey is a user-supplied API key for a specific provider.
|
||||
type UserProviderKey struct {
|
||||
ChatProviderID uuid.UUID
|
||||
APIKey string
|
||||
}
|
||||
|
||||
// ProviderAvailability describes whether a provider has a usable
|
||||
// API key and, if not, why.
|
||||
type ProviderAvailability struct {
|
||||
Available bool
|
||||
UnavailableReason codersdk.ChatModelProviderUnavailableReason
|
||||
}
|
||||
|
||||
// ConfiguredProvider is an enabled provider loaded from database config.
|
||||
type ConfiguredProvider struct {
|
||||
Provider string
|
||||
APIKey string
|
||||
BaseURL string
|
||||
ProviderID uuid.UUID
|
||||
Provider string
|
||||
APIKey string
|
||||
BaseURL string
|
||||
CentralAPIKeyEnabled bool
|
||||
AllowUserAPIKey bool
|
||||
AllowCentralAPIKeyFallback bool
|
||||
}
|
||||
|
||||
// ConfiguredModel is an enabled model loaded from database config.
|
||||
@@ -189,21 +206,146 @@ func MergeProviderAPIKeys(fallback ProviderAPIKeys, providers []ConfiguredProvid
|
||||
return merged
|
||||
}
|
||||
|
||||
type ModelCatalog struct {
|
||||
keys ProviderAPIKeys
|
||||
// ResolveUserProviderKeys computes effective API keys and per-provider
|
||||
// availability for a given user. It considers the provider's credential
|
||||
// policy flags alongside central (DB/deployment) keys and the user's
|
||||
// personal keys.
|
||||
func ResolveUserProviderKeys(
|
||||
fallback ProviderAPIKeys,
|
||||
providers []ConfiguredProvider,
|
||||
userKeys []UserProviderKey,
|
||||
) (ProviderAPIKeys, map[string]ProviderAvailability) {
|
||||
merged := ProviderAPIKeys{
|
||||
OpenAI: strings.TrimSpace(fallback.OpenAI),
|
||||
Anthropic: strings.TrimSpace(fallback.Anthropic),
|
||||
ByProvider: map[string]string{},
|
||||
BaseURLByProvider: map[string]string{},
|
||||
}
|
||||
for provider, apiKey := range fallback.ByProvider {
|
||||
normalizedProvider := NormalizeProvider(provider)
|
||||
if normalizedProvider == "" {
|
||||
continue
|
||||
}
|
||||
if key := strings.TrimSpace(apiKey); key != "" {
|
||||
merged.ByProvider[normalizedProvider] = key
|
||||
}
|
||||
}
|
||||
for provider, baseURL := range fallback.BaseURLByProvider {
|
||||
normalizedProvider := NormalizeProvider(provider)
|
||||
if normalizedProvider == "" {
|
||||
continue
|
||||
}
|
||||
if url := strings.TrimSpace(baseURL); url != "" {
|
||||
merged.BaseURLByProvider[normalizedProvider] = url
|
||||
}
|
||||
}
|
||||
if merged.OpenAI != "" {
|
||||
merged.ByProvider[fantasyopenai.Name] = merged.OpenAI
|
||||
}
|
||||
if merged.Anthropic != "" {
|
||||
merged.ByProvider[fantasyanthropic.Name] = merged.Anthropic
|
||||
}
|
||||
|
||||
userKeyByProviderID := make(map[uuid.UUID]string, len(userKeys))
|
||||
for _, userKey := range userKeys {
|
||||
if userKey.ChatProviderID == uuid.Nil {
|
||||
continue
|
||||
}
|
||||
if key := strings.TrimSpace(userKey.APIKey); key != "" {
|
||||
userKeyByProviderID[userKey.ChatProviderID] = key
|
||||
}
|
||||
}
|
||||
|
||||
availabilityByProvider := make(map[string]ProviderAvailability, len(providers))
|
||||
for _, provider := range providers {
|
||||
normalizedProvider := NormalizeProvider(provider.Provider)
|
||||
if normalizedProvider == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
if url := strings.TrimSpace(provider.BaseURL); url != "" {
|
||||
merged.BaseURLByProvider[normalizedProvider] = url
|
||||
}
|
||||
|
||||
var userKey string
|
||||
if provider.ProviderID != uuid.Nil {
|
||||
userKey = userKeyByProviderID[provider.ProviderID]
|
||||
}
|
||||
|
||||
var centralKey string
|
||||
if provider.CentralAPIKeyEnabled {
|
||||
if key := strings.TrimSpace(provider.APIKey); key != "" {
|
||||
centralKey = key
|
||||
} else {
|
||||
centralKey = fallback.APIKey(normalizedProvider)
|
||||
}
|
||||
}
|
||||
|
||||
resolved := ProviderAvailability{}
|
||||
chosenKey := ""
|
||||
switch {
|
||||
case provider.AllowUserAPIKey && userKey != "":
|
||||
chosenKey = userKey
|
||||
resolved.Available = true
|
||||
case centralKey != "":
|
||||
if !provider.AllowUserAPIKey || provider.AllowCentralAPIKeyFallback {
|
||||
chosenKey = centralKey
|
||||
resolved.Available = true
|
||||
} else {
|
||||
resolved.UnavailableReason = codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired
|
||||
}
|
||||
case provider.AllowUserAPIKey && provider.AllowCentralAPIKeyFallback && provider.CentralAPIKeyEnabled:
|
||||
// When users can add their own key, a missing central fallback key is
|
||||
// still something the user can remedy.
|
||||
resolved.UnavailableReason = codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired
|
||||
case provider.AllowUserAPIKey:
|
||||
resolved.UnavailableReason = codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired
|
||||
default:
|
||||
resolved.UnavailableReason = codersdk.ChatModelProviderUnavailableMissingAPIKey
|
||||
}
|
||||
|
||||
setResolvedProviderAPIKey(&merged, normalizedProvider, chosenKey)
|
||||
availabilityByProvider[normalizedProvider] = resolved
|
||||
}
|
||||
|
||||
return merged, availabilityByProvider
|
||||
}
|
||||
|
||||
func NewModelCatalog(keys ProviderAPIKeys) *ModelCatalog {
|
||||
return &ModelCatalog{
|
||||
keys: keys,
|
||||
func setResolvedProviderAPIKey(keys *ProviderAPIKeys, provider string, apiKey string) {
|
||||
normalizedProvider := NormalizeProvider(provider)
|
||||
if normalizedProvider == "" {
|
||||
return
|
||||
}
|
||||
if keys.ByProvider == nil {
|
||||
keys.ByProvider = map[string]string{}
|
||||
}
|
||||
|
||||
delete(keys.ByProvider, normalizedProvider)
|
||||
trimmedKey := strings.TrimSpace(apiKey)
|
||||
switch normalizedProvider {
|
||||
case fantasyopenai.Name:
|
||||
keys.OpenAI = trimmedKey
|
||||
case fantasyanthropic.Name:
|
||||
keys.Anthropic = trimmedKey
|
||||
}
|
||||
if trimmedKey != "" {
|
||||
keys.ByProvider[normalizedProvider] = trimmedKey
|
||||
}
|
||||
}
|
||||
|
||||
type ModelCatalog struct{}
|
||||
|
||||
func NewModelCatalog() *ModelCatalog {
|
||||
return &ModelCatalog{}
|
||||
}
|
||||
|
||||
// ListConfiguredModels returns a model catalog from enabled DB-backed model
|
||||
// configs. The second return value reports whether DB-backed models were used.
|
||||
func (c *ModelCatalog) ListConfiguredModels(
|
||||
func (*ModelCatalog) ListConfiguredModels(
|
||||
configuredProviders []ConfiguredProvider,
|
||||
configuredModels []ConfiguredModel,
|
||||
availabilityByProvider map[string]ProviderAvailability,
|
||||
enabledProviders map[string]struct{},
|
||||
) (codersdk.ChatModelsResponse, bool) {
|
||||
if len(configuredModels) == 0 {
|
||||
return codersdk.ChatModelsResponse{}, false
|
||||
@@ -247,11 +389,14 @@ func (c *ModelCatalog) ListConfiguredModels(
|
||||
return codersdk.ChatModelsResponse{}, false
|
||||
}
|
||||
|
||||
keys := MergeProviderAPIKeys(c.keys, configuredProviders)
|
||||
response := codersdk.ChatModelsResponse{
|
||||
Providers: make([]codersdk.ChatModelProvider, 0, len(providers)),
|
||||
}
|
||||
for _, provider := range providers {
|
||||
if _, ok := enabledProviders[provider]; !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
models := modelsByProvider[provider]
|
||||
sortChatModels(models)
|
||||
|
||||
@@ -259,11 +404,14 @@ func (c *ModelCatalog) ListConfiguredModels(
|
||||
Provider: provider,
|
||||
Models: models,
|
||||
}
|
||||
if keys.APIKey(provider) == "" {
|
||||
if avail, ok := availabilityByProvider[provider]; ok {
|
||||
result.Available = avail.Available
|
||||
if !avail.Available {
|
||||
result.UnavailableReason = avail.UnavailableReason
|
||||
}
|
||||
} else {
|
||||
result.Available = false
|
||||
result.UnavailableReason = codersdk.ChatModelProviderUnavailableMissingAPIKey
|
||||
} else {
|
||||
result.Available = true
|
||||
}
|
||||
|
||||
response.Providers = append(response.Providers, result)
|
||||
@@ -273,25 +421,32 @@ func (c *ModelCatalog) ListConfiguredModels(
|
||||
}
|
||||
|
||||
// ListConfiguredProviderAvailability returns provider availability derived from
|
||||
// deployment/env keys merged with enabled DB provider keys.
|
||||
func (c *ModelCatalog) ListConfiguredProviderAvailability(
|
||||
configuredProviders []ConfiguredProvider,
|
||||
// the policy-aware availability map for enabled providers.
|
||||
func (*ModelCatalog) ListConfiguredProviderAvailability(
|
||||
availabilityByProvider map[string]ProviderAvailability,
|
||||
enabledProviders map[string]struct{},
|
||||
) codersdk.ChatModelsResponse {
|
||||
keys := MergeProviderAPIKeys(c.keys, configuredProviders)
|
||||
response := codersdk.ChatModelsResponse{
|
||||
Providers: make([]codersdk.ChatModelProvider, 0, len(supportedProviderNames)),
|
||||
}
|
||||
|
||||
for _, provider := range supportedProviderNames {
|
||||
if _, ok := enabledProviders[provider]; !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
result := codersdk.ChatModelProvider{
|
||||
Provider: provider,
|
||||
Models: []codersdk.ChatModel{},
|
||||
}
|
||||
if keys.APIKey(provider) == "" {
|
||||
if avail, ok := availabilityByProvider[provider]; ok {
|
||||
result.Available = avail.Available
|
||||
if !avail.Available {
|
||||
result.UnavailableReason = avail.UnavailableReason
|
||||
}
|
||||
} else {
|
||||
result.Available = false
|
||||
result.UnavailableReason = codersdk.ChatModelProviderUnavailableMissingAPIKey
|
||||
} else {
|
||||
result.Available = true
|
||||
}
|
||||
|
||||
response.Providers = append(response.Providers, result)
|
||||
@@ -300,6 +455,27 @@ func (c *ModelCatalog) ListConfiguredProviderAvailability(
|
||||
return response
|
||||
}
|
||||
|
||||
// PruneDisabledProviderKeys removes entries from keys that do not
|
||||
// belong to an enabled provider. It clears ByProvider and
|
||||
// BaseURLByProvider entries for disabled providers and zeroes the
|
||||
// legacy OpenAI and Anthropic fields when those providers are not
|
||||
// enabled.
|
||||
func PruneDisabledProviderKeys(keys *ProviderAPIKeys, enabledProviders map[string]struct{}) {
|
||||
for provider := range keys.ByProvider {
|
||||
if _, ok := enabledProviders[provider]; ok {
|
||||
continue
|
||||
}
|
||||
delete(keys.ByProvider, provider)
|
||||
delete(keys.BaseURLByProvider, provider)
|
||||
}
|
||||
if _, ok := enabledProviders[NormalizeProvider("openai")]; !ok {
|
||||
keys.OpenAI = ""
|
||||
}
|
||||
if _, ok := enabledProviders[NormalizeProvider("anthropic")]; !ok {
|
||||
keys.Anthropic = ""
|
||||
}
|
||||
}
|
||||
|
||||
func newChatModel(provider, modelID, displayName string) codersdk.ChatModel {
|
||||
name := strings.TrimSpace(displayName)
|
||||
if name == "" {
|
||||
|
||||
@@ -21,6 +21,166 @@ import (
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestResolveUserProviderKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configuredProvider := func(id uuid.UUID, provider string, centralEnabled bool, centralKey string, allowUser bool, allowCentralFallback bool) chatprovider.ConfiguredProvider {
|
||||
return chatprovider.ConfiguredProvider{
|
||||
ProviderID: id,
|
||||
Provider: provider,
|
||||
APIKey: centralKey,
|
||||
CentralAPIKeyEnabled: centralEnabled,
|
||||
AllowUserAPIKey: allowUser,
|
||||
AllowCentralAPIKeyFallback: allowCentralFallback,
|
||||
}
|
||||
}
|
||||
|
||||
userProviderKey := func(id uuid.UUID, apiKey string) chatprovider.UserProviderKey {
|
||||
return chatprovider.UserProviderKey{
|
||||
ChatProviderID: id,
|
||||
APIKey: apiKey,
|
||||
}
|
||||
}
|
||||
|
||||
openAIProviderID := uuid.MustParse("00000000-0000-0000-0000-000000000001")
|
||||
anthropicProviderID := uuid.MustParse("00000000-0000-0000-0000-000000000002")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
fallback chatprovider.ProviderAPIKeys
|
||||
providers []chatprovider.ConfiguredProvider
|
||||
userKeys []chatprovider.UserProviderKey
|
||||
wantAvailability map[string]chatprovider.ProviderAvailability
|
||||
wantKeys map[string]string
|
||||
}{
|
||||
{
|
||||
name: "CentralOnlyKeyPresent",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "sk-central", false, false)},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "sk-central",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "CentralOnlyKeyMissing",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "", false, false)},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: false, UnavailableReason: codersdk.ChatModelProviderUnavailableMissingAPIKey},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "UserOnlyUserHasKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, false, "sk-central", true, false)},
|
||||
userKeys: []chatprovider.UserProviderKey{userProviderKey(openAIProviderID, "sk-user")},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "sk-user",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "UserOnlyUserHasNoKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, false, "sk-central", true, false)},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: false, UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "BothEnabledFallbackOffUserHasKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "sk-central", true, false)},
|
||||
userKeys: []chatprovider.UserProviderKey{userProviderKey(openAIProviderID, "sk-user")},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "sk-user",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "BothEnabledFallbackOffUserHasNoKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "sk-central", true, false)},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: false, UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "BothEnabledFallbackOnUserHasKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "sk-central", true, true)},
|
||||
userKeys: []chatprovider.UserProviderKey{userProviderKey(openAIProviderID, "sk-user")},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "sk-user",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "BothEnabledFallbackOnUserHasNoKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "sk-central", true, true)},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "sk-central",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "BothEnabledFallbackOnCentralKeyEmptyUserHasNoKey",
|
||||
providers: []chatprovider.ConfiguredProvider{configuredProvider(openAIProviderID, fantasyopenai.Name, true, "", true, true)},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: false, UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MultipleProvidersDifferentPolicies",
|
||||
providers: []chatprovider.ConfiguredProvider{
|
||||
configuredProvider(openAIProviderID, fantasyopenai.Name, true, "sk-central", false, false),
|
||||
configuredProvider(anthropicProviderID, fantasyanthropic.Name, false, "", true, false),
|
||||
},
|
||||
wantAvailability: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {Available: true},
|
||||
fantasyanthropic.Name: {Available: false, UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired},
|
||||
},
|
||||
wantKeys: map[string]string{
|
||||
fantasyopenai.Name: "sk-central",
|
||||
fantasyanthropic.Name: "",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
keys, availability := chatprovider.ResolveUserProviderKeys(tt.fallback, tt.providers, tt.userKeys)
|
||||
|
||||
require.Len(t, availability, len(tt.wantAvailability))
|
||||
for provider, wantAvailability := range tt.wantAvailability {
|
||||
gotAvailability, ok := availability[provider]
|
||||
require.True(t, ok, "expected availability for provider %q", provider)
|
||||
require.Equal(t, wantAvailability, gotAvailability)
|
||||
require.Equal(t, tt.wantKeys[provider], keys.APIKey(provider))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReasoningEffortFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -91,6 +251,413 @@ func TestReasoningEffortFromChat(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveUserProviderKeys_UnavailableReason(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
provider chatprovider.ConfiguredProvider
|
||||
wantReason codersdk.ChatModelProviderUnavailableReason
|
||||
}{
|
||||
{
|
||||
name: "FallbackConfiguredWithoutCentralKeyReturnsUserAPIKeyRequired",
|
||||
provider: chatprovider.ConfiguredProvider{
|
||||
Provider: "anthropic",
|
||||
CentralAPIKeyEnabled: true,
|
||||
AllowUserAPIKey: true,
|
||||
AllowCentralAPIKeyFallback: true,
|
||||
},
|
||||
wantReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired,
|
||||
},
|
||||
{
|
||||
name: "UserKeyRequiredWithoutFallback",
|
||||
provider: chatprovider.ConfiguredProvider{
|
||||
Provider: "anthropic",
|
||||
CentralAPIKeyEnabled: true,
|
||||
AllowUserAPIKey: true,
|
||||
},
|
||||
wantReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
keys, availability := chatprovider.ResolveUserProviderKeys(
|
||||
chatprovider.ProviderAPIKeys{},
|
||||
[]chatprovider.ConfiguredProvider{tt.provider},
|
||||
nil,
|
||||
)
|
||||
|
||||
require.Empty(t, keys.APIKey(tt.provider.Provider))
|
||||
resolved, ok := availability[tt.provider.Provider]
|
||||
require.True(t, ok)
|
||||
require.False(t, resolved.Available)
|
||||
require.Equal(t, tt.wantReason, resolved.UnavailableReason)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestListConfiguredModels_PolicyAwareAvailability(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configuredProvider := func(provider string, apiKey string) chatprovider.ConfiguredProvider {
|
||||
return chatprovider.ConfiguredProvider{
|
||||
ProviderID: uuid.New(),
|
||||
Provider: provider,
|
||||
APIKey: apiKey,
|
||||
}
|
||||
}
|
||||
enabledProviders := func(providers ...string) map[string]struct{} {
|
||||
result := make(map[string]struct{}, len(providers))
|
||||
for _, provider := range providers {
|
||||
result[chatprovider.NormalizeProvider(provider)] = struct{}{}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
catalog := chatprovider.NewModelCatalog()
|
||||
tests := []struct {
|
||||
name string
|
||||
configuredProviders []chatprovider.ConfiguredProvider
|
||||
configuredModels []chatprovider.ConfiguredModel
|
||||
availabilityByProvider map[string]chatprovider.ProviderAvailability
|
||||
enabledProviders map[string]struct{}
|
||||
want codersdk.ChatModelsResponse
|
||||
}{
|
||||
{
|
||||
name: "PolicyUnavailableOverridesConfiguredKey",
|
||||
configuredProviders: []chatprovider.ConfiguredProvider{
|
||||
configuredProvider(fantasyopenai.Name, "sk-central"),
|
||||
},
|
||||
configuredModels: []chatprovider.ConfiguredModel{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Model: "gpt-4",
|
||||
}},
|
||||
availabilityByProvider: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyopenai.Name: {
|
||||
Available: false,
|
||||
UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired,
|
||||
},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Available: false,
|
||||
UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired,
|
||||
Models: []codersdk.ChatModel{{
|
||||
ID: fantasyopenai.Name + ":gpt-4",
|
||||
Provider: fantasyopenai.Name,
|
||||
Model: "gpt-4",
|
||||
DisplayName: "gpt-4",
|
||||
}},
|
||||
}}},
|
||||
},
|
||||
{
|
||||
name: "PolicyAvailableMarksProviderAvailable",
|
||||
configuredProviders: []chatprovider.ConfiguredProvider{
|
||||
configuredProvider(fantasyanthropic.Name, "sk-central"),
|
||||
},
|
||||
configuredModels: []chatprovider.ConfiguredModel{{
|
||||
Provider: fantasyanthropic.Name,
|
||||
Model: "claude-3-5-sonnet",
|
||||
}},
|
||||
availabilityByProvider: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyanthropic.Name: {Available: true},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyanthropic.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{{
|
||||
Provider: fantasyanthropic.Name,
|
||||
Available: true,
|
||||
Models: []codersdk.ChatModel{{
|
||||
ID: fantasyanthropic.Name + ":claude-3-5-sonnet",
|
||||
Provider: fantasyanthropic.Name,
|
||||
Model: "claude-3-5-sonnet",
|
||||
DisplayName: "claude-3-5-sonnet",
|
||||
}},
|
||||
}}},
|
||||
},
|
||||
{
|
||||
name: "DisabledProviderOmitted",
|
||||
configuredProviders: []chatprovider.ConfiguredProvider{
|
||||
configuredProvider(fantasyanthropic.Name, "sk-anthropic"),
|
||||
configuredProvider(fantasyopenai.Name, "sk-openai"),
|
||||
},
|
||||
configuredModels: []chatprovider.ConfiguredModel{
|
||||
{Provider: fantasyanthropic.Name, Model: "claude-3-5-sonnet"},
|
||||
{Provider: fantasyopenai.Name, Model: "gpt-4"},
|
||||
},
|
||||
availabilityByProvider: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyanthropic.Name: {Available: true},
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Available: true,
|
||||
Models: []codersdk.ChatModel{{
|
||||
ID: fantasyopenai.Name + ":gpt-4",
|
||||
Provider: fantasyopenai.Name,
|
||||
Model: "gpt-4",
|
||||
DisplayName: "gpt-4",
|
||||
}},
|
||||
}}},
|
||||
},
|
||||
{
|
||||
name: "MissingAvailabilityDefaultsToMissingAPIKey",
|
||||
configuredProviders: []chatprovider.ConfiguredProvider{
|
||||
configuredProvider(fantasyopenai.Name, "sk-central"),
|
||||
},
|
||||
configuredModels: []chatprovider.ConfiguredModel{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Model: "gpt-4o",
|
||||
}},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Available: false,
|
||||
UnavailableReason: codersdk.ChatModelProviderUnavailableMissingAPIKey,
|
||||
Models: []codersdk.ChatModel{{
|
||||
ID: fantasyopenai.Name + ":gpt-4o",
|
||||
Provider: fantasyopenai.Name,
|
||||
Model: "gpt-4o",
|
||||
DisplayName: "gpt-4o",
|
||||
}},
|
||||
}}},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got, ok := catalog.ListConfiguredModels(
|
||||
tt.configuredProviders,
|
||||
tt.configuredModels,
|
||||
tt.availabilityByProvider,
|
||||
tt.enabledProviders,
|
||||
)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestListConfiguredProviderAvailability_PolicyAwareFiltering(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
enabledProviders := func(providers ...string) map[string]struct{} {
|
||||
result := make(map[string]struct{}, len(providers))
|
||||
for _, provider := range providers {
|
||||
result[chatprovider.NormalizeProvider(provider)] = struct{}{}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
catalog := chatprovider.NewModelCatalog()
|
||||
tests := []struct {
|
||||
name string
|
||||
availabilityByProvider map[string]chatprovider.ProviderAvailability
|
||||
enabledProviders map[string]struct{}
|
||||
want codersdk.ChatModelsResponse
|
||||
}{
|
||||
{
|
||||
name: "EnabledProvidersUsePolicyAvailability",
|
||||
availabilityByProvider: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyanthropic.Name: {
|
||||
Available: false,
|
||||
UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired,
|
||||
},
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyanthropic.Name, fantasyopenai.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{
|
||||
{
|
||||
Provider: fantasyanthropic.Name,
|
||||
Available: false,
|
||||
UnavailableReason: codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired,
|
||||
Models: []codersdk.ChatModel{},
|
||||
},
|
||||
{
|
||||
Provider: fantasyopenai.Name,
|
||||
Available: true,
|
||||
Models: []codersdk.ChatModel{},
|
||||
},
|
||||
}},
|
||||
},
|
||||
{
|
||||
name: "DisabledSupportedProviderOmitted",
|
||||
availabilityByProvider: map[string]chatprovider.ProviderAvailability{
|
||||
fantasyanthropic.Name: {Available: true},
|
||||
fantasyopenai.Name: {Available: true},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Available: true,
|
||||
Models: []codersdk.ChatModel{},
|
||||
}}},
|
||||
},
|
||||
{
|
||||
name: "MissingAvailabilityDefaultsToMissingAPIKey",
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: codersdk.ChatModelsResponse{Providers: []codersdk.ChatModelProvider{{
|
||||
Provider: fantasyopenai.Name,
|
||||
Available: false,
|
||||
UnavailableReason: codersdk.ChatModelProviderUnavailableMissingAPIKey,
|
||||
Models: []codersdk.ChatModel{},
|
||||
}}},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := catalog.ListConfiguredProviderAvailability(
|
||||
tt.availabilityByProvider,
|
||||
tt.enabledProviders,
|
||||
)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneDisabledProviderKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
enabledProviders := func(providers ...string) map[string]struct{} {
|
||||
result := make(map[string]struct{}, len(providers))
|
||||
for _, provider := range providers {
|
||||
result[chatprovider.NormalizeProvider(provider)] = struct{}{}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
keys chatprovider.ProviderAPIKeys
|
||||
enabledProviders map[string]struct{}
|
||||
want chatprovider.ProviderAPIKeys
|
||||
}{
|
||||
{
|
||||
name: "DisabledProviderEntriesRemoved",
|
||||
keys: chatprovider.ProviderAPIKeys{
|
||||
ByProvider: map[string]string{
|
||||
fantasyanthropic.Name: "sk-anthropic",
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyanthropic.Name: "https://anthropic.example.com",
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: chatprovider.ProviderAPIKeys{
|
||||
ByProvider: map[string]string{
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "OpenAIDisabledClearsLegacyField",
|
||||
keys: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "sk-openai",
|
||||
Anthropic: "sk-anthropic",
|
||||
ByProvider: map[string]string{
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
fantasyanthropic.Name: "sk-anthropic",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
fantasyanthropic.Name: "https://anthropic.example.com",
|
||||
},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyanthropic.Name),
|
||||
want: chatprovider.ProviderAPIKeys{
|
||||
Anthropic: "sk-anthropic",
|
||||
ByProvider: map[string]string{
|
||||
fantasyanthropic.Name: "sk-anthropic",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyanthropic.Name: "https://anthropic.example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "AnthropicDisabledClearsLegacyField",
|
||||
keys: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "sk-openai",
|
||||
Anthropic: "sk-anthropic",
|
||||
ByProvider: map[string]string{
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
fantasyanthropic.Name: "sk-anthropic",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
fantasyanthropic.Name: "https://anthropic.example.com",
|
||||
},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name),
|
||||
want: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "sk-openai",
|
||||
ByProvider: map[string]string{
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "AllEnabledLeavesKeysUnchanged",
|
||||
keys: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "sk-openai",
|
||||
Anthropic: "sk-anthropic",
|
||||
ByProvider: map[string]string{
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
fantasyanthropic.Name: "sk-anthropic",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
fantasyanthropic.Name: "https://anthropic.example.com",
|
||||
},
|
||||
},
|
||||
enabledProviders: enabledProviders(fantasyopenai.Name, fantasyanthropic.Name),
|
||||
want: chatprovider.ProviderAPIKeys{
|
||||
OpenAI: "sk-openai",
|
||||
Anthropic: "sk-anthropic",
|
||||
ByProvider: map[string]string{
|
||||
fantasyopenai.Name: "sk-openai",
|
||||
fantasyanthropic.Name: "sk-anthropic",
|
||||
},
|
||||
BaseURLByProvider: map[string]string{
|
||||
fantasyopenai.Name: "https://openai.example.com",
|
||||
fantasyanthropic.Name: "https://anthropic.example.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
keys := tt.keys
|
||||
chatprovider.PruneDisabledProviderKeys(&keys, tt.enabledProviders)
|
||||
require.Equal(t, tt.want, keys)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCoderHeaders(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -440,13 +440,14 @@ func seedModelConfig(
|
||||
t.Helper()
|
||||
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
APIKey: "test-key",
|
||||
BaseUrl: "",
|
||||
ApiKeyKeyID: sql.NullString{},
|
||||
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
Enabled: true,
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
APIKey: "test-key",
|
||||
BaseUrl: "",
|
||||
ApiKeyKeyID: sql.NullString{},
|
||||
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
@@ -122,13 +122,14 @@ func seedInternalChatDeps(
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
APIKey: "test-key",
|
||||
BaseUrl: "",
|
||||
ApiKeyKeyID: sql.NullString{},
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
APIKey: "test-key",
|
||||
BaseUrl: "",
|
||||
ApiKeyKeyID: sql.NullString{},
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
@@ -959,9 +959,10 @@ func TestWorker(t *testing.T) {
|
||||
|
||||
// 3. Set up FK chain: chat_providers -> chat_model_configs -> chats.
|
||||
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
Enabled: true,
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
Enabled: true,
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user