summaryrefslogtreecommitdiff
path: root/internal/controlplane/usage.go
blob: 6c69a9fc868b4b40853fe9203f17faa86628de50 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
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()
}