mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
Merge pull request #5581 from wucm667/fix/issue-5574-passthrough-model-discovery
fix(gateway): align passthrough model discovery
This commit is contained in:
@@ -602,6 +602,97 @@ func TestGetAvailableModels_ErrorAndGlobalListBranches(t *testing.T) {
|
||||
require.Equal(t, int64(1), okRepo.listAllCalls.Load())
|
||||
}
|
||||
|
||||
func TestGetAvailableModels_OpenAIPassthroughUsesDefaultFallback(t *testing.T) {
|
||||
groupID := int64(10)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
accounts []Account
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "passthrough only ignores stale mapping",
|
||||
accounts: []Account{
|
||||
{
|
||||
ID: 1,
|
||||
Platform: PlatformOpenAI,
|
||||
Credentials: map[string]any{"model_mapping": map[string]any{"stale-model": "upstream-model"}},
|
||||
Extra: map[string]any{"openai_passthrough": true},
|
||||
},
|
||||
},
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "passthrough wins over ordinary account mapping",
|
||||
accounts: []Account{
|
||||
{
|
||||
ID: 2,
|
||||
Platform: PlatformOpenAI,
|
||||
Credentials: map[string]any{"model_mapping": map[string]any{"configured-model": "configured-upstream"}},
|
||||
},
|
||||
{
|
||||
ID: 3,
|
||||
Platform: PlatformOpenAI,
|
||||
Credentials: map[string]any{"model_mapping": map[string]any{"stale-model": "upstream-model"}},
|
||||
Extra: map[string]any{"openai_passthrough": true},
|
||||
},
|
||||
},
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "ordinary accounts preserve mapped whitelist",
|
||||
accounts: []Account{
|
||||
{
|
||||
ID: 4,
|
||||
Platform: PlatformOpenAI,
|
||||
Credentials: map[string]any{"model_mapping": map[string]any{"configured-model": "configured-upstream"}},
|
||||
},
|
||||
},
|
||||
want: []string{"configured-model"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
repo := &modelsListAccountRepoStub{byGroup: map[int64][]Account{groupID: tt.accounts}}
|
||||
svc := &GatewayService{
|
||||
accountRepo: repo,
|
||||
modelsListCache: gocache.New(time.Minute, time.Minute),
|
||||
modelsListCacheTTL: time.Minute,
|
||||
}
|
||||
|
||||
require.Equal(t, tt.want, svc.GetAvailableModels(context.Background(), &groupID, PlatformOpenAI))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAvailableModels_GlobalListPreservesMappedModelsWithOpenAIPassthrough(t *testing.T) {
|
||||
groupID := int64(11)
|
||||
repo := &modelsListAccountRepoStub{
|
||||
byGroup: map[int64][]Account{
|
||||
groupID: {
|
||||
{
|
||||
ID: 1,
|
||||
Platform: PlatformOpenAI,
|
||||
Extra: map[string]any{"openai_passthrough": true},
|
||||
},
|
||||
{
|
||||
ID: 2,
|
||||
Platform: PlatformAnthropic,
|
||||
Credentials: map[string]any{"model_mapping": map[string]any{"claude-mapped": "claude-upstream"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
svc := &GatewayService{
|
||||
accountRepo: repo,
|
||||
modelsListCache: gocache.New(time.Minute, time.Minute),
|
||||
modelsListCacheTTL: time.Minute,
|
||||
}
|
||||
|
||||
require.Equal(t, []string{"claude-mapped"}, svc.GetAvailableModels(context.Background(), &groupID, ""))
|
||||
}
|
||||
|
||||
func TestGatewayHotpathHelpers_CacheTTLAndStickyContext(t *testing.T) {
|
||||
t.Run("resolve_user_group_rate_cache_ttl", func(t *testing.T) {
|
||||
require.Equal(t, defaultUserGroupRateCacheTTL, resolveUserGroupRateCacheTTL(nil))
|
||||
|
||||
@@ -1379,6 +1379,17 @@ func (s *GatewayService) GetAvailableModels(ctx context.Context, groupID *int64,
|
||||
hasAnyMapping := false
|
||||
|
||||
for _, acc := range accounts {
|
||||
// Passthrough routing accepts models independently of model_mapping. A stale
|
||||
// mapping on any eligible passthrough account therefore cannot define the
|
||||
// public whitelist; return nil so the handler uses its default model set.
|
||||
if platform == PlatformOpenAI && acc.IsOpenAIPassthroughEnabled() {
|
||||
if s.modelsListCache != nil {
|
||||
s.modelsListCache.Set(cacheKey, []string(nil), s.modelsListCacheTTL)
|
||||
modelsListCacheStoreTotal.Add(1)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
mapping := acc.GetModelMapping()
|
||||
if len(mapping) > 0 {
|
||||
hasAnyMapping = true
|
||||
|
||||
Reference in New Issue
Block a user