summaryrefslogtreecommitdiff
path: root/internal/controlplane/preferences.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/controlplane/preferences.go')
-rw-r--r--internal/controlplane/preferences.go156
1 files changed, 156 insertions, 0 deletions
diff --git a/internal/controlplane/preferences.go b/internal/controlplane/preferences.go
new file mode 100644
index 0000000..73bcce6
--- /dev/null
+++ b/internal/controlplane/preferences.go
@@ -0,0 +1,156 @@
+package controlplane
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+)
+
+const maxLowBalanceThresholdMicros int64 = 1_000_000_000_000_000
+
+func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defaultThresholdMicros int64) (TenantPreferences, error) {
+ if defaultThresholdMicros <= 0 {
+ defaultThresholdMicros = 5_000_000
+ }
+ result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, LowBalanceThresholdMicros: defaultThresholdMicros}
+ if result.TenantID == "" {
+ return result, nil
+ }
+ var defaultModel, fallbackModel *string
+ var updatedAt time.Time
+ err := s.db.QueryRow(ctx, `
+ SELECT default_model, fallback_model, low_balance_enabled, low_balance_threshold_micros, updated_at
+ FROM tenant_preferences WHERE tenant_id=$1`, result.TenantID).Scan(
+ &defaultModel, &fallbackModel, &result.LowBalanceEnabled, &result.LowBalanceThresholdMicros, &updatedAt)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return result, nil
+ }
+ if err != nil {
+ return TenantPreferences{}, fmt.Errorf("query tenant preferences: %w", err)
+ }
+ if defaultModel != nil {
+ result.DefaultModel = *defaultModel
+ }
+ if fallbackModel != nil {
+ result.FallbackModel = *fallbackModel
+ }
+ result.UpdatedAt = &updatedAt
+ return result, nil
+}
+
+func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput) (TenantPreferences, error) {
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ input.DefaultModel = strings.TrimSpace(input.DefaultModel)
+ input.FallbackModel = strings.TrimSpace(input.FallbackModel)
+ if input.TenantID == "" {
+ return TenantPreferences{}, errors.New("tenant_id is required")
+ }
+ if input.DefaultModel != "" && input.DefaultModel == input.FallbackModel {
+ return TenantPreferences{}, errors.New("default_model and fallback_model must be different")
+ }
+ available, err := s.ListDeveloperModels(ctx, input.TenantID)
+ if err != nil {
+ return TenantPreferences{}, err
+ }
+ allowed := make(map[string]struct{}, len(available))
+ for _, model := range available {
+ allowed[model.PublicID] = struct{}{}
+ }
+ for field, model := range map[string]string{"default_model": input.DefaultModel, "fallback_model": input.FallbackModel} {
+ if model != "" {
+ if _, ok := allowed[model]; !ok {
+ return TenantPreferences{}, fmt.Errorf("%s is not available to this tenant", field)
+ }
+ }
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return TenantPreferences{}, err
+ }
+ defer tx.Rollback(ctx)
+ if err := tx.QueryRow(ctx, `SELECT id FROM tenants WHERE id=$1`, input.TenantID).Scan(new(string)); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TenantPreferences{}, ErrNotFound
+ }
+ return TenantPreferences{}, err
+ }
+ var result TenantPreferences
+ var defaultModel, fallbackModel *string
+ var updatedAt time.Time
+ if err := tx.QueryRow(ctx, `
+ INSERT INTO tenant_preferences (tenant_id, default_model, fallback_model)
+ VALUES ($1, NULLIF($2,''), NULLIF($3,''))
+ ON CONFLICT (tenant_id) DO UPDATE SET default_model=EXCLUDED.default_model,
+ fallback_model=EXCLUDED.fallback_model, updated_at=now()
+ RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled,
+ low_balance_threshold_micros, updated_at`, input.TenantID, input.DefaultModel, input.FallbackModel).Scan(
+ &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled,
+ &result.LowBalanceThresholdMicros, &updatedAt); err != nil {
+ return TenantPreferences{}, fmt.Errorf("save developer preferences: %w", err)
+ }
+ if defaultModel != nil {
+ result.DefaultModel = *defaultModel
+ }
+ if fallbackModel != nil {
+ result.FallbackModel = *fallbackModel
+ }
+ result.UpdatedAt = &updatedAt
+ if err := tx.Commit(ctx); err != nil {
+ return TenantPreferences{}, err
+ }
+ return result, nil
+}
+
+func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaultThresholdMicros int64) (TenantPreferences, error) {
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ if input.TenantID == "" {
+ return TenantPreferences{}, errors.New("tenant_id is required")
+ }
+ if defaultThresholdMicros <= 0 {
+ defaultThresholdMicros = 5_000_000
+ }
+ if input.LowBalanceThresholdMicros != nil && (*input.LowBalanceThresholdMicros < 0 || *input.LowBalanceThresholdMicros > maxLowBalanceThresholdMicros) {
+ return TenantPreferences{}, errors.New("low balance threshold is outside the supported range")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return TenantPreferences{}, err
+ }
+ defer tx.Rollback(ctx)
+ if err := tx.QueryRow(ctx, `SELECT id FROM tenants WHERE id=$1`, input.TenantID).Scan(new(string)); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TenantPreferences{}, ErrNotFound
+ }
+ return TenantPreferences{}, err
+ }
+ var result TenantPreferences
+ var defaultModel, fallbackModel *string
+ var updatedAt time.Time
+ if err := tx.QueryRow(ctx, `
+ INSERT INTO tenant_preferences (tenant_id, low_balance_enabled, low_balance_threshold_micros)
+ VALUES ($1, COALESCE($2::boolean, TRUE), COALESCE($3::bigint, $4::bigint))
+ ON CONFLICT (tenant_id) DO UPDATE SET
+ low_balance_enabled=COALESCE($2::boolean, tenant_preferences.low_balance_enabled),
+ low_balance_threshold_micros=COALESCE($3::bigint, tenant_preferences.low_balance_threshold_micros), updated_at=now()
+ RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled,
+ low_balance_threshold_micros, updated_at`, input.TenantID, input.LowBalanceEnabled, input.LowBalanceThresholdMicros, defaultThresholdMicros).Scan(
+ &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled,
+ &result.LowBalanceThresholdMicros, &updatedAt); err != nil {
+ return TenantPreferences{}, fmt.Errorf("save billing preferences: %w", err)
+ }
+ if defaultModel != nil {
+ result.DefaultModel = *defaultModel
+ }
+ if fallbackModel != nil {
+ result.FallbackModel = *fallbackModel
+ }
+ result.UpdatedAt = &updatedAt
+ if err := tx.Commit(ctx); err != nil {
+ return TenantPreferences{}, err
+ }
+ return result, nil
+}