diff options
Diffstat (limited to 'internal/controlplane/usage.go')
| -rw-r--r-- | internal/controlplane/usage.go | 151 |
1 files changed, 151 insertions, 0 deletions
diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go new file mode 100644 index 0000000..6c69a9f --- /dev/null +++ b/internal/controlplane/usage.go @@ -0,0 +1,151 @@ +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() +} |
