summaryrefslogtreecommitdiff
path: root/internal/controlplane/queries.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/controlplane/queries.go')
-rw-r--r--internal/controlplane/queries.go191
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
+}