mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add personal chat model overrides (#24715)
This commit is contained in:
@@ -0,0 +1,75 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
// ChatPersonalModelOverrideKeyPrefix is the user config key prefix for
|
||||
// chat personal model overrides. Values under this prefix should be parsed
|
||||
// with ParseChatPersonalModelOverride so malformed values use one fallback.
|
||||
const ChatPersonalModelOverrideKeyPrefix = "chat_personal_model_override:"
|
||||
|
||||
// ChatPersonalModelOverrideKey returns the user config key for a chat
|
||||
// personal model override context. Values stored at the returned key should
|
||||
// use ParseChatPersonalModelOverride so malformed values fall back safely.
|
||||
func ChatPersonalModelOverrideKey(
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
) string {
|
||||
return ChatPersonalModelOverrideKeyPrefix + string(overrideContext)
|
||||
}
|
||||
|
||||
// ParsedChatPersonalModelOverride is a parsed personal model override value.
|
||||
// When Malformed is true, Mode is the provided default and ModelConfigID is
|
||||
// uuid.Nil.
|
||||
type ParsedChatPersonalModelOverride struct {
|
||||
Mode codersdk.ChatPersonalModelOverrideMode
|
||||
ModelConfigID uuid.UUID
|
||||
Malformed bool
|
||||
}
|
||||
|
||||
// ParseChatPersonalModelOverride parses a stored personal model override.
|
||||
// Empty values return defaultMode without marking the value malformed.
|
||||
// Malformed values return defaultMode, uuid.Nil, and Malformed true.
|
||||
func ParseChatPersonalModelOverride(
|
||||
raw string,
|
||||
defaultMode codersdk.ChatPersonalModelOverrideMode,
|
||||
) ParsedChatPersonalModelOverride {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ParsedChatPersonalModelOverride{Mode: defaultMode}
|
||||
}
|
||||
|
||||
switch trimmed {
|
||||
case string(codersdk.ChatPersonalModelOverrideModeChatDefault):
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
}
|
||||
case string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault):
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
}
|
||||
}
|
||||
|
||||
mode, rawModelConfigID, ok := strings.Cut(trimmed, ":")
|
||||
if !ok || mode != string(codersdk.ChatPersonalModelOverrideModeModel) {
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: defaultMode,
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
modelConfigID, err := uuid.Parse(rawModelConfigID)
|
||||
if err != nil {
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: defaultMode,
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func TestChatPersonalModelOverrideKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(
|
||||
t,
|
||||
"chat_personal_model_override:root",
|
||||
chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot),
|
||||
)
|
||||
}
|
||||
|
||||
func TestParseChatPersonalModelOverride(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.MustParse("11111111-1111-1111-1111-111111111111")
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
defaultMode codersdk.ChatPersonalModelOverrideMode
|
||||
want chatd.ParsedChatPersonalModelOverride
|
||||
}{
|
||||
{
|
||||
name: "EmptyUsesDefault",
|
||||
raw: "",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ChatDefault",
|
||||
raw: string(codersdk.ChatPersonalModelOverrideModeChatDefault),
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DeploymentDefault",
|
||||
raw: string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault),
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "Model",
|
||||
raw: "model:" + modelConfigID.String(),
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "InvalidModelUUID",
|
||||
raw: "model:not-a-uuid",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
Malformed: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "UnknownValue",
|
||||
raw: "unknown",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
Malformed: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "OuterWhitespace",
|
||||
raw: " \tmodel:" + modelConfigID.String() + "\n",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatd.ParseChatPersonalModelOverride(tt.raw, tt.defaultMode)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
+175
-5
@@ -140,6 +140,22 @@ func readSubagentModelOverride(
|
||||
}
|
||||
}
|
||||
|
||||
func personalModelOverrideContextForSubagent(
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (codersdk.ChatPersonalModelOverrideContext, error) {
|
||||
switch overrideContext {
|
||||
case codersdk.ChatModelOverrideContextGeneral:
|
||||
return codersdk.ChatPersonalModelOverrideContextGeneral, nil
|
||||
case codersdk.ChatModelOverrideContextExplore:
|
||||
return codersdk.ChatPersonalModelOverrideContextExplore, nil
|
||||
default:
|
||||
return "", xerrors.Errorf(
|
||||
"unknown subagent model override context %q",
|
||||
overrideContext,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func validateModelConfigAndResolveProvider(
|
||||
modelConfig database.ChatModelConfig,
|
||||
) (database.ChatModelConfig, string, error) {
|
||||
@@ -173,6 +189,15 @@ func enabledProviderContainsName(
|
||||
return false
|
||||
}
|
||||
|
||||
func userCanUseProviderKeys(
|
||||
providerKeys chatprovider.ProviderAPIKeys,
|
||||
providerName string,
|
||||
) bool {
|
||||
return providerKeys.APIKey(providerName) != "" ||
|
||||
(chatprovider.ProviderAllowsAmbientCredentials(providerName) &&
|
||||
providerKeys.HasProvider(providerName))
|
||||
}
|
||||
|
||||
type modelOverrideFailureMode int
|
||||
|
||||
const (
|
||||
@@ -274,9 +299,7 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
err,
|
||||
)
|
||||
}
|
||||
if providerKeys.APIKey(providerName) == "" &&
|
||||
!(chatprovider.ProviderAllowsAmbientCredentials(providerName) &&
|
||||
providerKeys.HasProvider(providerName)) {
|
||||
if !userCanUseProviderKeys(providerKeys, providerName) {
|
||||
if failureMode == modelOverrideFailureModeHard {
|
||||
return database.ChatModelConfig{}, true, xerrors.Errorf(
|
||||
"%s model override credentials are unavailable for provider %q",
|
||||
@@ -296,13 +319,160 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
return modelConfig, true, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (uuid.UUID, bool, error) {
|
||||
personalContext, err := personalModelOverrideContextForSubagent(overrideContext)
|
||||
if err != nil {
|
||||
return uuid.Nil, false, err
|
||||
}
|
||||
raw, err := p.db.GetUserChatPersonalModelOverride(
|
||||
ctx,
|
||||
database.GetUserChatPersonalModelOverrideParams{
|
||||
UserID: ownerID,
|
||||
Key: ChatPersonalModelOverrideKey(personalContext),
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
if !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return uuid.Nil, false, xerrors.Errorf(
|
||||
"get %s personal model override: %w",
|
||||
subagentModelOverrideLogLabel(overrideContext),
|
||||
err,
|
||||
)
|
||||
}
|
||||
raw = ""
|
||||
}
|
||||
|
||||
parsed := ParseChatPersonalModelOverride(
|
||||
raw,
|
||||
codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
)
|
||||
if parsed.Malformed {
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override is malformed, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("raw_model_config_id", strings.TrimSpace(raw)),
|
||||
)
|
||||
}
|
||||
switch parsed.Mode {
|
||||
case codersdk.ChatPersonalModelOverrideModeChatDefault:
|
||||
return uuid.Nil, true, nil
|
||||
case codersdk.ChatPersonalModelOverrideModeDeploymentDefault:
|
||||
case codersdk.ChatPersonalModelOverrideModeModel:
|
||||
modelConfig, ok, err := p.resolvePersonalModelOverride(
|
||||
ctx,
|
||||
overrideContext,
|
||||
ownerID,
|
||||
parsed.ModelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, false, err
|
||||
}
|
||||
if ok {
|
||||
return modelConfig.ID, true, nil
|
||||
}
|
||||
default:
|
||||
p.logger.Warn(ctx,
|
||||
"unsupported personal model override mode, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("mode", parsed.Mode),
|
||||
)
|
||||
}
|
||||
|
||||
return uuid.Nil, false, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolvePersonalModelOverride(
|
||||
ctx context.Context,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
ownerID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
) (database.ChatModelConfig, bool, error) {
|
||||
modelConfig, providerName, err := p.resolveModelConfigAndNormalizedProvider(
|
||||
ctx,
|
||||
modelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
switch {
|
||||
case xerrors.Is(err, sql.ErrNoRows):
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override is unavailable, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
)
|
||||
case errors.Is(err, errInvalidModelOverrideMetadata):
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override metadata is invalid, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
default:
|
||||
p.logger.Warn(ctx,
|
||||
"failed to resolve personal model override, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
}
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
}
|
||||
providerKeys, err := p.resolveUserProviderAPIKeys(ctx, ownerID)
|
||||
if err != nil {
|
||||
return database.ChatModelConfig{}, false, xerrors.Errorf(
|
||||
"resolve provider API keys: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
if !userCanUseProviderKeys(providerKeys, providerName) {
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override credentials are unavailable, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.F("provider", providerName),
|
||||
)
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
}
|
||||
return modelConfig, true, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolveSubagentModelConfigID(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (uuid.UUID, error) {
|
||||
//nolint:gocritic // Chatd needs its scoped deployment-config read access here.
|
||||
//nolint:gocritic // Chatd needs its scoped config and user-data access here.
|
||||
chatdCtx := dbauthz.AsChatd(ctx)
|
||||
personalOverridesEnabled, err := p.db.GetChatPersonalModelOverridesEnabled(chatdCtx)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf(
|
||||
"get chat personal model overrides enabled: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
if personalOverridesEnabled {
|
||||
modelConfigID, resolved, err := p.resolvePersonalSubagentModelConfigID(
|
||||
chatdCtx,
|
||||
ownerID,
|
||||
overrideContext,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, err
|
||||
}
|
||||
if resolved {
|
||||
return modelConfigID, nil
|
||||
}
|
||||
}
|
||||
|
||||
raw, err := readSubagentModelOverride(chatdCtx, p.db, overrideContext)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf(
|
||||
@@ -312,7 +482,7 @@ func (p *Server) resolveSubagentModelConfigID(
|
||||
)
|
||||
}
|
||||
modelConfig, ok, err := p.resolveConfiguredModelOverride(
|
||||
ctx,
|
||||
chatdCtx,
|
||||
string(overrideContext),
|
||||
raw,
|
||||
ownerID,
|
||||
|
||||
@@ -409,6 +409,46 @@ func chatdTestContext(t *testing.T) context.Context {
|
||||
return dbauthz.AsChatd(testutil.Context(t, testutil.WaitLong))
|
||||
}
|
||||
|
||||
func systemRestrictedTestContext(t *testing.T) context.Context {
|
||||
t.Helper()
|
||||
return dbauthz.AsSystemRestricted(testutil.Context(t, testutil.WaitLong))
|
||||
}
|
||||
|
||||
func enableInternalChatPersonalModelOverrides(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
) {
|
||||
t.Helper()
|
||||
require.NoError(
|
||||
t,
|
||||
db.UpsertChatPersonalModelOverridesEnabled(
|
||||
systemRestrictedTestContext(t),
|
||||
true,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
func upsertInternalUserChatPersonalModelOverride(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
raw string,
|
||||
) {
|
||||
t.Helper()
|
||||
require.NoError(
|
||||
t,
|
||||
db.UpsertUserChatPersonalModelOverride(
|
||||
systemRestrictedTestContext(t),
|
||||
database.UpsertUserChatPersonalModelOverrideParams{
|
||||
UserID: userID,
|
||||
Key: ChatPersonalModelOverrideKey(overrideContext),
|
||||
Value: raw,
|
||||
},
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -668,6 +708,204 @@ func TestSpawnAgent_GeneralUsesConfiguredModelOverride(t *testing.T) {
|
||||
require.False(t, childChat.PlanMode.Valid)
|
||||
}
|
||||
|
||||
func TestSpawnAgent_GeneralHonorsPersonalModelOverrides(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
enablePersonalOverride bool
|
||||
personalRaw func(database.ChatModelConfig) string
|
||||
personalModel func(context.Context, *testing.T, database.Store, uuid.UUID) database.ChatModelConfig
|
||||
wantModelID func(
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
) uuid.UUID
|
||||
}{
|
||||
{
|
||||
name: "UnsetUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DeploymentDefaultUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault)
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ChatDefaultBypassesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(parentModel, _, _ database.ChatModelConfig) uuid.UUID {
|
||||
return parentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ModelUsesPersonalOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, personalModel database.ChatModelConfig) uuid.UUID {
|
||||
return personalModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "AdminFlagOffIgnoresPersonalOverride",
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DisabledPersonalModelFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
return insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"general-personal-disabled-"+uuid.NewString(),
|
||||
false,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MissingCredentialsFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
insertInternalChatProvider(
|
||||
t,
|
||||
db,
|
||||
userID,
|
||||
"openai-compat",
|
||||
"",
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
return insertInternalChatModelConfigForProvider(
|
||||
t,
|
||||
db,
|
||||
"openai-compat",
|
||||
"gpt-4o-mini",
|
||||
true,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MalformedValueUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return "model:not-a-uuid"
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, parentModel := seedInternalChatDeps(t, db)
|
||||
deploymentModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"general-deployment-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, db.UpsertChatGeneralModelOverride(ctx, deploymentModel.ID.String()))
|
||||
personalModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"general-personal-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
if tt.personalModel != nil {
|
||||
personalModel = tt.personalModel(ctx, t, db, user.ID)
|
||||
}
|
||||
if tt.enablePersonalOverride {
|
||||
enableInternalChatPersonalModelOverrides(t, db)
|
||||
}
|
||||
if tt.personalRaw != nil {
|
||||
upsertInternalUserChatPersonalModelOverride(
|
||||
t,
|
||||
db,
|
||||
user.ID,
|
||||
codersdk.ChatPersonalModelOverrideContextGeneral,
|
||||
tt.personalRaw(personalModel),
|
||||
)
|
||||
}
|
||||
parentChat := createInternalParentChat(
|
||||
ctx,
|
||||
t,
|
||||
server,
|
||||
db,
|
||||
org.ID,
|
||||
user.ID,
|
||||
parentModel.ID,
|
||||
"parent-general-personal-override",
|
||||
)
|
||||
|
||||
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
|
||||
Type: subagentTypeGeneral,
|
||||
Prompt: "delegate general work",
|
||||
})
|
||||
childID := requireSpawnAgentChildChatID(t, resp)
|
||||
|
||||
childChat, err := db.GetChatByID(ctx, childID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(
|
||||
t,
|
||||
tt.wantModelID(parentModel, deploymentModel, personalModel),
|
||||
childChat.LastModelConfigID,
|
||||
)
|
||||
require.False(t, childChat.PlanMode.Valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -947,6 +1185,218 @@ func TestSpawnAgent_ExploreFallsBackToCurrentTurnModel(t *testing.T) {
|
||||
require.Equal(t, parentModel.ID, parentChat.LastModelConfigID)
|
||||
}
|
||||
|
||||
func TestSpawnAgent_ExploreHonorsPersonalModelOverrides(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
enablePersonalOverride bool
|
||||
personalRaw func(database.ChatModelConfig) string
|
||||
personalModel func(context.Context, *testing.T, database.Store, uuid.UUID) database.ChatModelConfig
|
||||
wantModelID func(
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
) uuid.UUID
|
||||
}{
|
||||
{
|
||||
name: "UnsetUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DeploymentDefaultUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault)
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ChatDefaultBypassesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(_, currentTurnModel, _, _ database.ChatModelConfig) uuid.UUID {
|
||||
return currentTurnModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ModelUsesPersonalOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, _, personalModel database.ChatModelConfig) uuid.UUID {
|
||||
return personalModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "AdminFlagOffIgnoresPersonalOverride",
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DisabledPersonalModelFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
return insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-personal-disabled-"+uuid.NewString(),
|
||||
false,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MissingCredentialsFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
insertInternalChatProvider(
|
||||
t,
|
||||
db,
|
||||
userID,
|
||||
"openai-compat",
|
||||
"",
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
return insertInternalChatModelConfigForProvider(
|
||||
t,
|
||||
db,
|
||||
"openai-compat",
|
||||
"gpt-4o-mini",
|
||||
true,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MalformedValueUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return "not-a-mode"
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, parentModel := seedInternalChatDeps(t, db)
|
||||
currentTurnModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-current-turn-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
deploymentModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-deployment-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, db.UpsertChatExploreModelOverride(ctx, deploymentModel.ID.String()))
|
||||
personalModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-personal-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
if tt.personalModel != nil {
|
||||
personalModel = tt.personalModel(ctx, t, db, user.ID)
|
||||
}
|
||||
if tt.enablePersonalOverride {
|
||||
enableInternalChatPersonalModelOverrides(t, db)
|
||||
}
|
||||
if tt.personalRaw != nil {
|
||||
upsertInternalUserChatPersonalModelOverride(
|
||||
t,
|
||||
db,
|
||||
user.ID,
|
||||
codersdk.ChatPersonalModelOverrideContextExplore,
|
||||
tt.personalRaw(personalModel),
|
||||
)
|
||||
}
|
||||
parentChat := createInternalParentChat(
|
||||
ctx,
|
||||
t,
|
||||
server,
|
||||
db,
|
||||
org.ID,
|
||||
user.ID,
|
||||
parentModel.ID,
|
||||
"parent-explore-personal-override",
|
||||
)
|
||||
|
||||
resp := runSubagentTool(
|
||||
ctx,
|
||||
t,
|
||||
server,
|
||||
parentChat,
|
||||
currentTurnModel.ID,
|
||||
spawnAgentToolName,
|
||||
spawnAgentArgs{Type: subagentTypeExplore, Prompt: "inspect the codebase"},
|
||||
)
|
||||
childID := requireSpawnAgentChildChatID(t, resp)
|
||||
|
||||
childChat, err := db.GetChatByID(ctx, childID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(
|
||||
t,
|
||||
tt.wantModelID(parentModel, currentTurnModel, deploymentModel, personalModel),
|
||||
childChat.LastModelConfigID,
|
||||
)
|
||||
require.True(t, childChat.Mode.Valid)
|
||||
require.Equal(t, database.ChatModeExplore, childChat.Mode.ChatMode)
|
||||
require.False(t, childChat.PlanMode.Valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateChat_ExploreRootStartsWithoutMCPSnapshot(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user