diff options
Diffstat (limited to 'internal/catalog')
| -rw-r--r-- | internal/catalog/catalog.go | 21 | ||||
| -rw-r--r-- | internal/catalog/catalog_test.go | 10 |
2 files changed, 30 insertions, 1 deletions
diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go index 09a6fba..362d5b0 100644 --- a/internal/catalog/catalog.go +++ b/internal/catalog/catalog.go @@ -23,9 +23,15 @@ type snapshot struct { func New(cfg config.Config) *Catalog { providers := make(map[string]domain.Provider, len(cfg.Providers)) for _, provider := range cfg.Providers { + slug := provider.Slug + if slug == "" { + slug = provider.ID + } providers[provider.ID] = domain.Provider{ ID: provider.ID, + Slug: slug, Protocol: provider.Protocol, + WireAPI: provider.WireAPI, BaseURL: strings.TrimRight(provider.BaseURL, "/"), APIKey: provider.APIKey, } @@ -124,7 +130,7 @@ func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { result := make([]domain.Model, 0, len(current.list)) for _, model := range current.list { for _, route := range model.Routes { - if route.Provider.Protocol == protocol { + if protocolCompatible(route.Provider, protocol) { result = append(result, model) break } @@ -133,6 +139,19 @@ func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { return result } +func protocolCompatible(provider domain.Provider, requestProtocol domain.Protocol) bool { + switch requestProtocol { + case domain.ProtocolOpenAI: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions" + case domain.ProtocolOpenAIResponses: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses" + case domain.ProtocolAnthropic: + return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages" + default: + return false + } +} + func (c *Catalog) Count() int { current := c.state.Load() if current == nil { diff --git a/internal/catalog/catalog_test.go b/internal/catalog/catalog_test.go index 07fb71f..2ce8779 100644 --- a/internal/catalog/catalog_test.go +++ b/internal/catalog/catalog_test.go @@ -41,3 +41,13 @@ func TestCatalogReplaceCopiesRouteSlices(t *testing.T) { t.Fatal("catalog snapshot aliases the caller's route slice") } } + +func TestModelRestrictionIsAppliedBeforeCatalogAccess(t *testing.T) { + principal := domain.Principal{AllowedModels: map[string]struct{}{"model/allowed": {}}} + if !(domain.Model{ID: "model/allowed"}).Allows(principal) { + t.Fatal("allowed model was rejected") + } + if (domain.Model{ID: "model/other"}).Allows(principal) { + t.Fatal("model outside the API key restriction was allowed") + } +} |
