diff options
Diffstat (limited to 'internal/controlplane/snapshot.go')
| -rw-r--r-- | internal/controlplane/snapshot.go | 108 |
1 files changed, 102 insertions, 6 deletions
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go index bc04e7c..b9bddaf 100644 --- a/internal/controlplane/snapshot.go +++ b/internal/controlplane/snapshot.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "strings" + "time" "aigw/internal/auth" "aigw/internal/domain" @@ -102,13 +103,23 @@ func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) { rows, err := tx.Query(ctx, ` - SELECT m.public_id, m.owned_by, m.input_price_micros_per_million, m.output_price_micros_per_million, - m.cache_read_price_micros_per_million, m.cache_write_price_micros_per_million, + SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by, + m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens, + m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at, + m.retired_at, COALESCE(m.replacement_model,''), + pv.id::text, pv.version, pv.currency, pv.effective_from, + pv.input_price_micros_per_million, pv.output_price_micros_per_million, + pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million, r.provider_id::text, r.upstream_model, r.priority, r.weight FROM models m + JOIN LATERAL ( + SELECT * FROM model_price_versions v WHERE v.model_id=m.id + AND v.effective_from <= now() AND (v.effective_to IS NULL OR v.effective_to > now()) + ORDER BY v.effective_from DESC LIMIT 1 + ) pv ON TRUE JOIN model_routes r ON r.model_id = m.id AND r.enabled = TRUE JOIN providers p ON p.id = r.provider_id AND p.enabled = TRUE - WHERE m.enabled = TRUE + WHERE m.enabled = TRUE AND m.lifecycle <> 'retired' ORDER BY m.public_id, r.priority, r.created_at`) if err != nil { return nil, fmt.Errorf("query model routes: %w", err) @@ -117,10 +128,21 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov models := make([]domain.Model, 0) index := make(map[string]int) for rows.Next() { - var publicID, ownedBy, providerID, upstreamModel string + var modelID, publicID, displayName, description, ownedBy, replacement, providerID, upstreamModel string + var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte + var contextWindow, maxOutput int64 + var lifecycle, priceVersionID, priceCurrency string + var releasedAt, deprecatedAt, retiredAt *time.Time + var priceVersion int + var priceEffectiveFrom time.Time var inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64 var priority, weight int - if err := rows.Scan(&publicID, &ownedBy, &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice, &providerID, &upstreamModel, &priority, &weight); err != nil { + if err := rows.Scan(&modelID, &publicID, &displayName, &description, &ownedBy, + &inputModalitiesJSON, &outputModalitiesJSON, &contextWindow, &maxOutput, + &capabilitiesJSON, ®ionsJSON, &lifecycle, &releasedAt, &deprecatedAt, &retiredAt, &replacement, + &priceVersionID, &priceVersion, &priceCurrency, &priceEffectiveFrom, + &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice, + &providerID, &upstreamModel, &priority, &weight); err != nil { return nil, fmt.Errorf("scan model route: %w", err) } provider, ok := providers[providerID] @@ -129,13 +151,32 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov } position, exists := index[publicID] if !exists { + var inputModalities, outputModalities, capabilities, regions []string + if err := json.Unmarshal(inputModalitiesJSON, &inputModalities); err != nil { + return nil, fmt.Errorf("decode model input modalities: %w", err) + } + if err := json.Unmarshal(outputModalitiesJSON, &outputModalities); err != nil { + return nil, fmt.Errorf("decode model output modalities: %w", err) + } + if err := json.Unmarshal(capabilitiesJSON, &capabilities); err != nil { + return nil, fmt.Errorf("decode model capabilities: %w", err) + } + if err := json.Unmarshal(regionsJSON, ®ions); err != nil { + return nil, fmt.Errorf("decode model regions: %w", err) + } position = len(models) index[publicID] = position models = append(models, domain.Model{ - ID: publicID, OwnedBy: ownedBy, + ID: publicID, DisplayName: displayName, Description: description, OwnedBy: ownedBy, + InputModalities: inputModalities, OutputModalities: outputModalities, + ContextWindow: contextWindow, MaxOutputTokens: maxOutput, Capabilities: capabilities, + Regions: regions, Lifecycle: lifecycle, ReleasedAt: releasedAt, DeprecatedAt: deprecatedAt, + RetiredAt: retiredAt, ReplacementModel: replacement, PriceVersionID: priceVersionID, + PriceVersion: priceVersion, PriceCurrency: priceCurrency, PriceEffectiveFrom: priceEffectiveFrom, InputPriceMicrosPerMillion: inputPrice, OutputPriceMicrosPerMillion: outputPrice, CacheReadPriceMicrosPerMillion: cacheReadPrice, CacheWritePriceMicrosPerMillion: cacheWritePrice, }) + _ = modelID } models[position].Routes = append(models[position].Routes, domain.Route{ Provider: provider, UpstreamModel: upstreamModel, Priority: priority, Weight: weight, @@ -144,6 +185,61 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov if err := rows.Err(); err != nil { return nil, fmt.Errorf("read model routes: %w", err) } + byID := make(map[string]int, len(models)) + for i := range models { + byID[models[i].ID] = i + } + aliasRows, err := tx.Query(ctx, `SELECT a.alias, m.public_id FROM model_aliases a JOIN models m ON m.id=a.model_id`) + if err != nil { + return nil, fmt.Errorf("query model aliases: %w", err) + } + for aliasRows.Next() { + var alias, publicID string + if err := aliasRows.Scan(&alias, &publicID); err != nil { + aliasRows.Close() + return nil, err + } + if position, ok := byID[publicID]; ok { + models[position].Aliases = append(models[position].Aliases, alias) + } + } + aliasRows.Close() + tenantRows, err := tx.Query(ctx, `SELECT m.public_id, a.tenant_id::text FROM model_tenant_allowlist a JOIN models m ON m.id=a.model_id`) + if err != nil { + return nil, fmt.Errorf("query model tenant allowlist: %w", err) + } + for tenantRows.Next() { + var publicID, tenantID string + if err := tenantRows.Scan(&publicID, &tenantID); err != nil { + tenantRows.Close() + return nil, err + } + if position, ok := byID[publicID]; ok { + if models[position].AllowedTenantIDs == nil { + models[position].AllowedTenantIDs = map[string]struct{}{} + } + models[position].AllowedTenantIDs[tenantID] = struct{}{} + } + } + tenantRows.Close() + keyRows, err := tx.Query(ctx, `SELECT m.public_id, a.api_key_id::text FROM api_key_model_allowlist a JOIN models m ON m.id=a.model_id`) + if err != nil { + return nil, fmt.Errorf("query model key allowlist: %w", err) + } + for keyRows.Next() { + var publicID, keyID string + if err := keyRows.Scan(&publicID, &keyID); err != nil { + keyRows.Close() + return nil, err + } + if position, ok := byID[publicID]; ok { + if models[position].AllowedKeyIDs == nil { + models[position].AllowedKeyIDs = map[string]struct{}{} + } + models[position].AllowedKeyIDs[keyID] = struct{}{} + } + } + keyRows.Close() return models, nil } |
