package catalog import ( "fmt" "sort" "strings" "sync/atomic" "aigw/internal/config" "aigw/internal/domain" ) type Catalog struct { state atomic.Pointer[snapshot] } type snapshot struct { models map[string]domain.Model aliases map[string]string list []domain.Model } 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, } } models := make([]domain.Model, 0, len(cfg.Models)) for _, modelCfg := range cfg.Models { model := domain.Model{ ID: modelCfg.ID, OwnedBy: modelCfg.OwnedBy, InputPriceMicrosPerMillion: modelCfg.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: modelCfg.OutputPriceMicrosPerMillion, CacheReadPriceMicrosPerMillion: modelCfg.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: modelCfg.CacheWritePriceMicrosPerMillion, } for _, route := range modelCfg.Routes { model.Routes = append(model.Routes, domain.Route{ Provider: providers[route.Provider], UpstreamModel: route.UpstreamModel, Priority: route.Priority, Weight: route.Weight, }) } models = append(models, model) } return NewModels(models) } func NewModels(models []domain.Model) *Catalog { catalog := &Catalog{} catalog.Replace(models) return catalog } func (c *Catalog) Replace(source []domain.Model) { models := make(map[string]domain.Model, len(source)) aliases := make(map[string]string) list := make([]domain.Model, 0, len(source)) for _, sourceModel := range source { model := sourceModel model.Routes = append([]domain.Route(nil), sourceModel.Routes...) models[model.ID] = model for _, alias := range model.Aliases { aliases[alias] = model.ID } list = append(list, model) } sort.Slice(list, func(i, j int) bool { return list[i].ID < list[j].ID }) c.state.Store(&snapshot{models: models, aliases: aliases, list: list}) } func (c *Catalog) Model(id string) (domain.Model, error) { current := c.state.Load() if current == nil { return domain.Model{}, fmt.Errorf("model %q not found", id) } model, ok := current.models[id] if !ok { if canonical, aliasOK := current.aliases[id]; aliasOK { model, ok = current.models[canonical] } } if !ok { return domain.Model{}, fmt.Errorf("model %q not found", id) } return model, nil } func (c *Catalog) ModelForPrincipal(id string, principal domain.Principal) (domain.Model, error) { model, err := c.Model(id) if err != nil { return domain.Model{}, err } if !model.Allows(principal) { return domain.Model{}, fmt.Errorf("model %q not allowed", id) } return model, nil } func (c *Catalog) ModelsFor(protocol domain.Protocol, principal domain.Principal) []domain.Model { models := c.Models(protocol) result := models[:0] for _, model := range models { if model.Allows(principal) { result = append(result, model) } } return result } func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { current := c.state.Load() if current == nil { return nil } result := make([]domain.Model, 0, len(current.list)) for _, model := range current.list { for _, route := range model.Routes { if protocolCompatible(route.Provider, protocol) { result = append(result, model) break } } } 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 { return 0 } return len(current.list) }