Files

277 lines
7.8 KiB
Go

// Package agent owns the Agent's file-backed execution and asset recovery
// state. It deliberately has no business database and never proxies audio to
// Dispatcher.
package agent
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"time"
)
type State struct {
SchemaVersion string `json:"schema_version"`
ExecutionID string `json:"execution_id"`
TaskRevision int64 `json:"task_revision"`
SessionID string `json:"session_id,omitempty"`
Status string `json:"status"`
Unknown bool `json:"unknown"`
Reason string `json:"reason,omitempty"`
UpdatedAt time.Time `json:"updated_at"`
}
type Spool struct {
root string
now func() time.Time
mu sync.Mutex
}
func NewSpool(root string, now func() time.Time) (*Spool, error) {
if strings.TrimSpace(root) == "" {
return nil, errors.New("spool root is required")
}
if now == nil {
now = time.Now
}
if err := os.MkdirAll(root, 0o700); err != nil {
return nil, fmt.Errorf("create spool root: %w", err)
}
return &Spool{root: root, now: now}, nil
}
func (s *Spool) Root() string { return s.root }
func (s *Spool) Start(executionID string, revision int64, sessionID string) (State, error) {
if err := validateName(executionID); err != nil {
return State{}, err
}
if revision < 1 {
return State{}, errors.New("task revision must be positive")
}
state := State{SchemaVersion: "1", ExecutionID: executionID, TaskRevision: revision, SessionID: sessionID, Status: "reserved", UpdatedAt: s.now().UTC()}
s.mu.Lock()
defer s.mu.Unlock()
path := s.statePath(executionID)
if _, err := os.Stat(path); err == nil {
return State{}, fmt.Errorf("execution state already exists: %s", executionID)
} else if !errors.Is(err, os.ErrNotExist) {
return State{}, err
}
for _, dir := range []string{s.executionDir(executionID), filepath.Join(s.executionDir(executionID), "transcript"), filepath.Join(s.executionDir(executionID), "assets")} {
if err := os.MkdirAll(dir, 0o700); err != nil {
return State{}, err
}
}
if err := writeJSONAtomic(path, state); err != nil {
return State{}, err
}
return state, nil
}
func (s *Spool) Load(executionID string) (State, error) {
if err := validateName(executionID); err != nil {
return State{}, err
}
data, err := os.ReadFile(s.statePath(executionID))
if err != nil {
return State{}, err
}
var state State
if err := json.Unmarshal(data, &state); err != nil {
return State{}, fmt.Errorf("decode execution state: %w", err)
}
return state, nil
}
func (s *Spool) Update(executionID, status, reason string) (State, error) {
if err := validateName(executionID); err != nil {
return State{}, err
}
if status == "" {
return State{}, errors.New("state status is required")
}
s.mu.Lock()
defer s.mu.Unlock()
data, err := os.ReadFile(s.statePath(executionID))
if err != nil {
return State{}, err
}
var state State
if err := json.Unmarshal(data, &state); err != nil {
return State{}, fmt.Errorf("decode execution state: %w", err)
}
state.Status, state.Reason, state.UpdatedAt = status, reason, s.now().UTC()
state.Unknown = status == "unknown"
if err := writeJSONAtomic(s.statePath(executionID), state); err != nil {
return State{}, err
}
return state, nil
}
// MarkUnknownOnBoot converts all in-flight local state to explicit unknown.
// It never deletes or releases a remote reservation; Dispatcher reconciliation
// must decide whether a recovered execution may proceed.
func (s *Spool) MarkUnknownOnBoot() (RecoveryReport, error) {
s.mu.Lock()
defer s.mu.Unlock()
entries, err := os.ReadDir(s.root)
if err != nil {
return RecoveryReport{}, err
}
var report RecoveryReport
for _, entry := range entries {
if !entry.IsDir() || strings.HasPrefix(entry.Name(), ".") {
continue
}
executionID := entry.Name()
path := s.statePath(executionID)
data, err := os.ReadFile(path)
if errors.Is(err, os.ErrNotExist) {
continue
}
if err != nil {
return report, err
}
var state State
if err := json.Unmarshal(data, &state); err != nil {
quarantine := path + ".corrupt-" + s.now().UTC().Format("20060102T150405.000000000Z")
if renameErr := os.Rename(path, quarantine); renameErr != nil {
return report, fmt.Errorf("quarantine corrupt state: %w (decode: %v)", renameErr, err)
}
report.Quarantined = append(report.Quarantined, executionID)
continue
}
if state.Status != "running" && state.Status != "reserved" && state.Status != "starting" && state.Status != "draining" {
continue
}
state.Status, state.Unknown, state.Reason, state.UpdatedAt = "unknown", true, "agent_boot_recovery", s.now().UTC()
if err := writeJSONAtomic(path, state); err != nil {
return report, err
}
report.Unknown = append(report.Unknown, executionID)
}
return report, nil
}
type RecoveryReport struct {
Unknown []string
Quarantined []string
}
func (s *Spool) AppendTranscript(executionID string, event []byte) error {
if err := validateName(executionID); err != nil {
return err
}
if len(event) == 0 {
return errors.New("transcript event is empty")
}
if !json.Valid(event) {
return errors.New("transcript event must be valid JSON")
}
path := filepath.Join(s.executionDir(executionID), "transcript", "events.jsonl")
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_APPEND, 0o600)
if err != nil {
return err
}
defer file.Close()
if _, err := file.Write(append(event, '\n')); err != nil {
return err
}
return file.Sync()
}
func (s *Spool) WriteAsset(executionID, assetID string, r io.Reader) (string, int64, string, error) {
if err := validateName(executionID); err != nil {
return "", 0, "", err
}
if err := validateName(assetID); err != nil {
return "", 0, "", err
}
if r == nil {
return "", 0, "", errors.New("asset reader is required")
}
dir := filepath.Join(s.executionDir(executionID), "assets")
if err := os.MkdirAll(dir, 0o700); err != nil {
return "", 0, "", err
}
part := filepath.Join(dir, assetID+".part")
final := filepath.Join(dir, assetID)
file, err := os.OpenFile(part, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o600)
if err != nil {
return "", 0, "", err
}
hash := sha256.New()
n, copyErr := io.Copy(io.MultiWriter(file, hash), r)
syncErr := file.Sync()
closeErr := file.Close()
if copyErr != nil || syncErr != nil || closeErr != nil {
_ = os.Remove(part)
return "", n, "", firstError(copyErr, syncErr, closeErr)
}
if err := os.Rename(part, final); err != nil {
_ = os.Remove(part)
return "", n, "", err
}
return final, n, hex.EncodeToString(hash.Sum(nil)), nil
}
func (s *Spool) executionDir(executionID string) string { return filepath.Join(s.root, executionID) }
func (s *Spool) statePath(executionID string) string {
return filepath.Join(s.executionDir(executionID), "state.json")
}
func writeJSONAtomic(path string, value any) error {
data, err := json.MarshalIndent(value, "", " ")
if err != nil {
return err
}
file, err := os.CreateTemp(filepath.Dir(path), ".state-*.tmp")
if err != nil {
return err
}
tmp := file.Name()
defer os.Remove(tmp)
if _, err := file.Write(append(data, '\n')); err != nil {
_ = file.Close()
return err
}
if err := file.Sync(); err != nil {
_ = file.Close()
_ = os.Remove(tmp)
return err
}
if err := file.Close(); err != nil {
_ = os.Remove(tmp)
return err
}
if err := os.Rename(tmp, path); err != nil {
_ = os.Remove(tmp)
return err
}
return syncDirectory(filepath.Dir(path))
}
func validateName(name string) error {
if name == "" || name == "." || name == ".." || strings.ContainsAny(name, `/\\`) || strings.Contains(name, "..") || strings.TrimSpace(name) != name {
return fmt.Errorf("unsafe file name %q", name)
}
return nil
}
func firstError(errs ...error) error {
for _, err := range errs {
if err != nil {
return err
}
}
return nil
}