package controlplane import ( "context" "encoding/base64" "encoding/json" "errors" "fmt" "strings" "time" "aigw/internal/domain" "github.com/jackc/pgx/v5" ) type UsageQuery struct { TenantID string ProjectID string KeyID string Model string Provider string Protocol string ErrorType string Stream *bool RequestID string Status string From time.Time To time.Time Limit int Cursor string } var ErrInvalidUsageCursor = errors.New("invalid usage cursor") type usageCursor struct { StartedAt time.Time `json:"started_at"` RequestID string `json:"request_id"` } 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, ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,0,0,0,$22,$23) 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.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, event.UsageReported, usageMeteringStatus(event)) if err != nil { return fmt.Errorf("persist usage event: %w", err) } if _, err := tx.Exec(ctx, `UPDATE api_keys SET last_used_at = GREATEST(COALESCE(last_used_at, $2), $2) WHERE id = $1`, event.KeyID, event.StartedAt); err != nil { return fmt.Errorf("update API key last used time: %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 usageMeteringStatus(event domain.UsageEvent) string { if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success { return "upstream_failed" } if event.UsageReported { return "reported" } return "missing" } 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) (UsagePage, error) { limit := query.Limit if limit < 1 || limit > 1000 { limit = 200 } where := []string{"1=1"} args := make([]any, 0, 16) index := 1 for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.KeyID, "key_id=$"}, {query.Model, "public_model=$"}, {query.RequestID, "request_id=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } for _, item := range []struct{ value, clause string }{{query.Protocol, "protocol=$"}, {query.ErrorType, "error_type=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } if strings.TrimSpace(query.Provider) != "" { where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=usage_events.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")") args = append(args, query.Provider) index++ } if query.Stream != nil { where = append(where, "stream=$"+fmt.Sprint(index)) args = append(args, *query.Stream) index++ } if query.Status == "success" { where = append(where, "success=TRUE") } else if query.Status == "error" { where = append(where, "success=FALSE") } if !query.From.IsZero() { where = append(where, "started_at >= $"+fmt.Sprint(index)) args = append(args, query.From) index++ } if !query.To.IsZero() { where = append(where, "started_at < $"+fmt.Sprint(index)) args = append(args, query.To) index++ } if strings.TrimSpace(query.Cursor) != "" { cursor, err := decodeUsageCursor(query.Cursor) if err != nil { return UsagePage{}, err } where = append(where, "(started_at, request_id) < ($"+fmt.Sprint(index)+",$"+fmt.Sprint(index+1)+")") args = append(args, cursor.StartedAt, cursor.RequestID) index += 2 } args = append(args, limit+1) rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, COALESCE((SELECT name FROM projects p WHERE p.id=usage_events.project_id),''), key_id::text, COALESCE((SELECT name FROM api_keys k WHERE k.id=usage_events.key_id),''), public_model, COALESCE(provider_id,''), COALESCE((SELECT name FROM providers p WHERE p.id::text=usage_events.provider_id),''), COALESCE(upstream_model,''), protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status FROM usage_events WHERE `+ strings.Join(where, " AND ")+` ORDER BY started_at DESC, request_id DESC LIMIT $`+fmt.Sprint(index), args...) if err != nil { return UsagePage{}, fmt.Errorf("query usage events: %w", err) } defer rows.Close() result := make([]UsageRecord, 0, limit+1) for rows.Next() { var item UsageRecord if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.ProjectName, &item.KeyID, &item.KeyName, &item.PublicModel, &item.ProviderID, &item.ProviderName, &item.UpstreamModel, &item.Protocol, &item.Stream, &item.StatusCode, &item.Success, &item.ErrorType, &item.Attempts, &item.StartedAt, &item.DurationMS, &item.TTFTMS, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros, &item.UsageReported, &item.MeteringStatus); err != nil { return UsagePage{}, fmt.Errorf("scan usage event: %w", err) } result = append(result, item) } if err := rows.Err(); err != nil { return UsagePage{}, err } page := UsagePage{Data: result} if len(result) > limit { page.Data = result[:limit] page.NextCursor = encodeUsageCursor(page.Data[len(page.Data)-1]) } return page, nil } func encodeUsageCursor(record UsageRecord) string { payload, _ := json.Marshal(usageCursor{StartedAt: record.StartedAt.UTC(), RequestID: record.RequestID}) return base64.RawURLEncoding.EncodeToString(payload) } func decodeUsageCursor(raw string) (usageCursor, error) { if len(raw) > 2048 { return usageCursor{}, ErrInvalidUsageCursor } payload, err := base64.RawURLEncoding.DecodeString(raw) if err != nil { return usageCursor{}, ErrInvalidUsageCursor } var cursor usageCursor if json.Unmarshal(payload, &cursor) != nil || cursor.StartedAt.IsZero() || strings.TrimSpace(cursor.RequestID) == "" { return usageCursor{}, ErrInvalidUsageCursor } return cursor, nil } func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyPoint, error) { where := []string{"1=1"} args := make([]any, 0, 12) index := 1 for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.KeyID, "key_id=$"}, {query.Model, "public_model=$"}, {query.RequestID, "request_id=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } for _, item := range []struct{ value, clause string }{{query.Protocol, "protocol=$"}, {query.ErrorType, "error_type=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } if strings.TrimSpace(query.Provider) != "" { where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=usage_events.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")") args = append(args, query.Provider) index++ } if query.Stream != nil { where = append(where, "stream=$"+fmt.Sprint(index)) args = append(args, *query.Stream) index++ } if query.Status == "success" { where = append(where, "success=TRUE") } else if query.Status == "error" { where = append(where, "success=FALSE") } if !query.From.IsZero() { where = append(where, "started_at >= $"+fmt.Sprint(index)) args = append(args, query.From) index++ } if !query.To.IsZero() { where = append(where, "started_at < $"+fmt.Sprint(index)) args = append(args, query.To) index++ } rows, err := s.db.Query(ctx, `SELECT date_trunc('day', started_at AT TIME ZONE 'UTC'), count(*), count(*) FILTER (WHERE success), COALESCE(sum(input_tokens),0), COALESCE(sum(output_tokens),0), COALESCE(sum(total_tokens),0), COALESCE(sum(charged_micros),0), COALESCE(sum(uncollected_micros),0), COALESCE(round(avg(duration_ms)),0)::bigint, COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint, COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint, COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY ttft_ms) FILTER (WHERE ttft_ms > 0)),0)::bigint, COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY ttft_ms) FILTER (WHERE ttft_ms > 0)),0)::bigint FROM usage_events WHERE `+strings.Join(where, " AND ")+` GROUP BY 1 ORDER BY 1`, args...) if err != nil { return nil, fmt.Errorf("query daily usage: %w", err) } defer rows.Close() result := make([]UsageDailyPoint, 0) for rows.Next() { var item UsageDailyPoint if err := rows.Scan(&item.Day, &item.RequestCount, &item.SuccessfulRequests, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.ChargedMicros, &item.UncollectedMicros, &item.AverageDurationMS, &item.P50DurationMS, &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan daily usage: %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() }