summaryrefslogtreecommitdiff
path: root/internal/controlplane/queries.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/controlplane/queries.go73
1 files changed, 68 insertions, 5 deletions
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go
index 9b76fad..73d869f 100644
--- a/internal/controlplane/queries.go
+++ b/internal/controlplane/queries.go
@@ -173,10 +173,17 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) {
func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
rows, err := s.db.Query(ctx, `
- SELECT id::text, public_id, owned_by, input_price_micros_per_million,
- output_price_micros_per_million, cache_read_price_micros_per_million,
- cache_write_price_micros_per_million, enabled, created_at
- FROM models ORDER BY public_id`)
+ 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,
+ m.enabled, m.created_at
+ FROM models m JOIN LATERAL (
+ SELECT * FROM model_price_versions v WHERE v.model_id=m.id
+ ORDER BY v.effective_from DESC LIMIT 1
+ ) pv ON TRUE ORDER BY m.public_id`)
if err != nil {
return nil, fmt.Errorf("query models: %w", err)
}
@@ -184,13 +191,23 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
positions := make(map[string]int)
for rows.Next() {
var item Model
- if err := rows.Scan(&item.ID, &item.PublicID, &item.OwnedBy, &item.InputPriceMicrosPerMillion,
+ var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte
+ if err := rows.Scan(&item.ID, &item.PublicID, &item.DisplayName, &item.Description, &item.OwnedBy,
+ &inputModalitiesJSON, &outputModalitiesJSON, &item.ContextWindow, &item.MaxOutputTokens,
+ &capabilitiesJSON, &regionsJSON, &item.Lifecycle, &item.ReleasedAt, &item.DeprecatedAt,
+ &item.RetiredAt, &item.ReplacementModel, &item.PriceVersionID, &item.PriceVersion,
+ &item.PriceCurrency, &item.PriceEffectiveFrom, &item.InputPriceMicrosPerMillion,
&item.OutputPriceMicrosPerMillion, &item.CacheReadPriceMicrosPerMillion,
&item.CacheWritePriceMicrosPerMillion, &item.Enabled, &item.CreatedAt); err != nil {
rows.Close()
return nil, fmt.Errorf("scan model: %w", err)
}
+ _ = json.Unmarshal(inputModalitiesJSON, &item.InputModalities)
+ _ = json.Unmarshal(outputModalitiesJSON, &item.OutputModalities)
+ _ = json.Unmarshal(capabilitiesJSON, &item.Capabilities)
+ _ = json.Unmarshal(regionsJSON, &item.Regions)
item.Routes = []Route{}
+ item.Aliases, item.AllowedTenantIDs, item.AllowedKeyIDs = []string{}, []string{}, []string{}
positions[item.ID] = len(models)
models = append(models, item)
}
@@ -200,6 +217,52 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
}
rows.Close()
+ aliasRows, err := s.db.Query(ctx, `SELECT model_id::text, alias FROM model_aliases ORDER BY alias`)
+ if err != nil {
+ return nil, fmt.Errorf("query model aliases: %w", err)
+ }
+ for aliasRows.Next() {
+ var modelID, alias string
+ if err := aliasRows.Scan(&modelID, &alias); err != nil {
+ aliasRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].Aliases = append(models[p].Aliases, alias)
+ }
+ }
+ aliasRows.Close()
+ tenantRows, err := s.db.Query(ctx, `SELECT model_id::text, tenant_id::text FROM model_tenant_allowlist`)
+ if err != nil {
+ return nil, fmt.Errorf("query tenant model allowlist: %w", err)
+ }
+ for tenantRows.Next() {
+ var modelID, tenantID string
+ if err := tenantRows.Scan(&modelID, &tenantID); err != nil {
+ tenantRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].AllowedTenantIDs = append(models[p].AllowedTenantIDs, tenantID)
+ }
+ }
+ tenantRows.Close()
+ keyRows, err := s.db.Query(ctx, `SELECT model_id::text, api_key_id::text FROM api_key_model_allowlist`)
+ if err != nil {
+ return nil, fmt.Errorf("query key model allowlist: %w", err)
+ }
+ for keyRows.Next() {
+ var modelID, keyID string
+ if err := keyRows.Scan(&modelID, &keyID); err != nil {
+ keyRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].AllowedKeyIDs = append(models[p].AllowedKeyIDs, keyID)
+ }
+ }
+ keyRows.Close()
+
routeRows, err := s.db.Query(ctx, `
SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol,
r.upstream_model, r.priority, r.weight, r.enabled