161 lines
6.3 KiB
Go
161 lines
6.3 KiB
Go
package creator
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/sirupsen/logrus"
|
|
)
|
|
|
|
func scanSettings(scanner interface{ Scan(...any) error }) (Settings, error) {
|
|
var result Settings
|
|
if err := scanner.Scan(&result.LookbackDays, &result.NewWorkIntervalSeconds, &result.MetricInitialIntervalSeconds,
|
|
&result.MetricMultiplier, &result.MetricMaxIntervalSeconds, &result.MetricAgeSeconds,
|
|
&result.AIBaseURL, &result.AIAPIKey, &result.AIModel, &result.UpdatedAt); err != nil {
|
|
return Settings{}, err
|
|
}
|
|
result.UpdatedAt = result.UpdatedAt.UTC()
|
|
return result, nil
|
|
}
|
|
|
|
const settingsColumns = `lookback_days, new_work_interval_seconds, metric_initial_interval_seconds,
|
|
metric_multiplier, metric_max_interval_seconds, metric_age_seconds, ai_base_url, ai_api_key, ai_model, updated_at`
|
|
const settingsSelect = `SELECT ` + settingsColumns + ` FROM creator_settings`
|
|
|
|
func (s *Store) GetSettings(ctx context.Context) (Settings, error) {
|
|
result, err := scanSettings(s.db.QueryRowContext(ctx, settingsSelect))
|
|
return result, rowError(err)
|
|
}
|
|
|
|
func validateAISettings(input SettingsUpdate) (SettingsUpdate, error) {
|
|
input.AIBaseURL = strings.TrimSpace(input.AIBaseURL)
|
|
input.AIAPIKey = strings.TrimSpace(input.AIAPIKey)
|
|
input.AIModel = strings.TrimSpace(input.AIModel)
|
|
if input.AIBaseURL == "" && input.AIAPIKey == "" && input.AIModel == "" {
|
|
return input, nil
|
|
}
|
|
connection, err := validateAIConnection(AIConnection{BaseURL: input.AIBaseURL, APIKey: input.AIAPIKey})
|
|
if err != nil {
|
|
return SettingsUpdate{}, err
|
|
}
|
|
if input.AIModel == "" {
|
|
return SettingsUpdate{}, fmt.Errorf("%w: 请从模型列表选择 MODEL", ErrInvalid)
|
|
}
|
|
input.AIBaseURL = connection.BaseURL
|
|
return input, nil
|
|
}
|
|
|
|
func (s *Store) UpdateSettings(ctx context.Context, input SettingsUpdate) (Settings, error) {
|
|
if err := ValidateSettings(input); err != nil {
|
|
return Settings{}, err
|
|
}
|
|
input, err := validateAISettings(input)
|
|
if err != nil {
|
|
return Settings{}, err
|
|
}
|
|
tx, err := s.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return Settings{}, fmt.Errorf("begin creator settings update: %w", err)
|
|
}
|
|
defer tx.Rollback()
|
|
current, err := scanSettings(tx.QueryRowContext(ctx, settingsSelect+` FOR UPDATE`))
|
|
if err != nil {
|
|
return Settings{}, rowError(err)
|
|
}
|
|
collectionChanged := current.LookbackDays != input.LookbackDays || current.NewWorkIntervalSeconds != input.NewWorkIntervalSeconds ||
|
|
current.MetricInitialIntervalSeconds != input.MetricInitialIntervalSeconds || current.MetricMultiplier != input.MetricMultiplier ||
|
|
current.MetricMaxIntervalSeconds != input.MetricMaxIntervalSeconds || current.MetricAgeSeconds != input.MetricAgeSeconds
|
|
value, err := scanSettings(tx.QueryRowContext(ctx, `
|
|
UPDATE creator_settings SET lookback_days=$1, new_work_interval_seconds=$2,
|
|
metric_initial_interval_seconds=$3, metric_multiplier=$4, metric_max_interval_seconds=$5,
|
|
metric_age_seconds=$6, ai_base_url=$7, ai_api_key=$8, ai_model=$9,
|
|
updated_at=now() RETURNING `+settingsColumns, input.LookbackDays, input.NewWorkIntervalSeconds,
|
|
input.MetricInitialIntervalSeconds, input.MetricMultiplier, input.MetricMaxIntervalSeconds,
|
|
input.MetricAgeSeconds, input.AIBaseURL, input.AIAPIKey, input.AIModel))
|
|
if err != nil {
|
|
return Settings{}, databaseError(err)
|
|
}
|
|
if collectionChanged {
|
|
var runningCheckpoints int
|
|
if err := tx.QueryRowContext(ctx, `SELECT count(*) FROM creator_collection_checkpoint WHERE status='running'`).Scan(&runningCheckpoints); err != nil {
|
|
return Settings{}, databaseError(err)
|
|
}
|
|
if runningCheckpoints > 0 {
|
|
return Settings{}, ErrConflict
|
|
}
|
|
rows, err := tx.QueryContext(ctx, `SELECT work_id, published_at FROM creator_work WHERE published_at IS NOT NULL FOR UPDATE`)
|
|
if err != nil {
|
|
return Settings{}, databaseError(err)
|
|
}
|
|
type metricSchedule struct {
|
|
workID string
|
|
publishedAt *time.Time
|
|
}
|
|
schedules := make([]metricSchedule, 0)
|
|
for rows.Next() {
|
|
var item metricSchedule
|
|
if err := rows.Scan(&item.workID, &item.publishedAt); err != nil {
|
|
_ = rows.Close()
|
|
return Settings{}, err
|
|
}
|
|
schedules = append(schedules, item)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
_ = rows.Close()
|
|
return Settings{}, err
|
|
}
|
|
if err := rows.Close(); err != nil {
|
|
return Settings{}, err
|
|
}
|
|
now := time.Now().UTC()
|
|
start, end, err := NewCollectionWindow(now, input.LookbackDays)
|
|
if err != nil {
|
|
return Settings{}, err
|
|
}
|
|
if _, err := tx.ExecContext(ctx, `
|
|
UPDATE creator_collection_checkpoint
|
|
SET window_start=$1, window_end=$2, cursor='', lease_token='', lease_until=NULL,
|
|
status=CASE WHEN status='blocked' THEN 'blocked' ELSE 'idle' END,
|
|
last_error=CASE WHEN status='blocked' THEN last_error ELSE '' END
|
|
WHERE status <> 'running'`, start, end); err != nil {
|
|
return Settings{}, databaseError(err)
|
|
}
|
|
for _, schedule := range schedules {
|
|
nextAt, reason := NextMetricAtValue(schedule.publishedAt, now, input)
|
|
stopped := nextAt.IsZero()
|
|
if _, err := tx.ExecContext(ctx, `
|
|
UPDATE creator_work SET next_metric_at=$2, metric_stop_reason=$3,
|
|
metric_plan_next_at=$2, metric_plan_interval_seconds=$4, metric_plan_multiplier=$5,
|
|
metric_plan_max_interval_seconds=$6, metric_plan_stopped=$7, updated_at=now()
|
|
WHERE work_id=$1`, schedule.workID, nullableArg(nextAt), reason, input.MetricInitialIntervalSeconds,
|
|
input.MetricMultiplier, input.MetricMaxIntervalSeconds, stopped); err != nil {
|
|
return Settings{}, databaseError(err)
|
|
}
|
|
}
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
return Settings{}, fmt.Errorf("commit creator settings update: %w", err)
|
|
}
|
|
logrus.WithFields(logrus.Fields{"collection_changed": collectionChanged, "ai_base_url": value.AIBaseURL, "ai_model": value.AIModel}).Info("Creator settings saved")
|
|
return value, nil
|
|
}
|
|
|
|
func NextMetricAtValue(publishedAt *time.Time, now time.Time, input SettingsUpdate) (time.Time, string) {
|
|
if publishedAt == nil {
|
|
return time.Time{}, "published_at_pending_verification"
|
|
}
|
|
return NextMetricAt(publishedAt.UTC(), now.UTC(),
|
|
time.Duration(input.MetricInitialIntervalSeconds)*time.Second,
|
|
time.Duration(input.MetricMaxIntervalSeconds)*time.Second,
|
|
input.MetricMultiplier, time.Duration(input.MetricAgeSeconds)*time.Second)
|
|
}
|
|
|
|
func (s *Store) EnsureSchema(ctx context.Context) error {
|
|
if _, err := s.db.ExecContext(ctx, `SELECT 1 FROM creator_settings`); err != nil {
|
|
return fmt.Errorf("check creator schema: %w", err)
|
|
}
|
|
return nil
|
|
}
|