Files
poimen-workflows/internal/recovery/checkpoint.go
T

258 lines
5.9 KiB
Go
Raw Normal View History

package recovery
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync"
"time"
)
// Checkpoint represents a saved workflow state
type Checkpoint struct {
WorkflowID string `json:"workflow_id"`
Timestamp time.Time `json:"timestamp"`
Stage string `json:"stage"` // e.g., "clone", "plan", "implement", "judge", "merge"
CompletedTasks []string `json:"completed_tasks"`
PendingTasks []string `json:"pending_tasks"`
FailedTasks []string `json:"failed_tasks"`
CurrentTaskID string `json:"current_task_id"`
CurrentActivityType string `json:"current_activity_type"`
Metadata map[string]any `json:"metadata"`
}
// CheckpointManager manages workflow checkpoints for recovery
type CheckpointManager struct {
mu sync.RWMutex
basePath string
interval time.Duration
stopChan chan struct{}
wg sync.WaitGroup
running bool
current *Checkpoint
lastSave time.Time
}
// NewCheckpointManager creates a new checkpoint manager
func NewCheckpointManager(basePath string, interval time.Duration) *CheckpointManager {
return &CheckpointManager{
basePath: basePath,
interval: interval,
stopChan: make(chan struct{}),
current: &Checkpoint{Metadata: make(map[string]any)},
}
}
// Start starts periodic checkpoint saving
func (cm *CheckpointManager) Start(workflowID string) error {
cm.mu.Lock()
defer cm.mu.Unlock()
if cm.running {
return fmt.Errorf("checkpoint manager already running")
}
cm.current.WorkflowID = workflowID
cm.current.Timestamp = time.Now()
cm.running = true
// Start periodic checkpoint save
cm.wg.Add(1)
go cm.periodicCheckpoint()
return nil
}
// Stop stops checkpoint saving and performs a final save
func (cm *CheckpointManager) Stop() error {
cm.mu.Lock()
defer cm.mu.Unlock()
if !cm.running {
return nil
}
cm.running = false
close(cm.stopChan)
cm.wg.Wait()
// Final checkpoint
return cm.saveLocked()
}
// Update updates the current checkpoint
func (cm *CheckpointManager) Update(checkpoint *Checkpoint) error {
cm.mu.Lock()
defer cm.mu.Unlock()
checkpoint.Timestamp = time.Now()
cm.current = checkpoint
return nil
}
// UpdateStage updates the current stage
func (cm *CheckpointManager) UpdateStage(stage string) error {
cm.mu.Lock()
defer cm.mu.Unlock()
cm.current.Stage = stage
cm.current.Timestamp = time.Now()
return nil
}
// AddCompletedTask adds a completed task to the checkpoint
func (cm *CheckpointManager) AddCompletedTask(taskID string) error {
cm.mu.Lock()
defer cm.mu.Unlock()
cm.current.CompletedTasks = append(cm.current.CompletedTasks, taskID)
cm.current.Timestamp = time.Now()
// Remove from pending if it's there
for i, id := range cm.current.PendingTasks {
if id == taskID {
cm.current.PendingTasks = append(cm.current.PendingTasks[:i], cm.current.PendingTasks[i+1:]...)
break
}
}
return nil
}
// AddFailedTask adds a failed task to the checkpoint
func (cm *CheckpointManager) AddFailedTask(taskID string) error {
cm.mu.Lock()
defer cm.mu.Unlock()
cm.current.FailedTasks = append(cm.current.FailedTasks, taskID)
cm.current.Timestamp = time.Now()
// Remove from pending if it's there
for i, id := range cm.current.PendingTasks {
if id == taskID {
cm.current.PendingTasks = append(cm.current.PendingTasks[:i], cm.current.PendingTasks[i+1:]...)
break
}
}
return nil
}
// SetPendingTasks sets the list of pending tasks
func (cm *CheckpointManager) SetPendingTasks(tasks []string) error {
cm.mu.Lock()
defer cm.mu.Unlock()
cm.current.PendingTasks = tasks
cm.current.Timestamp = time.Now()
return nil
}
// GetLatest retrieves the latest checkpoint from disk
func (cm *CheckpointManager) GetLatest(workflowID string) (*Checkpoint, error) {
cm.mu.RLock()
defer cm.mu.RUnlock()
path := cm.checkpointPath(workflowID)
data, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, err
}
var cp Checkpoint
if err := json.Unmarshal(data, &cp); err != nil {
return nil, err
}
return &cp, nil
}
// periodictCheckpoint periodically saves checkpoints
func (cm *CheckpointManager) periodicCheckpoint() {
defer cm.wg.Done()
ticker := time.NewTicker(cm.interval)
defer ticker.Stop()
for {
select {
case <-cm.stopChan:
return
case <-ticker.C:
cm.mu.Lock()
if cm.running {
_ = cm.saveLocked()
}
cm.mu.Unlock()
}
}
}
// saveLocked saves the current checkpoint to disk (must be called with lock held)
func (cm *CheckpointManager) saveLocked() error {
if !cm.running || cm.current == nil {
return nil
}
path := cm.checkpointPath(cm.current.WorkflowID)
// Create directory if it doesn't exist
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
return err
}
data, err := json.MarshalIndent(cm.current, "", " ")
if err != nil {
return err
}
cm.lastSave = time.Now()
return os.WriteFile(path, data, 0644)
}
// checkpointPath returns the path to a checkpoint file
func (cm *CheckpointManager) checkpointPath(workflowID string) string {
return filepath.Join(cm.basePath, "checkpoints", fmt.Sprintf("%s.checkpoint.json", workflowID))
}
// CleanupCheckpoint removes a checkpoint (after successful completion)
func (cm *CheckpointManager) CleanupCheckpoint(workflowID string) error {
path := cm.checkpointPath(workflowID)
if _, err := os.Stat(path); err == nil {
return os.Remove(path)
}
return nil
}
// HasCheckpoint checks if a checkpoint exists
func (cm *CheckpointManager) HasCheckpoint(workflowID string) (bool, error) {
path := cm.checkpointPath(workflowID)
_, err := os.Stat(path)
if err == nil {
return true, nil
}
if os.IsNotExist(err) {
return false, nil
}
return false, err
}
// GetCurrent returns the current checkpoint in memory (non-persistent)
func (cm *CheckpointManager) GetCurrent() *Checkpoint {
cm.mu.RLock()
defer cm.mu.RUnlock()
if cm.current == nil {
return nil
}
// Return a copy to avoid external mutations
cpCopy := *cm.current
return &cpCopy
}