|
| 1 | +package service |
| 2 | + |
| 3 | +import ( |
| 4 | + "context" |
| 5 | + "errors" |
| 6 | + "fmt" |
| 7 | + "sort" |
| 8 | + "strings" |
| 9 | +) |
| 10 | + |
| 11 | +// AccountShareModels describes the current key's bound account and permitted |
| 12 | +// model IDs. A non-nil result with no models must remain an empty model list. |
| 13 | +type AccountShareModels struct { |
| 14 | + Account *Account |
| 15 | + Models []string |
| 16 | +} |
| 17 | + |
| 18 | +// GetAccountShareModels bypasses the ordinary group model cache: mode groups |
| 19 | +// contain accounts from multiple rooms, while permissions belong to each key. |
| 20 | +func (s *GatewayService) GetAccountShareModels(ctx context.Context, apiKey *APIKey, platform string) (*AccountShareModels, error) { |
| 21 | + return s.accountShareModeService.modelsForRequest(ctx, apiKey, platform, s.channelService) |
| 22 | +} |
| 23 | + |
| 24 | +func (s *OpenAIGatewayService) GetAccountShareModels(ctx context.Context, apiKey *APIKey) (*AccountShareModels, error) { |
| 25 | + return s.accountShareModeService.modelsForRequest(ctx, apiKey, PlatformOpenAI, s.channelService) |
| 26 | +} |
| 27 | + |
| 28 | +func (s *AccountShareModeService) modelsForRequest(ctx context.Context, apiKey *APIKey, platform string, channels *ChannelService) (*AccountShareModels, error) { |
| 29 | + if s == nil || apiKey == nil || apiKey.GroupID == nil { |
| 30 | + return nil, nil |
| 31 | + } |
| 32 | + isMode, err := s.IsModeGroupChecked(ctx, *apiKey.GroupID) |
| 33 | + if err != nil { |
| 34 | + return nil, fmt.Errorf("check account share model group: %w", err) |
| 35 | + } |
| 36 | + if !isMode { |
| 37 | + return nil, nil |
| 38 | + } |
| 39 | + if apiKey.UserID <= 0 || apiKey.ID <= 0 { |
| 40 | + return nil, ErrAccountShareModeGroupUnbound |
| 41 | + } |
| 42 | + // This read applies the member's effective terms without activating queued |
| 43 | + // rooms, renewing paid seats, touching idle time, or rebinding accounts. |
| 44 | + membership, listing, err := s.repo.GetActiveMembershipForRequest(ctx, apiKey.UserID, apiKey.ID, *apiKey.GroupID) |
| 45 | + if errors.Is(err, ErrAccountShareListingNotFound) { |
| 46 | + return nil, ErrAccountShareModeGroupUnbound |
| 47 | + } |
| 48 | + if err != nil { |
| 49 | + return nil, fmt.Errorf("read account share model binding: %w", err) |
| 50 | + } |
| 51 | + if membership == nil || listing == nil || membership.AccountID <= 0 { |
| 52 | + return nil, ErrAccountShareModeGroupUnbound |
| 53 | + } |
| 54 | + if s.accountRepo == nil { |
| 55 | + return nil, ErrServiceUnavailable |
| 56 | + } |
| 57 | + account, err := s.accountRepo.GetByID(ctx, membership.AccountID) |
| 58 | + if err != nil { |
| 59 | + return nil, fmt.Errorf("read account share model account: %w", err) |
| 60 | + } |
| 61 | + if account == nil || account.ID != membership.AccountID || account.Platform != platform || listing.Platform != platform { |
| 62 | + return nil, ErrAccountShareModeSelection |
| 63 | + } |
| 64 | + if s.pricedModelCatalog == nil { |
| 65 | + return nil, ErrOwnedAccountModelCatalogUnavailable |
| 66 | + } |
| 67 | + candidates, err := s.pricedModelCatalog.ListSelectablePricedModelIDs(ctx, PricedModelQuery{Platform: platform}) |
| 68 | + if err != nil { |
| 69 | + return nil, ErrOwnedAccountModelCatalogUnavailable.WithCause(err) |
| 70 | + } |
| 71 | + // Wildcard pricing can authorize concrete room/account models that the |
| 72 | + // selectable catalog cannot enumerate. Never expose the patterns themselves. |
| 73 | + candidates = append(append([]string(nil), candidates...), listing.AllowedModels...) |
| 74 | + for model := range account.GetModelMapping() { |
| 75 | + candidates = append(candidates, model) |
| 76 | + } |
| 77 | + result := &AccountShareModels{Account: account, Models: make([]string, 0, len(candidates))} |
| 78 | + for _, model := range normalizeAllowedModels(candidates) { |
| 79 | + if strings.ContainsAny(model, "*?") { |
| 80 | + continue |
| 81 | + } |
| 82 | + selectionModel, err := accountShareDiscoverySelectionModel(ctx, channels, *apiKey.GroupID, account, model) |
| 83 | + if err != nil { |
| 84 | + return nil, ErrOwnedAccountModelCatalogUnavailable.WithCause(err) |
| 85 | + } |
| 86 | + if !accountShareListingAllowsModel(listing, selectionModel) || !account.IsModelSupported(selectionModel) { |
| 87 | + continue |
| 88 | + } |
| 89 | + priced, err := s.pricedModelCatalog.IsModelPriced(ctx, PricedModelQuery{Platform: platform}, model) |
| 90 | + if err != nil { |
| 91 | + return nil, ErrOwnedAccountModelCatalogUnavailable.WithCause(err) |
| 92 | + } |
| 93 | + if !priced { |
| 94 | + continue |
| 95 | + } |
| 96 | + if selectionModel != model { |
| 97 | + priced, err = s.pricedModelCatalog.IsModelPriced(ctx, PricedModelQuery{Platform: platform}, selectionModel) |
| 98 | + if err != nil { |
| 99 | + return nil, ErrOwnedAccountModelCatalogUnavailable.WithCause(err) |
| 100 | + } |
| 101 | + if !priced { |
| 102 | + continue |
| 103 | + } |
| 104 | + } |
| 105 | + restricted, err := accountShareDiscoveryModelRestricted(ctx, channels, *apiKey.GroupID, account, selectionModel) |
| 106 | + if err != nil { |
| 107 | + return nil, ErrOwnedAccountModelCatalogUnavailable.WithCause(err) |
| 108 | + } |
| 109 | + if !restricted { |
| 110 | + result.Models = append(result.Models, model) |
| 111 | + } |
| 112 | + } |
| 113 | + sort.Strings(result.Models) |
| 114 | + return result, nil |
| 115 | +} |
| 116 | + |
| 117 | +// OpenAI-compatible handlers apply channel mapping before selecting the room |
| 118 | +// account. Anthropic's native handler checks the original requested model. |
| 119 | +func accountShareDiscoverySelectionModel(ctx context.Context, channels *ChannelService, groupID int64, account *Account, model string) (string, error) { |
| 120 | + if channels == nil || !account.IsOpenAICompatible() { |
| 121 | + return model, nil |
| 122 | + } |
| 123 | + mapping, err := channels.ResolveChannelMappingChecked(ctx, groupID, model) |
| 124 | + if err != nil { |
| 125 | + return "", err |
| 126 | + } |
| 127 | + if mapping.Mapped && strings.TrimSpace(mapping.MappedModel) != "" { |
| 128 | + return strings.TrimSpace(mapping.MappedModel), nil |
| 129 | + } |
| 130 | + return model, nil |
| 131 | +} |
| 132 | + |
| 133 | +// Use the same billing-model basis as dispatch, including account mappings |
| 134 | +// when the channel restricts upstream models. An unbound channel adds no limit. |
| 135 | +func accountShareDiscoveryModelRestricted(ctx context.Context, channels *ChannelService, groupID int64, account *Account, model string) (bool, error) { |
| 136 | + if channels == nil { |
| 137 | + return false, nil |
| 138 | + } |
| 139 | + mapping, err := channels.ResolveChannelMappingChecked(ctx, groupID, model) |
| 140 | + if err != nil { |
| 141 | + return false, err |
| 142 | + } |
| 143 | + billingModel := billingModelForRestriction(mapping.BillingModelSource, model, mapping.MappedModel) |
| 144 | + if mapping.BillingModelSource == BillingModelSourceUpstream { |
| 145 | + if account.IsOpenAICompatible() { |
| 146 | + billingModel = resolveOpenAIAccountUpstreamModelForRequest(account, model, false) |
| 147 | + } else { |
| 148 | + billingModel = resolveAccountUpstreamModel(account, model) |
| 149 | + } |
| 150 | + } |
| 151 | + if billingModel == "" { |
| 152 | + return false, nil |
| 153 | + } |
| 154 | + return channels.IsModelRestrictedChecked(ctx, groupID, billingModel) |
| 155 | +} |
0 commit comments