Files
AIRouter/internal/task/store.go
T
freedakgmail e7e98271d4
CI / lint (push) Has been cancelled
CI / test (push) Has been cancelled
CI / build (push) Has been cancelled
CI / security-scan (push) Has been cancelled
feat: 添加过载保护、背压和熔断机制
2026-08-03 08:12:28 +08:00

175 lines
4.8 KiB
Go

package task
import (
"database/sql"
"encoding/json"
"fmt"
"sync"
"time"
_ "github.com/mattn/go-sqlite3"
)
// Store manages task state persistence with SQLite.
type Store struct {
mu sync.Mutex
db *sql.DB
}
// NewStore creates a new task store.
func NewStore(dbPath string) (*Store, error) {
db, err := sql.Open("sqlite3", dbPath)
if err != nil {
return nil, fmt.Errorf("open task db: %w", err)
}
if err := initTaskDB(db); err != nil {
return nil, fmt.Errorf("init task db: %w", err)
}
return &Store{db: db}, nil
}
func initTaskDB(db *sql.DB) error {
schema := `
CREATE TABLE IF NOT EXISTS tasks (
id TEXT PRIMARY KEY,
request_id TEXT NOT NULL,
session_id TEXT,
app_id TEXT NOT NULL,
tenant_id TEXT NOT NULL,
logical_model TEXT NOT NULL,
actual_model TEXT,
priority INTEGER NOT NULL DEFAULT 2,
state TEXT NOT NULL,
stream INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL,
started_at TEXT,
completed_at TEXT,
cancel_reason TEXT,
error_message TEXT,
input_tokens INTEGER DEFAULT 0,
output_tokens INTEGER DEFAULT 0,
node_id TEXT,
degraded INTEGER DEFAULT 0
);
CREATE INDEX IF NOT EXISTS idx_tasks_state ON tasks(state);
CREATE INDEX IF NOT EXISTS idx_tasks_app ON tasks(app_id);
CREATE INDEX IF NOT EXISTS idx_tasks_tenant ON tasks(tenant_id);`
_, err := db.Exec(schema)
return err
}
// Save persists a task to the database.
func (s *Store) Save(t *Task) error {
s.mu.Lock()
defer s.mu.Unlock()
var startedAt, completedAt interface{}
if t.StartedAt != nil {
startedAt = t.StartedAt.Format(time.RFC3339)
}
if t.CompletedAt != nil {
completedAt = t.CompletedAt.Format(time.RFC3339)
}
streamInt := 0
if t.Stream {
streamInt = 1
}
degradedInt := 0
if t.Degraded {
degradedInt = 1
}
_, err := s.db.Exec(
`INSERT OR REPLACE INTO tasks
(id, request_id, session_id, app_id, tenant_id, logical_model, actual_model, priority, state, stream, created_at, started_at, completed_at, cancel_reason, error_message, input_tokens, output_tokens, node_id, degraded)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
t.ID, t.RequestID, t.SessionID, t.AppID, t.TenantID, t.LogicalModel, t.ActualModel,
int(t.Priority), string(t.State), streamInt, t.CreatedAt.Format(time.RFC3339),
startedAt, completedAt, t.CancelReason, t.ErrorMessage,
t.InputTokens, t.OutputTokens, t.NodeID, degradedInt,
)
return err
}
// Get retrieves a task by ID.
func (s *Store) Get(id string) (*Task, error) {
s.mu.Lock()
defer s.mu.Unlock()
var (
requestID, sessionID, appID, tenantID, logicalModel, actualModel, state string
priority int
streamInt int
createdAtStr, startedAt, completedAt, cancelReason, errorMessage, nodeID sql.NullString
inputTokens, outputTokens, degradedInt int
)
err := s.db.QueryRow(
`SELECT request_id, session_id, app_id, tenant_id, logical_model, actual_model, priority, state, stream, created_at, started_at, completed_at, cancel_reason, error_message, input_tokens, output_tokens, node_id, degraded FROM tasks WHERE id = ?`,
id,
).Scan(&requestID, &sessionID, &appID, &tenantID, &logicalModel, &actualModel, &priority, &state, &streamInt, &createdAtStr, &startedAt, &completedAt, &cancelReason, &errorMessage, &inputTokens, &outputTokens, &nodeID, &degradedInt)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
t := &Task{
ID: id,
RequestID: requestID,
SessionID: sessionID,
AppID: appID,
TenantID: tenantID,
LogicalModel: logicalModel,
ActualModel: actualModel,
Priority: TaskPriority(priority),
State: TaskState(state),
Stream: streamInt == 1,
InputTokens: inputTokens,
OutputTokens: outputTokens,
NodeID: nodeID.String,
Degraded: degradedInt == 1,
CancelReason: cancelReason.String,
ErrorMessage: errorMessage.String,
cancelCh: make(chan struct{}),
}
t.CreatedAt, _ = time.Parse(time.RFC3339, createdAtStr.String)
if startedAt.Valid {
tt, _ := time.Parse(time.RFC3339, startedAt.String)
t.StartedAt = &tt
}
if completedAt.Valid {
tt, _ := time.Parse(time.RFC3339, completedAt.String)
t.CompletedAt = &tt
}
return t, nil
}
// RecoverPendingTasks marks RUNNING/STREAMING tasks as FAILED on startup.
func (s *Store) RecoverPendingTasks() (int, error) {
s.mu.Lock()
defer s.mu.Unlock()
result, err := s.db.Exec(
`UPDATE tasks SET state = 'FAILED', error_message = 'gateway restart' WHERE state IN ('RUNNING', 'STREAMING', 'QUEUED')`)
if err != nil {
return 0, err
}
n, _ := result.RowsAffected()
return int(n), nil
}
// Close closes the database connection.
func (s *Store) Close() error {
return s.db.Close()
}
// Ensure json is imported for future use.
var _ = json.Marshal