diff options
Diffstat (limited to 'internal/catalog/catalog.go')
| -rw-r--r-- | internal/catalog/catalog.go | 21 |
1 files changed, 20 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 { |
