mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
fix(openai): validate bulk account settings
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user