package controlplane import ( "context" "fmt" "strings" "time" "aigw/internal/domain" "github.com/jackc/pgx/v5" ) type UsageQuery struct { TenantID string ProjectID string Model string Limit int } func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error { tx, err := s.db.Begin(ctx) if err != nil { return fmt.Errorf("begin usage record: %w", err) } defer tx.Rollback(ctx) command, err := tx.Exec(ctx, ` INSERT INTO usage_events ( request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros) VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, event.TenantID, event.ProjectID, event.KeyID, event.PublicModel, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens) if err != nil { return fmt.Errorf("persist usage event: %w", err) } if command.RowsAffected() == 0 { return tx.Commit(ctx) } if err := upsertUsageRollup(ctx, tx, event, 0, 0, 0); err != nil { return err } if err := tx.Commit(ctx); err != nil { return fmt.Errorf("commit usage record: %w", err) } return nil } func upsertUsageRollup(ctx context.Context, tx pgx.Tx, event domain.UsageEvent, cost, charged, uncollected int64) error { period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC) _, err := tx.Exec(ctx, ` INSERT INTO usage_monthly_rollups (period_start, tenant_id, project_id, request_count, successful_requests, input_tokens, output_tokens, total_tokens, cost_micros, charged_micros, uncollected_micros) VALUES ($1,$2,$3,1,$4,$5,$6,$7,$8,$9,$10) ON CONFLICT (project_id, period_start) DO UPDATE SET request_count=usage_monthly_rollups.request_count+1, successful_requests=usage_monthly_rollups.successful_requests+EXCLUDED.successful_requests, input_tokens=usage_monthly_rollups.input_tokens+EXCLUDED.input_tokens, output_tokens=usage_monthly_rollups.output_tokens+EXCLUDED.output_tokens, total_tokens=usage_monthly_rollups.total_tokens+EXCLUDED.total_tokens, cost_micros=usage_monthly_rollups.cost_micros+EXCLUDED.cost_micros, charged_micros=usage_monthly_rollups.charged_micros+EXCLUDED.charged_micros, uncollected_micros=usage_monthly_rollups.uncollected_micros+EXCLUDED.uncollected_micros, updated_at=now()`, period, event.TenantID, event.ProjectID, boolInt(event.Success), event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, cost, charged, uncollected) if err != nil { return fmt.Errorf("update usage monthly rollup: %w", err) } return nil } func boolInt(value bool) int { if value { return 1 } return 0 } func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, error) { limit := query.Limit if limit < 1 || limit > 1000 { limit = 200 } where := []string{"1=1"} args := make([]any, 0, 5) index := 1 for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.Model, "public_model=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } args = append(args, limit) rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, key_id::text, public_model, COALESCE(provider_id,''), COALESCE(upstream_model,''), protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros FROM usage_events WHERE `+ strings.Join(where, " AND ")+` ORDER BY created_at DESC LIMIT $`+fmt.Sprint(index), args...) if err != nil { return nil, fmt.Errorf("query usage events: %w", err) } defer rows.Close() result := make([]UsageRecord, 0) for rows.Next() { var item UsageRecord if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.KeyID, &item.PublicModel, &item.ProviderID, &item.UpstreamModel, &item.Protocol, &item.Stream, &item.StatusCode, &item.Success, &item.ErrorType, &item.Attempts, &item.StartedAt, &item.DurationMS, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros); err != nil { return nil, fmt.Errorf("scan usage event: %w", err) } result = append(result, item) } return result, rows.Err() } func (s *Store) UsageSummary(ctx context.Context, tenantID, projectID string) ([]UsageSummary, error) { where := []string{"1=1"} args := make([]any, 0, 2) index := 1 for _, item := range []struct{ value, clause string }{{tenantID, "r.tenant_id=$"}, {projectID, "r.project_id=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } rows, err := s.db.Query(ctx, `SELECT r.period_start, r.tenant_id::text, r.project_id::text, p.name, r.request_count, r.successful_requests, r.input_tokens, r.output_tokens, r.total_tokens, r.cost_micros, r.charged_micros, r.uncollected_micros FROM usage_monthly_rollups r JOIN projects p ON p.id=r.project_id WHERE `+ strings.Join(where, " AND ")+` ORDER BY r.period_start DESC, p.name`, args...) if err != nil { return nil, fmt.Errorf("query usage summary: %w", err) } defer rows.Close() result := make([]UsageSummary, 0) for rows.Next() { var item UsageSummary if err := rows.Scan(&item.PeriodStart, &item.TenantID, &item.ProjectID, &item.ProjectName, &item.RequestCount, &item.SuccessfulRequests, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros); err != nil { return nil, fmt.Errorf("scan usage summary: %w", err) } result = append(result, item) } return result, rows.Err() }