Files
go-sip/internal/store/store.go
T

818 lines
27 KiB
Go

package store
import (
"crypto/sha256"
"database/sql"
"embed"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
iofs "io/fs"
"sort"
"strings"
"sync"
"time"
"git.ipao.vip/rogee/go-sip/internal/contract"
"git.ipao.vip/rogee/go-sip/internal/tenant"
_ "modernc.org/sqlite"
)
//go:embed migrations/*.sql
var migrations embed.FS
var (
ErrDuplicateCommand = errors.New("duplicate command")
ErrCommandConflict = errors.New("command id reused with different body")
ErrNoCapacity = errors.New("quota capacity unavailable")
ErrCASConflict = errors.New("control revision conflict")
ErrLeaseHeld = errors.New("active lease held by another dispatcher")
)
type Store struct {
db *sql.DB
now func() time.Time
mu sync.Mutex
}
func Open(path string) (*Store, error) {
if path == "" {
path = "file:dispatcher.db?_pragma=busy_timeout(5000)"
}
db, err := sql.Open("sqlite", path)
if err != nil {
return nil, fmt.Errorf("open sqlite: %w", err)
}
db.SetMaxOpenConns(1)
db.SetMaxIdleConns(1)
store := &Store{db: db, now: time.Now}
if err := store.migrate(); err != nil {
_ = db.Close()
return nil, err
}
if err := store.RecoverOutbox(); err != nil {
_ = db.Close()
return nil, err
}
return store, nil
}
func New(db *sql.DB, now func() time.Time) (*Store, error) {
if db == nil {
return nil, errors.New("nil database")
}
db.SetMaxOpenConns(1)
db.SetMaxIdleConns(1)
if now == nil {
now = time.Now
}
store := &Store{db: db, now: now}
if err := store.migrate(); err != nil {
return nil, err
}
if err := store.RecoverOutbox(); err != nil {
return nil, err
}
return store, nil
}
func (s *Store) migrate() error {
files, err := iofs.Glob(migrations, "migrations/*.sql")
if err != nil {
return err
}
sort.Strings(files)
for _, file := range files {
data, err := migrations.ReadFile(file)
if err != nil {
return err
}
if _, err := s.db.Exec(string(data)); err != nil {
return fmt.Errorf("migrate sqlite %s: %w", file, err)
}
}
return nil
}
func (s *Store) RecoverOutbox() error {
_, err := s.db.Exec(`UPDATE outbox SET status = 'retry', last_error = 'recovered_after_restart' WHERE status = 'dispatching'`)
return err
}
func (s *Store) Close() error { return s.db.Close() }
func (s *Store) DB() *sql.DB { return s.db }
func (s *Store) LoadSchedulerCursor(scope string, tenants []string) (int, error) {
if scope == "" {
return 0, errors.New("scheduler scope is required")
}
encoded, err := json.Marshal(tenants)
if err != nil {
return 0, err
}
s.mu.Lock()
defer s.mu.Unlock()
var cursor int
err = s.db.QueryRow(`SELECT cursor FROM scheduler_state WHERE scope = ?`, scope).Scan(&cursor)
if errors.Is(err, sql.ErrNoRows) {
_, err = s.db.Exec(`INSERT INTO scheduler_state(scope, tenants_json, cursor, updated_at) VALUES(?, ?, 0, ?)`, scope, encoded, s.now().UTC().Format(time.RFC3339Nano))
return 0, err
}
if err != nil {
return 0, err
}
if len(tenants) == 0 {
cursor = 0
} else {
cursor %= len(tenants)
if cursor < 0 {
cursor += len(tenants)
}
}
_, err = s.db.Exec(`UPDATE scheduler_state SET tenants_json = ?, cursor = ?, updated_at = ? WHERE scope = ?`, encoded, cursor, s.now().UTC().Format(time.RFC3339Nano), scope)
return cursor, err
}
func (s *Store) SaveSchedulerCursor(scope string, tenants []string, cursor int) error {
if scope == "" {
return errors.New("scheduler scope is required")
}
if len(tenants) == 0 {
cursor = 0
} else {
cursor %= len(tenants)
if cursor < 0 {
cursor += len(tenants)
}
}
encoded, err := json.Marshal(tenants)
if err != nil {
return err
}
s.mu.Lock()
defer s.mu.Unlock()
_, err = s.db.Exec(`INSERT INTO scheduler_state(scope, tenants_json, cursor, updated_at) VALUES(?, ?, ?, ?)
ON CONFLICT(scope) DO UPDATE SET tenants_json = excluded.tenants_json, cursor = excluded.cursor, updated_at = excluded.updated_at`,
scope, encoded, cursor, s.now().UTC().Format(time.RFC3339Nano))
return err
}
type IngestResult struct {
CommandID string
ExecutionID string
Duplicate bool
PersistedAt time.Time
}
func (s *Store) IngestCommand(raw []byte, routingKey string) (IngestResult, error) {
envelope, payload, err := contract.DecodeExecute(raw)
if err != nil {
return IngestResult{}, err
}
if routingKey != "" {
if err := verifyRouting(envelope.TenantKey, routingKey); err != nil {
return IngestResult{}, err
}
}
now := s.now().UTC()
expired, err := contract.NotAfterExpired(envelope.NotAfter, now)
if err != nil {
return IngestResult{}, err
}
if expired {
return IngestResult{}, fmt.Errorf("command admission deadline has expired")
}
hash := sha256.Sum256(raw)
bodyHash := hex.EncodeToString(hash[:])
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return IngestResult{}, fmt.Errorf("begin ingest: %w", err)
}
defer tx.Rollback()
var existingHash, status string
err = tx.QueryRow(`SELECT body_hash, status FROM inbox WHERE command_id = ?`, envelope.CommandID).Scan(&existingHash, &status)
switch {
case err == nil:
if existingHash != bodyHash {
return IngestResult{}, fmt.Errorf("%w: %s", ErrCommandConflict, envelope.CommandID)
}
return IngestResult{CommandID: envelope.CommandID, ExecutionID: payload.ExecutionID, Duplicate: true, PersistedAt: now}, nil
case !errors.Is(err, sql.ErrNoRows):
return IngestResult{}, fmt.Errorf("lookup inbox: %w", err)
}
if _, err := tx.Exec(`INSERT INTO inbox(command_id, tenant_id, tenant_key, command_type, body_hash, body, status, received_at)
VALUES(?, ?, ?, ?, ?, ?, 'received', ?)`, envelope.CommandID, envelope.TenantID, envelope.TenantKey, envelope.CommandType, bodyHash, raw, now.Format(time.RFC3339Nano)); err != nil {
return IngestResult{}, fmt.Errorf("persist inbox: %w", err)
}
variables, err := json.Marshal(payload.Variables)
if err != nil {
return IngestResult{}, fmt.Errorf("encode variables: %w", err)
}
result, err := tx.Exec(`INSERT INTO tasks(
execution_id, tenant_key, tenant_id, task_id, task_item_id, task_revision, trace_id,
callee, route_policy_id, caller_profile_id, agent_version_id, variables,
ring_timeout_ms, max_call_duration_ms, status, created_at, updated_at)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'accepted', ?, ?)
ON CONFLICT(tenant_key, task_id, task_item_id, task_revision) DO NOTHING`,
payload.ExecutionID, envelope.TenantKey, envelope.TenantID, payload.TaskID, payload.TaskItemID, payload.TaskRevision,
envelope.TraceID, payload.Callee, payload.RoutePolicyID, payload.CallerProfileID, payload.AgentVersionID,
variables, payload.RingTimeoutMS, payload.MaxCallDurationMS, now.Format(time.RFC3339Nano), now.Format(time.RFC3339Nano))
if err != nil {
return IngestResult{}, fmt.Errorf("persist task: %w", err)
}
inserted, err := result.RowsAffected()
if err != nil {
return IngestResult{}, fmt.Errorf("inspect task insert: %w", err)
}
if inserted == 0 {
if _, err := tx.Exec(`UPDATE inbox SET status = 'persisted', persisted_at = ? WHERE command_id = ?`, now.Format(time.RFC3339Nano), envelope.CommandID); err != nil {
return IngestResult{}, fmt.Errorf("mark duplicate inbox: %w", err)
}
if err := tx.Commit(); err != nil {
return IngestResult{}, fmt.Errorf("commit duplicate command: %w", err)
}
return IngestResult{CommandID: envelope.CommandID, ExecutionID: payload.ExecutionID, Duplicate: true, PersistedAt: now}, nil
}
event, err := (contract.EventBuilder{
TenantID: envelope.TenantID, TenantKey: envelope.TenantKey, TraceID: envelope.TraceID,
EventType: "command.result", Aggregate: "command", AggregateID: envelope.CommandID, Version: 1,
Payload: map[string]any{
"command_id": envelope.CommandID, "command_type": envelope.CommandType,
"execution_id": payload.ExecutionID,
"status": "accepted", "reason_code": "accepted",
"requested_task_revision": payload.TaskRevision,
},
}).Marshal(now, envelope.CommandID+"-result")
if err != nil {
return IngestResult{}, fmt.Errorf("build command.result: %w", err)
}
if _, err := tx.Exec(`INSERT INTO outbox(event_id, tenant_key, exchange, routing_key, body, status, created_at)
VALUES(?, ?, ?, ?, ?, 'pending', ?)`, envelope.CommandID+"-result", envelope.TenantKey, tenant.EventExchange, "agent-call.command.result", event, now.Format(time.RFC3339Nano)); err != nil {
return IngestResult{}, fmt.Errorf("persist outbox: %w", err)
}
if _, err := tx.Exec(`UPDATE inbox SET status = 'persisted', persisted_at = ? WHERE command_id = ?`, now.Format(time.RFC3339Nano), envelope.CommandID); err != nil {
return IngestResult{}, fmt.Errorf("mark inbox persisted: %w", err)
}
if err := tx.Commit(); err != nil {
return IngestResult{}, fmt.Errorf("commit ingest: %w", err)
}
return IngestResult{CommandID: envelope.CommandID, ExecutionID: payload.ExecutionID, PersistedAt: now}, nil
}
func verifyRouting(tenantKey, routingKey string) error {
want, err := tenant.CommandRoutingKey(tenantKey)
if err != nil {
return err
}
if want != routingKey {
return fmt.Errorf("tenant routing mismatch: expected %q got %q", want, routingKey)
}
return nil
}
type OutboxRecord struct {
ID int64
EventID string
TenantKey string
Exchange string
RoutingKey string
Body []byte
Attempts int
}
func (s *Store) ClaimOutbox(limit int) ([]OutboxRecord, error) {
if limit <= 0 {
return nil, errors.New("outbox limit must be positive")
}
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return nil, err
}
defer tx.Rollback()
rows, err := tx.Query(`SELECT id, event_id, tenant_key, exchange, routing_key, body, attempts
FROM outbox WHERE status IN ('pending', 'retry') ORDER BY id LIMIT ?`, limit)
if err != nil {
return nil, err
}
var records []OutboxRecord
for rows.Next() {
var r OutboxRecord
if err := rows.Scan(&r.ID, &r.EventID, &r.TenantKey, &r.Exchange, &r.RoutingKey, &r.Body, &r.Attempts); err != nil {
rows.Close()
return nil, err
}
records = append(records, r)
}
if err := rows.Close(); err != nil {
return nil, err
}
for _, r := range records {
if _, err := tx.Exec(`UPDATE outbox SET status = 'dispatching', attempts = attempts + 1 WHERE id = ?`, r.ID); err != nil {
return nil, err
}
}
if err := tx.Commit(); err != nil {
return nil, err
}
return records, nil
}
func (s *Store) MarkOutboxPublished(id int64) error {
now := s.now().UTC().Format(time.RFC3339Nano)
_, err := s.db.Exec(`UPDATE outbox SET status = 'published', published_at = ?, last_error = NULL WHERE id = ?`, now, id)
return err
}
func (s *Store) MarkOutboxRetry(id int64, cause error) error {
message := "retry"
if cause != nil {
message = cause.Error()
}
_, err := s.db.Exec(`UPDATE outbox SET status = 'retry', last_error = ? WHERE id = ?`, message, id)
return err
}
type Task struct {
ExecutionID string
TenantKey string
TenantID string
TaskID string
TaskItemID string
TaskRevision int64
TraceID string
Callee string
RoutePolicyID string
CallerProfileID string
AgentVersionID string
Variables map[string]any
RingTimeoutMS int64
MaxCallDurationMS int64
Status string
CreatedAt time.Time
UpdatedAt time.Time
}
func (s *Store) FindTask(tenantID, taskID string) (Task, error) {
if tenantID == "" || taskID == "" {
return Task{}, errors.New("tenant id and task id are required")
}
s.mu.Lock()
defer s.mu.Unlock()
row := s.db.QueryRow(`SELECT execution_id, tenant_key, tenant_id, task_id, task_item_id, task_revision,
trace_id, callee, route_policy_id, caller_profile_id, agent_version_id, variables,
ring_timeout_ms, max_call_duration_ms, status, created_at, updated_at
FROM tasks WHERE tenant_id = ? AND task_id = ? ORDER BY task_revision DESC LIMIT 1`, tenantID, taskID)
return scanTask(row)
}
func (s *Store) NextTask(tenantKey string) (Task, error) {
if strings.TrimSpace(tenantKey) == "" {
return Task{}, errors.New("tenant key is required")
}
s.mu.Lock()
defer s.mu.Unlock()
row := s.db.QueryRow(`SELECT execution_id, tenant_key, tenant_id, task_id, task_item_id, task_revision,
trace_id, callee, route_policy_id, caller_profile_id, agent_version_id, variables,
ring_timeout_ms, max_call_duration_ms, status, created_at, updated_at
FROM tasks WHERE tenant_key = ? AND status = 'accepted' ORDER BY created_at, execution_id LIMIT 1`, tenantKey)
return scanTask(row)
}
func (s *Store) MarkTaskReserved(executionID string) error {
if executionID == "" {
return errors.New("execution id is required")
}
s.mu.Lock()
defer s.mu.Unlock()
result, err := s.db.Exec(`UPDATE tasks SET status = 'reserved', updated_at = ? WHERE execution_id = ? AND status = 'accepted'`, s.now().UTC().Format(time.RFC3339Nano), executionID)
if err != nil {
return err
}
count, err := result.RowsAffected()
if err != nil {
return err
}
if count != 1 {
return ErrCASConflict
}
return nil
}
func (s *Store) MarkTaskRunning(executionID string) error {
if executionID == "" {
return errors.New("execution id is required")
}
s.mu.Lock()
defer s.mu.Unlock()
result, err := s.db.Exec(`UPDATE tasks SET status = 'running', updated_at = ? WHERE execution_id = ? AND status = 'reserved'`, s.now().UTC().Format(time.RFC3339Nano), executionID)
if err != nil {
return err
}
count, err := result.RowsAffected()
if err != nil {
return err
}
if count != 1 {
return ErrCASConflict
}
return nil
}
// FinalizeReservation releases a failed remote attempt and updates its task in
// the same SQLite transaction. Unknown attempts stay counted in unknown_value;
// only a failure proven before remote submission is requeued as accepted.
func (s *Store) FinalizeReservation(reservationID, executionID string, unknown bool) error {
if reservationID == "" || executionID == "" {
return errors.New("reservation and execution ids are required")
}
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var state, storedExecutionID string
var scopesJSON []byte
if err := tx.QueryRow(`SELECT state, execution_id, scopes FROM reservations WHERE reservation_id = ?`, reservationID).Scan(&state, &storedExecutionID, &scopesJSON); err != nil {
return err
}
if storedExecutionID != executionID {
return fmt.Errorf("reservation execution mismatch: got %q, want %q", storedExecutionID, executionID)
}
if state != "held" {
return nil
}
var scopes []string
if err := json.Unmarshal(scopesJSON, &scopes); err != nil {
return fmt.Errorf("decode reservation scopes: %w", err)
}
if len(scopes) == 0 {
return errors.New("quota scopes are required to finalize a reservation")
}
now := s.now().UTC().Format(time.RFC3339Nano)
for _, scope := range scopes {
var result sql.Result
if unknown {
result, err = tx.Exec(`UPDATE quotas SET reserved_value = reserved_value - 1, unknown_value = unknown_value + 1, updated_at = ? WHERE scope = ? AND reserved_value > 0`, now, scope)
} else {
result, err = tx.Exec(`UPDATE quotas SET reserved_value = reserved_value - 1, updated_at = ? WHERE scope = ? AND reserved_value > 0`, now, scope)
}
if err != nil {
return err
}
count, err := result.RowsAffected()
if err != nil {
return err
}
if count != 1 {
return ErrCASConflict
}
}
newState := "released"
newTaskStatus := "accepted"
if unknown {
newState = "unknown"
newTaskStatus = "unknown"
}
if _, err := tx.Exec(`UPDATE reservations SET state = ?, released_at = ? WHERE reservation_id = ? AND state = 'held'`, newState, now, reservationID); err != nil {
return err
}
result, err := tx.Exec(`UPDATE tasks SET status = ?, updated_at = ? WHERE execution_id = ? AND status = 'reserved'`, newTaskStatus, now, executionID)
if err != nil {
return err
}
count, err := result.RowsAffected()
if err != nil {
return err
}
if count != 1 {
return ErrCASConflict
}
return tx.Commit()
}
type CommandRecord struct {
CommandID string
TenantID string
TenantKey string
CommandType string
Status string
Body []byte
ReceivedAt time.Time
PersistedAt *time.Time
}
func (s *Store) GetCommand(tenantID, commandID string) (CommandRecord, error) {
if tenantID == "" || commandID == "" {
return CommandRecord{}, errors.New("tenant id and command id are required")
}
var record CommandRecord
var received, persisted sql.NullString
if err := s.db.QueryRow(`SELECT command_id, tenant_id, tenant_key, command_type, status, body, received_at, persisted_at
FROM inbox WHERE tenant_id = ? AND command_id = ?`, tenantID, commandID).Scan(&record.CommandID, &record.TenantID, &record.TenantKey, &record.CommandType, &record.Status, &record.Body, &received, &persisted); err != nil {
return CommandRecord{}, err
}
var err error
record.ReceivedAt, err = time.Parse(time.RFC3339Nano, received.String)
if err != nil {
return CommandRecord{}, err
}
if persisted.Valid {
value, err := time.Parse(time.RFC3339Nano, persisted.String)
if err != nil {
return CommandRecord{}, err
}
record.PersistedAt = &value
}
return record, nil
}
func (s *Store) ReplayCommand(idempotencyKey, tenantID, sourceCommandID, reason string) error {
if idempotencyKey == "" || tenantID == "" || sourceCommandID == "" || strings.TrimSpace(reason) == "" {
return errors.New("replay identity and reason are required")
}
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var tenantKey string
var body []byte
if err := tx.QueryRow(`SELECT tenant_key, body FROM inbox WHERE tenant_id = ? AND command_id = ?`, tenantID, sourceCommandID).Scan(&tenantKey, &body); err != nil {
return err
}
var existing string
if err := tx.QueryRow(`SELECT idempotency_key FROM replays WHERE idempotency_key = ?`, idempotencyKey).Scan(&existing); err == nil {
return tx.Commit()
} else if !errors.Is(err, sql.ErrNoRows) {
return err
}
now := s.now().UTC().Format(time.RFC3339Nano)
if _, err := tx.Exec(`INSERT INTO replays(idempotency_key, tenant_id, source_command_id, tenant_key, reason, created_at) VALUES(?, ?, ?, ?, ?, ?)`, idempotencyKey, tenantID, sourceCommandID, tenantKey, reason, now); err != nil {
return err
}
routingKey := "agent-call.tenant." + tenantKey + ".call.execute"
eventID := "replay-" + idempotencyKey
if _, err := tx.Exec(`INSERT INTO outbox(event_id, tenant_key, exchange, routing_key, body, status, created_at) VALUES(?, ?, ?, ?, ?, 'pending', ?)`, eventID, tenantKey, tenant.CommandExchange, routingKey, body, now); err != nil {
return err
}
return tx.Commit()
}
func scanTask(row *sql.Row) (Task, error) {
var t Task
var variables []byte
var created, updated string
if err := row.Scan(&t.ExecutionID, &t.TenantKey, &t.TenantID, &t.TaskID, &t.TaskItemID, &t.TaskRevision,
&t.TraceID, &t.Callee, &t.RoutePolicyID, &t.CallerProfileID, &t.AgentVersionID, &variables,
&t.RingTimeoutMS, &t.MaxCallDurationMS, &t.Status, &created, &updated); err != nil {
return Task{}, err
}
if err := json.Unmarshal(variables, &t.Variables); err != nil {
return Task{}, fmt.Errorf("decode task variables: %w", err)
}
var err error
t.CreatedAt, err = time.Parse(time.RFC3339Nano, created)
if err != nil {
return Task{}, err
}
t.UpdatedAt, err = time.Parse(time.RFC3339Nano, updated)
if err != nil {
return Task{}, err
}
return t, nil
}
func (s *Store) SetQuota(scope string, limit int64) error {
if scope == "" || limit < 0 {
return errors.New("invalid quota")
}
_, err := s.db.Exec(`INSERT INTO quotas(scope, limit_value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(scope) DO UPDATE SET limit_value = excluded.limit_value, updated_at = excluded.updated_at`, scope, limit, s.now().UTC().Format(time.RFC3339Nano))
return err
}
func (s *Store) Reserve(reservationID, executionID, tenantKey string, scopes []string) error {
if reservationID == "" || executionID == "" || tenantKey == "" || len(scopes) == 0 {
return errors.New("reservation identity and scopes are required")
}
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
now := s.now().UTC().Format(time.RFC3339Nano)
scopesJSON, err := json.Marshal(scopes)
if err != nil {
return err
}
for _, scope := range scopes {
var limit, reserved, unknown int64
if err := tx.QueryRow(`SELECT limit_value, reserved_value, unknown_value FROM quotas WHERE scope = ?`, scope).Scan(&limit, &reserved, &unknown); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return fmt.Errorf("%w: quota %s is not configured", ErrNoCapacity, scope)
}
return err
}
if reserved+unknown >= limit {
return fmt.Errorf("%w: %s", ErrNoCapacity, scope)
}
}
if _, err := tx.Exec(`INSERT INTO reservations(reservation_id, execution_id, tenant_key, scopes, state, created_at) VALUES(?, ?, ?, ?, 'held', ?)`, reservationID, executionID, tenantKey, scopesJSON, now); err != nil {
return err
}
for _, scope := range scopes {
if _, err := tx.Exec(`UPDATE quotas SET reserved_value = reserved_value + 1, updated_at = ? WHERE scope = ?`, now, scope); err != nil {
return err
}
}
if err := tx.Commit(); err != nil {
return err
}
return nil
}
func (s *Store) ReleaseReservation(reservationID string, unknown bool) error {
if reservationID == "" {
return errors.New("reservation id is required")
}
return s.releaseReservation(reservationID, unknown, nil)
}
func (s *Store) ReleaseReservationWithScopes(reservationID string, scopes []string, unknown bool) error {
return s.releaseReservation(reservationID, unknown, scopes)
}
func (s *Store) releaseReservation(reservationID string, unknown bool, scopes []string) error {
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var state string
var scopesJSON []byte
if err := tx.QueryRow(`SELECT state, scopes FROM reservations WHERE reservation_id = ?`, reservationID).Scan(&state, &scopesJSON); err != nil {
return err
}
if state != "held" {
return nil
}
if len(scopes) == 0 {
if err := json.Unmarshal(scopesJSON, &scopes); err != nil {
return fmt.Errorf("decode reservation scopes: %w", err)
}
}
if len(scopes) == 0 {
return errors.New("quota scopes are required to release a reservation")
}
now := s.now().UTC().Format(time.RFC3339Nano)
for _, scope := range scopes {
if unknown {
if _, err := tx.Exec(`UPDATE quotas SET reserved_value = reserved_value - 1, unknown_value = unknown_value + 1, updated_at = ? WHERE scope = ? AND reserved_value > 0`, now, scope); err != nil {
return err
}
} else if _, err := tx.Exec(`UPDATE quotas SET reserved_value = reserved_value - 1, updated_at = ? WHERE scope = ? AND reserved_value > 0`, now, scope); err != nil {
return err
}
}
newState := "released"
if unknown {
newState = "unknown"
}
if _, err := tx.Exec(`UPDATE reservations SET state = ?, released_at = ? WHERE reservation_id = ?`, newState, now, reservationID); err != nil {
return err
}
return tx.Commit()
}
type Lease struct {
LeaseID string
Scope string
HolderID string
ExpiresAt time.Time
State string
}
func (s *Store) AcquireLease(leaseID, scope, holderID string, ttl time.Duration) (Lease, error) {
if leaseID == "" || scope == "" || holderID == "" || ttl <= 0 {
return Lease{}, errors.New("lease id, scope, holder, and positive ttl are required")
}
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return Lease{}, err
}
defer tx.Rollback()
now := s.now().UTC()
var existingID, existingHolder, existingExpiry string
err = tx.QueryRow(`SELECT lease_id, holder_id, expires_at FROM leases WHERE scope = ? AND state = 'active' LIMIT 1`, scope).Scan(&existingID, &existingHolder, &existingExpiry)
if err == nil {
expires, parseErr := time.Parse(time.RFC3339Nano, existingExpiry)
if parseErr != nil {
return Lease{}, parseErr
}
if expires.After(now) && existingHolder != holderID {
return Lease{}, fmt.Errorf("%w: scope %s", ErrLeaseHeld, scope)
}
if expires.After(now) && existingHolder == holderID && existingID != leaseID {
return Lease{}, fmt.Errorf("%w: holder already owns scope %s", ErrLeaseHeld, scope)
}
if !expires.After(now) {
if _, err := tx.Exec(`UPDATE leases SET state = 'expired' WHERE lease_id = ?`, existingID); err != nil {
return Lease{}, err
}
}
} else if !errors.Is(err, sql.ErrNoRows) {
return Lease{}, err
}
expires := now.Add(ttl)
if _, err := tx.Exec(`INSERT INTO leases(lease_id, scope, holder_id, expires_at, state) VALUES(?, ?, ?, ?, 'active')
ON CONFLICT(lease_id) DO UPDATE SET scope = excluded.scope, holder_id = excluded.holder_id, expires_at = excluded.expires_at, state = 'active'`, leaseID, scope, holderID, expires.Format(time.RFC3339Nano)); err != nil {
return Lease{}, err
}
if err := tx.Commit(); err != nil {
return Lease{}, err
}
return Lease{LeaseID: leaseID, Scope: scope, HolderID: holderID, ExpiresAt: expires, State: "active"}, nil
}
func (s *Store) ReleaseLease(leaseID string) error {
if leaseID == "" {
return errors.New("lease id is required")
}
_, err := s.db.Exec(`UPDATE leases SET state = 'released' WHERE lease_id = ? AND state = 'active'`, leaseID)
return err
}
func (s *Store) ApplyControl(executionID string, expectedRevision int64, action string) error {
return s.ApplyControlDetailed(executionID, expectedRevision, action, "", "", "")
}
func (s *Store) ApplyControlDetailed(executionID string, expectedRevision int64, action, activeCallPolicy, reason, idempotencyKey string) error {
if executionID == "" || expectedRevision < 1 {
return errors.New("execution id and expected revision are required")
}
if action != "pause" && action != "resume" && action != "drain" && action != "stop" && action != "hangup" {
return fmt.Errorf("unsupported control action %q", action)
}
s.mu.Lock()
defer s.mu.Unlock()
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var status string
if err := tx.QueryRow(`SELECT status FROM tasks WHERE execution_id = ? AND task_revision = ?`, executionID, expectedRevision).Scan(&status); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return ErrCASConflict
}
return err
}
if status == "stopped" && action != "stop" {
return fmt.Errorf("%w: stopped execution cannot resume", ErrCASConflict)
}
if _, err := tx.Exec(`INSERT INTO controls(execution_id, task_revision, action, active_call_policy, reason, idempotency_key, applied, requested_at)
VALUES(?, ?, ?, ?, ?, ?, 0, ?)
ON CONFLICT(execution_id) DO UPDATE SET task_revision = excluded.task_revision, action = excluded.action, active_call_policy = excluded.active_call_policy, reason = excluded.reason, idempotency_key = excluded.idempotency_key, applied = 0, requested_at = excluded.requested_at`, executionID, expectedRevision, action, activeCallPolicy, reason, idempotencyKey, s.now().UTC().Format(time.RFC3339Nano)); err != nil {
return err
}
newStatus := status
switch action {
case "pause":
newStatus = "paused"
case "resume":
if status != "paused" {
return fmt.Errorf("%w: execution is %s", ErrCASConflict, status)
}
newStatus = "accepted"
case "drain":
newStatus = "draining"
case "stop", "hangup":
newStatus = "stopped"
}
if _, err := tx.Exec(`UPDATE tasks SET status = ?, updated_at = ? WHERE execution_id = ? AND task_revision = ?`, newStatus, s.now().UTC().Format(time.RFC3339Nano), executionID, expectedRevision); err != nil {
return err
}
return tx.Commit()
}