From 76b70b1685e45e6c722a71887b7ff0e1dd5681ad Mon Sep 17 00:00:00 2001 From: lyy0709 <65712338+lyy0709@users.noreply.github.com> Date: Mon, 17 Aug 2026 13:34:41 +0800 Subject: [PATCH] fix(openai): validate bulk account settings --- .../account_long_context_billing_test.go | 12 +- backend/internal/service/admin_account.go | 38 ++- backend/internal/service/admin_service.go | 11 +- .../service/admin_service_bulk_update_test.go | 290 +++++++++++++++++- .../service/openai_bulk_account_settings.go | 225 ++++++++++++++ 5 files changed, 545 insertions(+), 31 deletions(-) create mode 100644 backend/internal/service/openai_bulk_account_settings.go diff --git a/backend/internal/service/account_long_context_billing_test.go b/backend/internal/service/account_long_context_billing_test.go index 709559d932..08bd3f0b3f 100644 --- a/backend/internal/service/account_long_context_billing_test.go +++ b/backend/internal/service/account_long_context_billing_test.go @@ -258,18 +258,20 @@ func TestAdminServiceBulkUpdateAccountsRejectsMalformedOpenAILongContextBillingV require.Zero(t, repo.bulkUpdateCalls) } -func TestAdminServiceBulkUpdateAccountsAllowsProviderOwnedValueForNonOpenAIAccounts(t *testing.T) { +func TestAdminServiceBulkUpdateAccountsRejectsOpenAILongContextKeyForNonOpenAIAccounts(t *testing.T) { repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformGrok}} svc := &adminServiceImpl{accountRepo: repo} result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ AccountIDs: []int64{1}, - Extra: map[string]any{openAILongContextBillingEnabledKey: []string{"provider-owned"}}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, }) - require.NoError(t, err) - require.NotNil(t, result) - require.Equal(t, 1, repo.bulkUpdateCalls) + require.Nil(t, result) + var appErr *infraerrors.ApplicationError + require.ErrorAs(t, err, &appErr) + require.Equal(t, "OPENAI_BULK_TARGET_INVALID", appErr.Reason) + require.Zero(t, repo.bulkUpdateCalls) } func TestAdminServiceBulkUpdateAccountsRejectsMalformedValueForMixedTargetsIncludingOpenAI(t *testing.T) { diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index 9958846658..77d9f35328 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -906,26 +906,36 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp return nil, err } } + openAISettings, err := normalizeBulkOpenAISettings(input) + if err != nil { + return nil, err + } needMixedChannelCheck := input.GroupIDs != nil && !input.SkipMixedChannelCheck - _, hasLongContextBillingUpdate := input.Extra[openAILongContextBillingEnabledKey] // 预取所有目标账号,供凭据守卫/代理守卫/混合渠道检查共用,避免多次 DB 查询。 var cachedTargets []*Account - if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || hasLongContextBillingUpdate || input.ProbeEnabled != nil || input.RateMultiplier != nil { + if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || openAISettings.any() || input.ProbeEnabled != nil || input.RateMultiplier != nil { loaded, err := s.accountRepo.GetByIDs(ctx, input.AccountIDs) if err != nil { return nil, err } cachedTargets = loaded } - if input.ProbeEnabled != nil { - targetsByID := make(map[int64]*Account, len(cachedTargets)) - for _, account := range cachedTargets { - if account != nil { - targetsByID[account.ID] = account - } + targetsByID := make(map[int64]*Account, len(cachedTargets)) + for _, account := range cachedTargets { + if account != nil { + targetsByID[account.ID] = account } + } + if openAISettings.any() { + inheritedCount, err := validateBulkOpenAISettingsTargets(input, openAISettings, targetsByID) + if err != nil { + return nil, err + } + result.LongContextInheritedCount = inheritedCount + } + if input.ProbeEnabled != nil { for _, accountID := range input.AccountIDs { account, ok := targetsByID[accountID] if !ok { @@ -936,18 +946,6 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp } } } - if hasLongContextBillingUpdate { - for _, account := range cachedTargets { - if account == nil || account.Platform != PlatformOpenAI { - continue - } - if err := ValidateOpenAILongContextBillingExtra(account.Platform, input.Extra); err != nil { - return nil, err - } - break - } - } - // 影子账号绝不持有凭据:批量更新携带凭据时,目标中不得含影子(外审 G5,与单账号 // UpdateAccount 守卫对齐)。覆盖显式 IDs 与 filter 解析出的 IDs(此处 AccountIDs 已解析完成)。 if len(input.Credentials) > 0 { diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 3030b945b9..be56680764 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -477,11 +477,12 @@ type UserGroupRPMStatus struct { // BulkUpdateAccountsResult is the aggregated response for bulk updates. type BulkUpdateAccountsResult struct { - Success int `json:"success"` - Failed int `json:"failed"` - SuccessIDs []int64 `json:"success_ids"` - FailedIDs []int64 `json:"failed_ids"` - Results []BulkUpdateAccountResult `json:"results"` + Success int `json:"success"` + Failed int `json:"failed"` + SuccessIDs []int64 `json:"success_ids"` + FailedIDs []int64 `json:"failed_ids"` + Results []BulkUpdateAccountResult `json:"results"` + LongContextInheritedCount int `json:"long_context_inherited_count,omitempty"` } type CreateProxyInput struct { diff --git a/backend/internal/service/admin_service_bulk_update_test.go b/backend/internal/service/admin_service_bulk_update_test.go index 206efb21f1..e935beaa98 100644 --- a/backend/internal/service/admin_service_bulk_update_test.go +++ b/backend/internal/service/admin_service_bulk_update_test.go @@ -18,6 +18,8 @@ type accountRepoStubForBulkUpdate struct { accountRepoStub bulkUpdateErr error bulkUpdateIDs []int64 + bulkUpdateCalls int + lastBulkUpdate AccountBulkUpdate bindGroupErrByID map[int64]error bindGroupsCalls []int64 bindGroupsByAccount map[int64][]int64 @@ -50,14 +52,23 @@ type accountRepoStubForBulkUpdate struct { } } -func (s *accountRepoStubForBulkUpdate) BulkUpdate(_ context.Context, ids []int64, _ AccountBulkUpdate) (int64, error) { +func (s *accountRepoStubForBulkUpdate) BulkUpdate(_ context.Context, ids []int64, updates AccountBulkUpdate) (int64, error) { + s.bulkUpdateCalls++ s.bulkUpdateIDs = append([]int64{}, ids...) + s.lastBulkUpdate = updates if s.bulkUpdateErr != nil { return 0, s.bulkUpdateErr } return int64(len(ids)), nil } +func requireApplicationErrorReason(t *testing.T, err error, reason string) { + t.Helper() + var appErr *infraerrors.ApplicationError + require.ErrorAs(t, err, &appErr) + require.Equal(t, reason, appErr.Reason) +} + func (s *accountRepoStubForBulkUpdate) Create(_ context.Context, account *Account) error { s.createAccount = account if s.createID > 0 { @@ -307,3 +318,280 @@ func TestAdminServiceBulkUpdateAccounts_ResolvesIDsFromFilters(t *testing.T) { require.Equal(t, 0, result.Failed) require.Equal(t, []int64{7, 11}, result.SuccessIDs) } + +func TestAdminServiceBulkUpdateAccounts_NormalizesOpenAISettings(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{ + {ID: 1, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, + {ID: 2, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1, 2}, + Credentials: map[string]any{ + openAIEndpointCapabilitiesCredentialKey: []any{"chat_completions", "embeddings"}, + }, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: true, + "openai_responses_mode": "auto", + }, + }) + + require.NoError(t, err) + require.Equal(t, 2, result.Success) + require.Zero(t, result.LongContextInheritedCount) + require.Equal(t, 1, repo.bulkUpdateCalls) + require.Contains(t, repo.lastBulkUpdate.Credentials, openAIEndpointCapabilitiesCredentialKey) + require.Nil(t, repo.lastBulkUpdate.Credentials[openAIEndpointCapabilitiesCredentialKey]) + require.Equal(t, true, repo.lastBulkUpdate.Extra[openAILongContextBillingEnabledKey]) + require.Contains(t, repo.lastBulkUpdate.Extra, "openai_responses_mode") + require.Nil(t, repo.lastBulkUpdate.Extra["openai_responses_mode"]) +} + +func TestAdminServiceBulkUpdateAccounts_AcceptsLongContextAccountTypes(t *testing.T) { + for _, accountType := range []string{AccountTypeOAuth, AccountTypeSetupToken, AccountTypeAPIKey} { + t.Run(accountType, func(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{{ + ID: 1, Platform: PlatformOpenAI, Type: accountType, + }}} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: false}, + }) + + require.NoError(t, err) + require.Equal(t, 1, result.Success) + require.Equal(t, 1, repo.bulkUpdateCalls) + }) + } +} + +func TestAdminServiceBulkUpdateAccounts_EmbeddingsOnlyResetsResponsesMode(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{ + {ID: 1, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + _, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Credentials: map[string]any{ + openAIEndpointCapabilitiesCredentialKey: []string{"embeddings"}, + }, + }) + + require.NoError(t, err) + require.Equal(t, []string{"embeddings"}, repo.lastBulkUpdate.Credentials[openAIEndpointCapabilitiesCredentialKey]) + require.Contains(t, repo.lastBulkUpdate.Extra, "openai_responses_mode") + require.Nil(t, repo.lastBulkUpdate.Extra["openai_responses_mode"]) +} + +func TestAdminServiceBulkUpdateAccounts_RejectsInvalidOpenAISettingValuesBeforeWrite(t *testing.T) { + tests := []struct { + name string + credentials map[string]any + extra map[string]any + reason string + }{ + {name: "long context type", extra: map[string]any{openAILongContextBillingEnabledKey: "true"}, reason: "OPENAI_LONG_CONTEXT_BILLING_INVALID"}, + {name: "empty capabilities", credentials: map[string]any{openAIEndpointCapabilitiesCredentialKey: []any{}}, reason: "OPENAI_ENDPOINT_CAPABILITIES_INVALID"}, + {name: "unknown capability", credentials: map[string]any{openAIEndpointCapabilitiesCredentialKey: []any{"responses"}}, reason: "OPENAI_ENDPOINT_CAPABILITIES_INVALID"}, + {name: "capabilities type", credentials: map[string]any{openAIEndpointCapabilitiesCredentialKey: "chat_completions"}, reason: "OPENAI_ENDPOINT_CAPABILITIES_INVALID"}, + {name: "responses mode", extra: map[string]any{"openai_responses_mode": "sometimes"}, reason: "OPENAI_RESPONSES_MODE_INVALID"}, + {name: "responses type", extra: map[string]any{"openai_responses_mode": true}, reason: "OPENAI_RESPONSES_MODE_INVALID"}, + { + name: "embeddings conflict", + credentials: map[string]any{openAIEndpointCapabilitiesCredentialKey: []any{"embeddings"}}, + extra: map[string]any{"openai_responses_mode": "force_responses"}, + reason: "OPENAI_RESPONSES_MODE_INVALID", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{} + svc := &adminServiceImpl{accountRepo: repo} + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Credentials: tt.credentials, + Extra: tt.extra, + }) + require.Nil(t, result) + requireApplicationErrorReason(t, err, tt.reason) + require.Zero(t, repo.bulkUpdateCalls) + }) + } +} + +func TestAdminServiceBulkUpdateAccounts_RejectsInvalidOpenAITargetsBeforeWrite(t *testing.T) { + tests := []struct { + name string + accounts []*Account + input *BulkUpdateAccountsInput + }{ + { + name: "missing account", + accounts: []*Account{{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}}, + input: &BulkUpdateAccountsInput{ + AccountIDs: []int64{1, 2}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, + }, + }, + { + name: "mixed platform long context", + accounts: []*Account{{ID: 1, Platform: PlatformAnthropic, Type: AccountTypeOAuth}}, + input: &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, + }, + }, + { + name: "oauth endpoint capabilities", + accounts: []*Account{{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}}, + input: &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Credentials: map[string]any{openAIEndpointCapabilitiesCredentialKey: nil}, + }, + }, + { + name: "unsupported OpenAI long context account type", + accounts: []*Account{{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeServiceAccount}}, + input: &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: tt.accounts} + svc := &adminServiceImpl{accountRepo: repo} + result, err := svc.BulkUpdateAccounts(context.Background(), tt.input) + require.Nil(t, result) + requireApplicationErrorReason(t, err, "OPENAI_BULK_TARGET_INVALID") + require.Zero(t, repo.bulkUpdateCalls) + }) + } +} + +func TestAdminServiceBulkUpdateAccounts_ForcedResponsesRequiresChatCapability(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + openAIEndpointCapabilitiesCredentialKey: []any{"embeddings"}, + }, + }}} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{"openai_responses_mode": "force_chat_completions"}, + }) + + require.Nil(t, result) + requireApplicationErrorReason(t, err, "OPENAI_BULK_TARGET_INVALID") + require.Zero(t, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccounts_ForcedResponsesAcceptsChatCapabilityUpdate(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + openAIEndpointCapabilitiesCredentialKey: []any{"embeddings"}, + }, + }}} + svc := &adminServiceImpl{accountRepo: repo} + + _, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Credentials: map[string]any{ + openAIEndpointCapabilitiesCredentialKey: []any{"chat_completions"}, + }, + Extra: map[string]any{"openai_responses_mode": "force_responses"}, + }) + + require.NoError(t, err) + require.Equal(t, 1, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccounts_ReportsLongContextShadowInheritance(t *testing.T) { + parentID := int64(1) + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{ + {ID: parentID, Platform: PlatformOpenAI, Type: AccountTypeOAuth}, + {ID: 2, Platform: PlatformOpenAI, Type: AccountTypeOAuth, ParentAccountID: &parentID}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{parentID, 2}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, + }) + + require.NoError(t, err) + require.Equal(t, 1, result.LongContextInheritedCount) + require.Equal(t, 1, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccounts_RequiresParentForShadowOnlyLongContextUpdate(t *testing.T) { + parentID := int64(10) + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{ + {ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, ParentAccountID: &parentID}, + {ID: 2, Platform: PlatformOpenAI, Type: AccountTypeOAuth, ParentAccountID: &parentID}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1, 2}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, + }) + + require.Nil(t, result) + requireApplicationErrorReason(t, err, "OPENAI_LONG_CONTEXT_PARENT_REQUIRED") + require.Zero(t, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccounts_ShadowLongContextAllowsOtherUpdates(t *testing.T) { + parentID := int64(10) + repo := &accountRepoStubForBulkUpdate{getByIDsAccounts: []*Account{{ + ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, ParentAccountID: &parentID, + }}} + svc := &adminServiceImpl{accountRepo: repo} + status := StatusDisabled + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Status: status, + Extra: map[string]any{openAILongContextBillingEnabledKey: false}, + }) + + require.NoError(t, err) + require.Equal(t, 1, result.LongContextInheritedCount) + require.Equal(t, 1, repo.bulkUpdateCalls) + require.NotNil(t, repo.lastBulkUpdate.Status) + require.Equal(t, status, *repo.lastBulkUpdate.Status) +} + +func TestAdminServiceBulkUpdateAccounts_ValidatesFilterResolvedOpenAITargets(t *testing.T) { + repo := &accountRepoStubForBulkUpdate{ + listData: []Account{{ID: 7}}, + listResult: &pagination.PaginationResult{Total: 1}, + getByIDsAccounts: []*Account{{ID: 7, Platform: PlatformAnthropic, Type: AccountTypeOAuth}}, + } + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + Filters: &BulkUpdateAccountFilters{Platform: PlatformOpenAI}, + Extra: map[string]any{openAILongContextBillingEnabledKey: true}, + }) + + require.Nil(t, result) + requireApplicationErrorReason(t, err, "OPENAI_BULK_TARGET_INVALID") + require.Equal(t, []int64{7}, repo.getByIDsIDs) + require.Zero(t, repo.bulkUpdateCalls) +} diff --git a/backend/internal/service/openai_bulk_account_settings.go b/backend/internal/service/openai_bulk_account_settings.go new file mode 100644 index 0000000000..f418e575f4 --- /dev/null +++ b/backend/internal/service/openai_bulk_account_settings.go @@ -0,0 +1,225 @@ +package service + +import ( + "fmt" + "strconv" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" +) + +type bulkOpenAISettings struct { + longContextBilling bool + endpointCapabilities bool + responsesMode bool + capabilitiesIncludeChat bool + forcedResponsesMode bool +} + +func (s bulkOpenAISettings) any() bool { + return s.longContextBilling || s.endpointCapabilities || s.responsesMode +} + +func normalizeBulkOpenAISettings(input *BulkUpdateAccountsInput) (bulkOpenAISettings, error) { + var settings bulkOpenAISettings + if input == nil { + return settings, nil + } + + if _, exists := input.Extra[openAILongContextBillingEnabledKey]; exists { + settings.longContextBilling = true + if err := ValidateOpenAILongContextBillingExtra(PlatformOpenAI, input.Extra); err != nil { + return settings, err + } + } + + if raw, exists := input.Credentials[openAIEndpointCapabilitiesCredentialKey]; exists { + settings.endpointCapabilities = true + capabilities, includeChat, err := normalizeBulkOpenAIEndpointCapabilities(raw) + if err != nil { + return settings, err + } + settings.capabilitiesIncludeChat = includeChat + input.Credentials[openAIEndpointCapabilitiesCredentialKey] = capabilities + } + + if raw, exists := input.Extra[openai_compat.ExtraKeyResponsesMode]; exists { + settings.responsesMode = true + mode, forced, err := normalizeBulkOpenAIResponsesMode(raw) + if err != nil { + return settings, err + } + settings.forcedResponsesMode = forced + input.Extra[openai_compat.ExtraKeyResponsesMode] = mode + } + + if settings.endpointCapabilities && !settings.capabilitiesIncludeChat { + if settings.forcedResponsesMode { + return settings, infraerrors.BadRequest( + "OPENAI_RESPONSES_MODE_INVALID", + "a forced Responses route requires the chat_completions endpoint capability", + ) + } + if input.Extra == nil { + input.Extra = make(map[string]any, 1) + } + input.Extra[openai_compat.ExtraKeyResponsesMode] = nil + settings.responsesMode = true + } + + return settings, nil +} + +func normalizeBulkOpenAIEndpointCapabilities(raw any) (any, bool, error) { + if raw == nil { + return nil, true, nil + } + + values := make([]string, 0, 2) + switch typed := raw.(type) { + case []any: + for _, item := range typed { + value, ok := item.(string) + if !ok { + return nil, false, invalidBulkOpenAIEndpointCapabilities() + } + values = append(values, value) + } + case []string: + values = append(values, typed...) + default: + return nil, false, invalidBulkOpenAIEndpointCapabilities() + } + + selected := make(map[string]bool, 2) + for _, value := range values { + switch OpenAIEndpointCapability(value) { + case OpenAIEndpointCapabilityChatCompletions, OpenAIEndpointCapabilityEmbeddings: + selected[value] = true + default: + return nil, false, invalidBulkOpenAIEndpointCapabilities() + } + } + if len(selected) == 0 { + return nil, false, invalidBulkOpenAIEndpointCapabilities() + } + + includeChat := selected[string(OpenAIEndpointCapabilityChatCompletions)] + if includeChat && selected[string(OpenAIEndpointCapabilityEmbeddings)] { + return nil, true, nil + } + if includeChat { + return []string{string(OpenAIEndpointCapabilityChatCompletions)}, true, nil + } + return []string{string(OpenAIEndpointCapabilityEmbeddings)}, false, nil +} + +func invalidBulkOpenAIEndpointCapabilities() error { + return infraerrors.BadRequest( + "OPENAI_ENDPOINT_CAPABILITIES_INVALID", + "openai_capabilities must contain chat_completions, embeddings, or both", + ) +} + +func normalizeBulkOpenAIResponsesMode(raw any) (any, bool, error) { + if raw == nil { + return nil, false, nil + } + mode, ok := raw.(string) + if !ok { + return nil, false, invalidBulkOpenAIResponsesMode() + } + switch openai_compat.ResponsesSupportMode(mode) { + case openai_compat.ResponsesSupportModeAuto: + return nil, false, nil + case openai_compat.ResponsesSupportModeForceResponses, + openai_compat.ResponsesSupportModeForceChatCompletions: + return mode, true, nil + default: + return nil, false, invalidBulkOpenAIResponsesMode() + } +} + +func invalidBulkOpenAIResponsesMode() error { + return infraerrors.BadRequest( + "OPENAI_RESPONSES_MODE_INVALID", + "openai_responses_mode must be auto, force_responses, force_chat_completions, or null", + ) +} + +func validateBulkOpenAISettingsTargets( + input *BulkUpdateAccountsInput, + settings bulkOpenAISettings, + targetsByID map[int64]*Account, +) (int, error) { + if input == nil || !settings.any() { + return 0, nil + } + + inheritedCount := 0 + for _, accountID := range input.AccountIDs { + account, ok := targetsByID[accountID] + if !ok || account == nil { + return 0, invalidBulkOpenAITarget(accountID, "account does not exist") + } + + if settings.longContextBilling { + if account.Platform != PlatformOpenAI || !supportsOpenAILongContextBilling(account.Type) { + return 0, invalidBulkOpenAITarget(accountID, "long-context billing requires an OpenAI OAuth, setup-token, or API-key account") + } + if account.IsShadow() { + inheritedCount++ + } + } + + if settings.endpointCapabilities || settings.responsesMode { + if account.Platform != PlatformOpenAI || account.Type != AccountTypeAPIKey { + return 0, invalidBulkOpenAITarget(accountID, "endpoint capabilities and Responses routing require an OpenAI API-key account") + } + } + + if settings.forcedResponsesMode && !settings.capabilitiesIncludeChat && + !settings.endpointCapabilities && + !account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions) { + return 0, invalidBulkOpenAITarget(accountID, "a forced Responses route requires the chat_completions endpoint capability") + } + } + + if settings.longContextBilling && inheritedCount == len(input.AccountIDs) && bulkUpdateOnlyChangesLongContext(input) { + return 0, infraerrors.BadRequest( + "OPENAI_LONG_CONTEXT_PARENT_REQUIRED", + "long-context billing is owned by parent accounts; select at least one parent account", + ) + } + return inheritedCount, nil +} + +func supportsOpenAILongContextBilling(accountType string) bool { + switch accountType { + case AccountTypeOAuth, AccountTypeSetupToken, AccountTypeAPIKey: + return true + default: + return false + } +} + +func invalidBulkOpenAITarget(accountID int64, message string) error { + return infraerrors.BadRequest( + "OPENAI_BULK_TARGET_INVALID", + fmt.Sprintf("account %d: %s", accountID, message), + ).WithMetadata(map[string]string{"account_id": strconv.FormatInt(accountID, 10)}) +} + +func bulkUpdateOnlyChangesLongContext(input *BulkUpdateAccountsInput) bool { + if input == nil || input.Name != "" || input.ProxyID != nil || input.Concurrency != nil || + input.Priority != nil || input.RateMultiplier != nil || input.LoadFactor != nil || + input.Status != "" || input.Schedulable != nil || input.GroupIDs != nil || + len(input.Credentials) != 0 || input.ProbeEnabled != nil { + return false + } + if len(input.Extra) != 1 { + return false + } + _, ok := input.Extra[openAILongContextBillingEnabledKey] + return ok +}