diff options
| author | Chia <Chia@93.nz> | 2026-08-04 19:58:52 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-04 20:43:23 +1200 |
| commit | 5b651488b081b65fda8a323f228e139adb79a35d (patch) | |
| tree | 08baf40efb8fe103b32721cd991ff712323e3173 /internal/catalog | |
Build AI gateway control plane and admin UI
Diffstat (limited to 'internal/catalog')
| -rw-r--r-- | internal/catalog/catalog.go | 103 | ||||
| -rw-r--r-- | internal/catalog/catalog_test.go | 43 |
2 files changed, 146 insertions, 0 deletions
diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go new file mode 100644 index 0000000..76e33b1 --- /dev/null +++ b/internal/catalog/catalog.go @@ -0,0 +1,103 @@ +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 + list []domain.Model +} + +func New(cfg config.Config) *Catalog { + providers := make(map[string]domain.Provider, len(cfg.Providers)) + for _, provider := range cfg.Providers { + providers[provider.ID] = domain.Provider{ + ID: provider.ID, + Protocol: provider.Protocol, + 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} + 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)) + 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 + 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, 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 { + return domain.Model{}, fmt.Errorf("model %q not found", id) + } + return model, nil +} + +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 route.Provider.Protocol == protocol { + result = append(result, model) + break + } + } + } + return result +} + +func (c *Catalog) Count() int { + current := c.state.Load() + if current == nil { + return 0 + } + return len(current.list) +} diff --git a/internal/catalog/catalog_test.go b/internal/catalog/catalog_test.go new file mode 100644 index 0000000..07fb71f --- /dev/null +++ b/internal/catalog/catalog_test.go @@ -0,0 +1,43 @@ +package catalog + +import ( + "testing" + + "aigw/internal/domain" +) + +func TestCatalogReplaceSwapsModelSnapshot(t *testing.T) { + openAI := domain.Provider{ID: "openai", Protocol: domain.ProtocolOpenAI} + anthropic := domain.Provider{ID: "anthropic", Protocol: domain.ProtocolAnthropic} + catalog := NewModels([]domain.Model{{ + ID: "old", Routes: []domain.Route{{Provider: openAI, UpstreamModel: "old-upstream"}}, + }}) + + catalog.Replace([]domain.Model{{ + ID: "new", Routes: []domain.Route{{Provider: anthropic, UpstreamModel: "new-upstream"}}, + }}) + if _, err := catalog.Model("old"); err == nil { + t.Fatal("old model remained after snapshot replacement") + } + model, err := catalog.Model("new") + if err != nil || len(model.Routes) != 1 || model.Routes[0].Provider.ID != "anthropic" { + t.Fatalf("new model was not loaded: model=%+v err=%v", model, err) + } + if got := catalog.Models(domain.ProtocolOpenAI); len(got) != 0 { + t.Fatalf("unexpected OpenAI models after replacement: %+v", got) + } +} + +func TestCatalogReplaceCopiesRouteSlices(t *testing.T) { + models := []domain.Model{{ID: "model", Routes: []domain.Route{{UpstreamModel: "before"}}}} + catalog := NewModels(models) + models[0].Routes[0].UpstreamModel = "after" + + model, err := catalog.Model("model") + if err != nil { + t.Fatal(err) + } + if model.Routes[0].UpstreamModel != "before" { + t.Fatal("catalog snapshot aliases the caller's route slice") + } +} |
