diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/controlplane/queries.go | 191 |
1 files changed, 180 insertions, 11 deletions
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go index 73d869f..48610c7 100644 --- a/internal/controlplane/queries.go +++ b/internal/controlplane/queries.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "time" ) func (s *Store) Overview(ctx context.Context) (Overview, error) { @@ -84,15 +85,35 @@ func (s *Store) ListAPIKeys(ctx context.Context) ([]APIKey, error) { } func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, error) { + periodStart := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + periodEnd := periodStart.AddDate(0, 1, 0) query := ` - SELECT id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at - FROM api_keys` - args := []any{} + SELECT k.id::text, k.tenant_id::text, k.project_id::text, k.name, k.key_prefix, k.scopes, + k.tags, k.monthly_spend_micros, k.status, k.expires_at, k.last_used_at, k.created_at, + usage.month_spend, usage.month_requests, pending.month_reserved, + COALESCE(( + SELECT jsonb_agg(m.public_id ORDER BY m.public_id) + FROM api_key_model_restrictions r + JOIN models m ON m.id = r.model_id + WHERE r.api_key_id = k.id + ), '[]'::jsonb) + FROM api_keys k + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(u.cost_micros), 0)::bigint AS month_spend, COUNT(*)::bigint AS month_requests + FROM usage_events u WHERE u.key_id = k.id AND u.started_at >= $1 AND u.started_at < $2 + ) usage + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(b.reserved_micros), 0)::bigint AS month_reserved + FROM billing_reservations b + WHERE b.key_id = k.id AND b.status IN ('pending', 'metering_failed') + AND b.created_at >= $1 AND b.created_at < $2 + ) pending` + args := []any{periodStart, periodEnd} if tenantID != "" { - query += ` WHERE tenant_id=$1` + query += ` WHERE k.tenant_id=$3` args = append(args, tenantID) } - query += ` ORDER BY created_at DESC` + query += ` ORDER BY k.created_at DESC` rows, err := s.db.Query(ctx, query, args...) if err != nil { return nil, fmt.Errorf("query API keys: %w", err) @@ -101,13 +122,22 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, result := make([]APIKey, 0) for rows.Next() { var item APIKey - var scopesJSON []byte - if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, &scopesJSON, &item.Status, &item.CreatedAt); err != nil { + var scopesJSON, tagsJSON, allowedModelsJSON []byte + if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, + &scopesJSON, &tagsJSON, &item.MonthlySpendMicros, &item.Status, &item.ExpiresAt, + &item.LastUsedAt, &item.CreatedAt, &item.CurrentMonthSpendMicros, &item.CurrentMonthRequests, + &item.CurrentMonthReservedMicros, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if err := json.Unmarshal(scopesJSON, &item.Scopes); err != nil { return nil, fmt.Errorf("decode API key scopes: %w", err) } + if err := json.Unmarshal(tagsJSON, &item.Tags); err != nil { + return nil, fmt.Errorf("decode API key tags: %w", err) + } + if err := json.Unmarshal(allowedModelsJSON, &item.AllowedModels); err != nil { + return nil, fmt.Errorf("decode API key model restrictions: %w", err) + } result = append(result, item) } return result, rows.Err() @@ -153,7 +183,7 @@ func (s *Store) ResourceTenantID(ctx context.Context, resource, id string) (stri func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) { rows, err := s.db.Query(ctx, ` - SELECT p.id::text, p.name, p.protocol, p.base_url, p.enabled, count(r.id), p.created_at + SELECT p.id::text, p.slug, p.name, p.protocol, p.wire_api, p.base_url, p.enabled, count(r.id), p.created_at FROM providers p LEFT JOIN model_routes r ON r.provider_id = p.id GROUP BY p.id ORDER BY p.created_at DESC`) if err != nil { @@ -163,7 +193,7 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) { result := make([]Provider, 0) for rows.Next() { var item Provider - if err := rows.Scan(&item.ID, &item.Name, &item.Protocol, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil { + if err := rows.Scan(&item.ID, &item.Slug, &item.Name, &item.Protocol, &item.WireAPI, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil { return nil, fmt.Errorf("scan provider: %w", err) } result = append(result, item) @@ -182,6 +212,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { m.enabled, m.created_at 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 ORDER BY m.public_id`) if err != nil { @@ -264,7 +295,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { keyRows.Close() routeRows, err := s.db.Query(ctx, ` - SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol, + SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol, p.wire_api, p.enabled, r.upstream_model, r.priority, r.weight, r.enabled FROM model_routes r JOIN providers p ON p.id = r.provider_id ORDER BY r.priority, r.created_at`) @@ -275,7 +306,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { for routeRows.Next() { var route Route var modelID string - if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil { + if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.WireAPI, &route.ProviderEnabled, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil { return nil, fmt.Errorf("scan model route: %w", err) } if position, ok := positions[modelID]; ok { @@ -284,3 +315,141 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { } return models, routeRows.Err() } + +func (s *Store) ListDeveloperModels(ctx context.Context, tenantID string) ([]DeveloperModel, error) { + models, err := s.ListModels(ctx) + if err != nil { + return nil, err + } + keyIDs := map[string]struct{}{} + if tenantID != "" { + rows, err := s.db.Query(ctx, `SELECT id::text FROM api_keys WHERE tenant_id=$1 AND status='active'`, tenantID) + if err != nil { + return nil, fmt.Errorf("query developer API keys: %w", err) + } + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + rows.Close() + return nil, err + } + keyIDs[id] = struct{}{} + } + if err := rows.Err(); err != nil { + rows.Close() + return nil, err + } + rows.Close() + } + return developerModelsFor(models, tenantID, keyIDs), nil +} + +func (s *Store) ListPublicModels(ctx context.Context) ([]PublicModel, error) { + models, err := s.ListModels(ctx) + if err != nil { + return nil, err + } + return publicModelsFor(models), nil +} + +func publicModelsFor(models []Model) []PublicModel { + result := make([]PublicModel, 0, len(models)) + for _, model := range models { + if !model.Enabled || model.Lifecycle == "retired" || len(model.AllowedTenantIDs) != 0 || len(model.AllowedKeyIDs) != 0 { + continue + } + wireSet := make(map[string]struct{}) + providerSet := make(map[string]struct{}) + wireAPIs := make([]string, 0, len(model.Routes)) + for _, route := range model.Routes { + if !route.Enabled || !route.ProviderEnabled { + continue + } + wireAPI := route.WireAPI + if wireAPI == "" && route.Protocol == "anthropic" { + wireAPI = "messages" + } else if wireAPI == "" { + wireAPI = "chat_completions" + } + if _, exists := wireSet[wireAPI]; !exists { + wireSet[wireAPI] = struct{}{} + wireAPIs = append(wireAPIs, wireAPI) + } + providerSet[route.ProviderID] = struct{}{} + } + if len(wireAPIs) == 0 { + continue + } + result = append(result, PublicModel{PublicID: model.PublicID, DisplayName: model.DisplayName, + Description: model.Description, OwnedBy: model.OwnedBy, InputModalities: model.InputModalities, + OutputModalities: model.OutputModalities, ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens, + Capabilities: model.Capabilities, Regions: model.Regions, Lifecycle: model.Lifecycle, ReleasedAt: model.ReleasedAt, + ReplacementModel: model.ReplacementModel, Aliases: model.Aliases, PriceCurrency: model.PriceCurrency, + InputPriceMicrosPerMillion: model.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: model.OutputPriceMicrosPerMillion, + CacheReadPriceMicrosPerMillion: model.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: model.CacheWritePriceMicrosPerMillion, + SupportedWireAPIs: wireAPIs, ProviderCount: len(providerSet), AvailableProviderCount: len(providerSet), HealthStatus: "available"}) + } + return result +} + +func developerModelsFor(models []Model, tenantID string, keyIDs map[string]struct{}) []DeveloperModel { + result := make([]DeveloperModel, 0, len(models)) + for _, model := range models { + if !model.Enabled || model.Lifecycle == "retired" || !stringAllowed(model.AllowedTenantIDs, tenantID) || !keyAllowed(model.AllowedKeyIDs, keyIDs) { + continue + } + wireSet := map[string]struct{}{} + wireAPIs := make([]string, 0, len(model.Routes)) + for _, route := range model.Routes { + if !route.Enabled || !route.ProviderEnabled { + continue + } + wireAPI := route.WireAPI + if wireAPI == "" && route.Protocol == "anthropic" { + wireAPI = "messages" + } else if wireAPI == "" { + wireAPI = "chat_completions" + } + if _, exists := wireSet[wireAPI]; !exists { + wireSet[wireAPI] = struct{}{} + wireAPIs = append(wireAPIs, wireAPI) + } + } + if len(wireAPIs) == 0 { + continue + } + result = append(result, DeveloperModel{ID: model.ID, PublicID: model.PublicID, DisplayName: model.DisplayName, + Description: model.Description, OwnedBy: model.OwnedBy, InputModalities: model.InputModalities, + OutputModalities: model.OutputModalities, ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens, + Capabilities: model.Capabilities, Regions: model.Regions, Lifecycle: model.Lifecycle, ReleasedAt: model.ReleasedAt, + ReplacementModel: model.ReplacementModel, Aliases: model.Aliases, PriceCurrency: model.PriceCurrency, + InputPriceMicrosPerMillion: model.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: model.OutputPriceMicrosPerMillion, + CacheReadPriceMicrosPerMillion: model.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: model.CacheWritePriceMicrosPerMillion, + SupportedWireAPIs: wireAPIs}) + } + return result +} + +func stringAllowed(allowed []string, value string) bool { + if len(allowed) == 0 || value == "" { + return true + } + for _, item := range allowed { + if item == value { + return true + } + } + return false +} + +func keyAllowed(allowed []string, keys map[string]struct{}) bool { + if len(allowed) == 0 { + return true + } + for _, id := range allowed { + if _, ok := keys[id]; ok { + return true + } + } + return false +} |
