summaryrefslogtreecommitdiff
path: root/internal/controlplane/snapshot.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/controlplane/snapshot.go
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal/controlplane/snapshot.go')
-rw-r--r--internal/controlplane/snapshot.go108
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, &regionsJSON, &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, &regions); 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
}