summaryrefslogtreecommitdiff
path: root/internal/controlplane/limits.go
blob: 5b5825e41d7748f97ce4069a699fca77f4b3d42c (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
package controlplane

import (
	"context"
	"errors"
	"fmt"
	"strings"

	"github.com/jackc/pgx/v5"
)

func (s *Store) ListProjectLimits(ctx context.Context, tenantID string) ([]ProjectLimit, error) {
	query := `SELECT l.tenant_id::text, l.project_id::text, p.name, t.name,
		l.requests_per_minute, l.tokens_per_minute, l.concurrent_requests, l.monthly_spend_micros, l.updated_at
		FROM project_limits l JOIN projects p ON p.id=l.project_id JOIN tenants t ON t.id=l.tenant_id`
	args := []any{}
	if strings.TrimSpace(tenantID) != "" {
		query += ` WHERE l.tenant_id=$1`
		args = append(args, tenantID)
	}
	query += ` ORDER BY t.name, p.name`
	rows, err := s.db.Query(ctx, query, args...)
	if err != nil {
		return nil, fmt.Errorf("query project limits: %w", err)
	}
	defer rows.Close()
	result := make([]ProjectLimit, 0)
	for rows.Next() {
		var item ProjectLimit
		if err := rows.Scan(&item.TenantID, &item.ProjectID, &item.ProjectName, &item.TenantName,
			&item.RequestsPerMinute, &item.TokensPerMinute, &item.Concurrent, &item.MonthlySpendMicros, &item.UpdatedAt); err != nil {
			return nil, fmt.Errorf("scan project limit: %w", err)
		}
		result = append(result, item)
	}
	return result, rows.Err()
}

func (s *Store) SetProjectLimit(ctx context.Context, projectID string, input SetProjectLimitInput) (ProjectLimit, int64, error) {
	if strings.TrimSpace(projectID) == "" || input.RequestsPerMinute < 0 || input.TokensPerMinute < 0 || input.ConcurrentRequests < 0 || input.MonthlySpendMicros < 0 {
		return ProjectLimit{}, 0, errors.New("limit values cannot be negative and project_id is required")
	}
	tx, err := s.db.Begin(ctx)
	if err != nil {
		return ProjectLimit{}, 0, err
	}
	defer tx.Rollback(ctx)
	var result ProjectLimit
	err = tx.QueryRow(ctx, `
		WITH project AS (
			SELECT p.id, p.tenant_id, p.name AS project_name, t.name AS tenant_name
			FROM projects p JOIN tenants t ON t.id=p.tenant_id WHERE p.id=$1
		), updated AS (
			INSERT INTO project_limits (project_id, tenant_id, requests_per_minute, tokens_per_minute, concurrent_requests, monthly_spend_micros)
			SELECT p.id, p.tenant_id, $2, $3, $4, $5 FROM project p
			ON CONFLICT (project_id) DO UPDATE SET requests_per_minute=EXCLUDED.requests_per_minute,
				tokens_per_minute=EXCLUDED.tokens_per_minute, concurrent_requests=EXCLUDED.concurrent_requests,
				monthly_spend_micros=EXCLUDED.monthly_spend_micros, updated_at=now()
			RETURNING tenant_id, project_id, requests_per_minute, tokens_per_minute, concurrent_requests, monthly_spend_micros, updated_at
		)
		SELECT u.tenant_id::text, u.project_id::text, p.project_name, p.tenant_name,
			u.requests_per_minute, u.tokens_per_minute, u.concurrent_requests, u.monthly_spend_micros, u.updated_at
		FROM updated u JOIN project p ON p.id=u.project_id`,
		projectID, input.RequestsPerMinute, input.TokensPerMinute, input.ConcurrentRequests, input.MonthlySpendMicros,
	).Scan(&result.TenantID, &result.ProjectID, &result.ProjectName, &result.TenantName,
		&result.RequestsPerMinute, &result.TokensPerMinute, &result.Concurrent, &result.MonthlySpendMicros, &result.UpdatedAt)
	if errors.Is(err, pgx.ErrNoRows) {
		return ProjectLimit{}, 0, ErrNotFound
	}
	if err != nil {
		return ProjectLimit{}, 0, fmt.Errorf("set project limit: %w", err)
	}
	generation, err := bumpGeneration(ctx, tx)
	if err != nil {
		return ProjectLimit{}, 0, err
	}
	if err := tx.Commit(ctx); err != nil {
		return ProjectLimit{}, 0, err
	}
	return result, generation, nil
}