277 lines
7.8 KiB
Go
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
|
|
}
|