175 lines
4.8 KiB
Go
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, °radedInt)
|
|
|
|
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')`)
|
|
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
|