From 363d6969063637f3643602f964e0390074278527 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aur=C3=A9lien=20Sibiril?= <81782+aureliensibiril@users.noreply.github.com> Date: Fri, 10 Apr 2026 18:53:13 +0200 Subject: [PATCH] Add Restore function for agent checkpoint recovery MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Restore loads a checkpoint from the store, resolves the agent from a registry, and re-enters coreLoop. Handles suspended, nested suspended (concurrent inner restore), and awaiting-approval states. Partial progress is saved when some inner agents complete while others remain suspended. Signed-off-by: Aurélien Sibiril <81782+aureliensibiril@users.noreply.github.com> --- pkg/agent/restore.go | 354 ++++++++++++++++++++++++++ pkg/agent/restore_test.go | 515 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 869 insertions(+) create mode 100644 pkg/agent/restore.go create mode 100644 pkg/agent/restore_test.go diff --git a/pkg/agent/restore.go b/pkg/agent/restore.go new file mode 100644 index 000000000..665e010b3 --- /dev/null +++ b/pkg/agent/restore.go @@ -0,0 +1,354 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "errors" + "fmt" + "sync" + + "go.gearno.de/kit/log" + "go.probo.inc/probo/pkg/llm" +) + +// Restore continues a previously suspended or approval-interrupted agent run +// from its last persisted checkpoint. The registry must contain all agents +// that may have been active (including handoff targets). +func Restore( + ctx context.Context, + store CheckpointStore, + runID string, + registry AgentRegistry, +) (*Result, error) { + cp, err := store.Load(ctx, runID) + if err != nil { + return nil, fmt.Errorf("cannot load checkpoint: %w", err) + } + if cp == nil { + return nil, fmt.Errorf("cannot restore: no checkpoint for run %s", runID) + } + if cp.Version != CheckpointVersion { + return nil, fmt.Errorf("cannot restore: unsupported checkpoint version %d", cp.Version) + } + + agent, err := registry.Agent(cp.AgentName) + if err != nil { + return nil, fmt.Errorf("cannot resolve agent %q: %w", cp.AgentName, err) + } + + return restoreCheckpoint(ctx, agent, cp, store, runID, registry) +} + +func restoreCheckpoint( + ctx context.Context, + agent *Agent, + cp *Checkpoint, + store CheckpointStore, + runID string, + registry AgentRegistry, +) (*Result, error) { + if cp.Version != CheckpointVersion { + return nil, fmt.Errorf("cannot restore: unsupported checkpoint version %d", cp.Version) + } + emitHook(agent, func(h RunHooks) { h.OnRunRestore(ctx, agent, cp) }) + + switch cp.Status { + case CheckpointStatusSuspended: + return restoreSuspended(ctx, agent, cp, store, runID, registry) + + case CheckpointStatusAwaitingApproval: + return restoreAwaitingApproval(ctx, agent, cp, store, runID, registry) + + default: + return nil, fmt.Errorf("cannot restore: unknown checkpoint status %q", cp.Status) + } +} + +func restoreSuspended( + ctx context.Context, + agent *Agent, + cp *Checkpoint, + store CheckpointStore, + runID string, + registry AgentRegistry, +) (*Result, error) { + if len(cp.InnerCheckpoints) > 0 { + return restoreNestedSuspended(ctx, agent, cp, store, runID, registry) + } + + return continueFromMessages(ctx, agent, cp.Messages, cp, store, runID) +} + +func continueFromMessages( + ctx context.Context, + agent *Agent, + messages []llm.Message, + cp *Checkpoint, + store CheckpointStore, + runID string, +) (*Result, error) { + messagesCopy := make([]llm.Message, len(messages)) + copy(messagesCopy, messages) + + if cp.Turns >= agent.maxTurns { + agent.logger.WarnCtx( + ctx, + "restored agent run has already reached max turns", + log.String("agent", agent.name), + log.Int("turns", cp.Turns), + log.Int("max_turns", agent.maxTurns), + ) + } + + return coreLoop( + ctx, + agent, + messagesCopy, + runOpts{ + callLLM: blockingCallLLM, + onEvent: noopEvent, + skipInputGuardrails: true, + skipSessionLoad: true, + initialUsage: cp.Usage, + initialTurns: cp.Turns, + checkpointStore: store, + runID: runID, + toolUsedInRun: cp.ToolUsedInRun, + }, + ) +} + +func restoreNestedSuspended( + ctx context.Context, + agent *Agent, + cp *Checkpoint, + store CheckpointStore, + runID string, + registry AgentRegistry, +) (*Result, error) { + type nestedRestoreEntry struct { + toolCall llm.ToolCall + originalCheckpoint *Checkpoint + suspendedCheckpoint *Checkpoint + result ToolResult + completed bool + err error + } + + completedByID := make(map[string]ToolResult, len(cp.CompletedCalls)) + for _, cc := range cp.CompletedCalls { + completedByID[cc.ToolCallID] = cc.Result + } + + entries := make([]nestedRestoreEntry, len(cp.AllToolCalls)) + var wg sync.WaitGroup + for i, tc := range cp.AllToolCalls { + entries[i].toolCall = tc + result, ok := completedByID[tc.ID] + if ok { + entries[i].result = result + entries[i].completed = true + continue + } + + innerCP, ok := cp.InnerCheckpoints[tc.ID] + if !ok { + entries[i].err = fmt.Errorf("cannot restore nested tool call %q: missing inner checkpoint", tc.ID) + continue + } + entries[i].originalCheckpoint = innerCP + + innerAgent, err := registry.Agent(innerCP.AgentName) + if err != nil { + entries[i].err = fmt.Errorf("cannot resolve inner agent %q: %w", innerCP.AgentName, err) + continue + } + + wg.Add(1) + go func(i int, tc llm.ToolCall, innerAgent *Agent, innerCP *Checkpoint) { + defer wg.Done() + + result, err := restoreCheckpoint(ctx, innerAgent, innerCP, nil, "", registry) + if err != nil { + if se, ok := errors.AsType[*SuspendedError](err); ok { + if se.Checkpoint == nil { + entries[i].err = fmt.Errorf("cannot restore nested tool call %q: missing suspension checkpoint", tc.ID) + return + } + entries[i].suspendedCheckpoint = se.Checkpoint + return + } + entries[i].err = fmt.Errorf("cannot restore nested tool call %q: %w", tc.ID, err) + return + } + + entries[i].result = ToolResult{Content: result.FinalMessage().Text()} + entries[i].completed = true + }(i, tc, innerAgent, innerCP) + } + wg.Wait() + + messages := make([]llm.Message, len(cp.Messages)) + copy(messages, cp.Messages) + + completedCalls := make([]CompletedCall, 0, len(cp.AllToolCalls)) + remainingInner := make(map[string]*Checkpoint) + var restoreErr error + for _, entry := range entries { + switch { + case entry.err != nil: + if entry.originalCheckpoint != nil { + remainingInner[entry.toolCall.ID] = entry.originalCheckpoint + } + if restoreErr == nil { + restoreErr = entry.err + } + continue + + case entry.suspendedCheckpoint != nil: + remainingInner[entry.toolCall.ID] = entry.suspendedCheckpoint + continue + + case !entry.completed: + if restoreErr == nil { + restoreErr = fmt.Errorf("cannot restore nested tool call %q: no result", entry.toolCall.ID) + } + continue + } + + completedCalls = append( + completedCalls, + CompletedCall{ + ToolCallID: entry.toolCall.ID, + Result: entry.result, + }, + ) + messages = append( + messages, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: entry.toolCall.ID, + Parts: []llm.Part{llm.TextPart{Text: entry.result.Content}}, + }, + ) + } + + saveProgress := func() (*Checkpoint, error) { + next := *cp + next.InnerCheckpoints = remainingInner + next.CompletedCalls = completedCalls + if store != nil && runID != "" { + if err := store.Save(ctx, runID, &next); err != nil { + return nil, fmt.Errorf("cannot save nested restore progress: %w", err) + } + } + return &next, nil + } + + if restoreErr != nil { + if _, err := saveProgress(); err != nil { + return nil, errors.Join(restoreErr, err) + } + return nil, restoreErr + } + + if len(remainingInner) > 0 { + next, err := saveProgress() + if err != nil { + return nil, err + } + return nil, &SuspendedError{RunID: runID, Checkpoint: next} + } + + return continueFromMessages(ctx, agent, messages, cp, store, runID) +} + +func restoreAwaitingApproval( + ctx context.Context, + agent *Agent, + cp *Checkpoint, + store CheckpointStore, + runID string, + registry AgentRegistry, +) (*Result, error) { + // Reconstruct an InterruptedError from the checkpoint. + ie := &InterruptedError{ + ToolCalls: cp.PendingToolCalls, + PendingApprovals: cp.PendingApprovals, + Agent: agent, + Messages: cp.Messages, + Usage: cp.Usage, + Turns: cp.Turns, + } + + // Reconstruct outerState if this was a nested interruption. + if len(cp.InnerCheckpoints) > 0 { + if len(cp.InnerCheckpoints) > 1 { + return nil, fmt.Errorf("cannot restore approval checkpoint: expected one inner checkpoint, got %d", len(cp.InnerCheckpoints)) + } + for toolCallID, innerCP := range cp.InnerCheckpoints { + innerAgent, err := registry.Agent(innerCP.AgentName) + if err != nil { + return nil, fmt.Errorf("cannot resolve inner agent %q: %w", innerCP.AgentName, err) + } + + innerIE := &InterruptedError{ + ToolCalls: innerCP.PendingToolCalls, + PendingApprovals: innerCP.PendingApprovals, + Agent: innerAgent, + Messages: innerCP.Messages, + Usage: innerCP.Usage, + Turns: innerCP.Turns, + } + + ie.Agent = innerAgent + ie.Messages = innerCP.Messages + ie.Usage = innerCP.Usage + ie.Turns = innerCP.Turns + ie.ToolCalls = innerCP.PendingToolCalls + ie.PendingApprovals = innerCP.PendingApprovals + + ie.outerState = &outerLoopState{ + agent: agent, + messages: cp.Messages, + usage: cp.Usage, + turns: cp.Turns, + allToolCalls: cp.AllToolCalls, + toolCallID: toolCallID, + completedCalls: cp.CompletedCalls, + innerInterrupt: innerIE, + } + break + } + } + + if len(cp.ApprovalInput) > 0 { + return resumeWithOpts( + ctx, + ie, + ResumeInput{Approvals: cp.ApprovalInput}, + runOpts{ + callLLM: blockingCallLLM, + onEvent: noopEvent, + checkpointStore: store, + runID: runID, + toolUsedInRun: cp.ToolUsedInRun, + }, + ) + } + + return nil, ie +} diff --git a/pkg/agent/restore_test.go b/pkg/agent/restore_test.go new file mode 100644 index 000000000..410d7522c --- /dev/null +++ b/pkg/agent/restore_test.go @@ -0,0 +1,515 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/llm" +) + +type memoryCheckpointStore struct { + mu sync.Mutex + checkpoints map[string]*agent.Checkpoint +} + +func newMemoryCheckpointStore() *memoryCheckpointStore { + return &memoryCheckpointStore{ + checkpoints: make(map[string]*agent.Checkpoint), + } +} + +func (s *memoryCheckpointStore) Save(_ context.Context, runID string, cp *agent.Checkpoint) error { + s.mu.Lock() + defer s.mu.Unlock() + + clone := *cp + s.checkpoints[runID] = &clone + return nil +} + +func (s *memoryCheckpointStore) Load(_ context.Context, runID string) (*agent.Checkpoint, error) { + s.mu.Lock() + defer s.mu.Unlock() + + cp, ok := s.checkpoints[runID] + if !ok { + return nil, nil + } + + clone := *cp + return &clone, nil +} + +func (s *memoryCheckpointStore) Delete(_ context.Context, runID string) error { + s.mu.Lock() + defer s.mu.Unlock() + + delete(s.checkpoints, runID) + return nil +} + +type simpleRegistry struct { + agents map[string]*agent.Agent +} + +func (r *simpleRegistry) Agent(name string) (*agent.Agent, error) { + a, ok := r.agents[name] + if !ok { + return nil, fmt.Errorf("agent %q not found", name) + } + return a, nil +} + +func TestRestore(t *testing.T) { + t.Parallel() + + t.Run( + "no checkpoint returns error", + func(t *testing.T) { + t.Parallel() + + store := newMemoryCheckpointStore() + registry := &simpleRegistry{agents: map[string]*agent.Agent{}} + + _, err := agent.Restore( + context.Background(), + store, + "nonexistent-run", + registry, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "no checkpoint") + }, + ) + + t.Run( + "unsupported checkpoint version returns error", + func(t *testing.T) { + t.Parallel() + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-1", &agent.Checkpoint{ + Version: 999, + Status: agent.CheckpointStatusSuspended, + AgentName: "test-agent", + }) + require.NoError(t, err) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{ + "test-agent": agent.New( + "test-agent", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ), + }, + } + + _, err = agent.Restore( + context.Background(), + store, + "run-1", + registry, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "unsupported checkpoint version") + }, + ) + + t.Run( + "suspended checkpoint restores and completes", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Restored successfully."), + }, + } + + ag := agent.New( + "test-agent", + newTestClient(provider), + agent.WithInstructions("You are a test agent."), + agent.WithModel("test-model"), + ) + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-suspended", &agent.Checkpoint{ + Version: 1, + Status: agent.CheckpointStatusSuspended, + AgentName: "test-agent", + Messages: []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "Hello"}}, + }, + { + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: "Working on it..."}}, + }, + }, + Usage: llm.Usage{InputTokens: 20, OutputTokens: 10}, + Turns: 1, + }) + require.NoError(t, err) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{ + "test-agent": ag, + }, + } + + result, err := agent.Restore( + context.Background(), + store, + "run-suspended", + registry, + ) + + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "Restored successfully.", result.FinalMessage().Text()) + assert.Equal(t, 2, result.Turns, "turns should include initial plus restored") + assert.Equal(t, 30, result.Usage.InputTokens, "usage should accumulate") + assert.Equal(t, 15, result.Usage.OutputTokens, "usage should accumulate") + }, + ) + + t.Run( + "awaiting approval without input returns InterruptedError", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Done."), + }, + } + + ag := agent.New( + "test-agent", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"dangerous_tool"}, + }), + ) + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-approval", &agent.Checkpoint{ + Version: 1, + Status: agent.CheckpointStatusAwaitingApproval, + AgentName: "test-agent", + Messages: []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "Do the thing"}}, + }, + }, + PendingToolCalls: []llm.ToolCall{ + { + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "dangerous_tool", + Arguments: `{}`, + }, + }, + }, + PendingApprovals: []llm.ToolCall{ + { + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "dangerous_tool", + Arguments: `{}`, + }, + }, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + Turns: 1, + }) + require.NoError(t, err) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{ + "test-agent": ag, + }, + } + + _, err = agent.Restore( + context.Background(), + store, + "run-approval", + registry, + ) + + require.Error(t, err) + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.Len(t, interrupted.PendingApprovals, 1) + assert.Equal(t, "dangerous_tool", interrupted.PendingApprovals[0].Function.Name) + assert.Equal(t, 1, interrupted.Turns) + assert.Equal(t, 10, interrupted.Usage.InputTokens) + }, + ) + + t.Run( + "awaiting approval with input resumes execution", + func(t *testing.T) { + t.Parallel() + + dangerousTool := agent.FunctionTool[struct{}]( + "dangerous_tool", + "A dangerous operation", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + return agent.ToolResult{Content: "executed"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Operation approved and completed."), + }, + } + + ag := agent.New( + "test-agent", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(dangerousTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"dangerous_tool"}, + }), + ) + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-approved", &agent.Checkpoint{ + Version: 1, + Status: agent.CheckpointStatusAwaitingApproval, + AgentName: "test-agent", + Messages: []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "Do the thing"}}, + }, + }, + PendingToolCalls: []llm.ToolCall{ + { + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "dangerous_tool", + Arguments: `{}`, + }, + }, + }, + PendingApprovals: []llm.ToolCall{ + { + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "dangerous_tool", + Arguments: `{}`, + }, + }, + }, + ApprovalInput: map[string]agent.ApprovalResult{ + "tc_1": {Approved: true}, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + Turns: 1, + }) + require.NoError(t, err) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{ + "test-agent": ag, + }, + } + + result, err := agent.Restore( + context.Background(), + store, + "run-approved", + registry, + ) + + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "Operation approved and completed.", result.FinalMessage().Text()) + }, + ) + + t.Run( + "nested approval rejects multiple inner checkpoints", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "test-agent", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-nested", &agent.Checkpoint{ + Version: 1, + Status: agent.CheckpointStatusAwaitingApproval, + AgentName: "test-agent", + Messages: []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "Do things"}}, + }, + }, + PendingToolCalls: []llm.ToolCall{ + { + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "inner_tool", + Arguments: `{}`, + }, + }, + }, + PendingApprovals: []llm.ToolCall{ + { + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "inner_tool", + Arguments: `{}`, + }, + }, + }, + InnerCheckpoints: map[string]*agent.Checkpoint{ + "tc_inner_1": { + Version: 1, + Status: agent.CheckpointStatusAwaitingApproval, + AgentName: "inner-agent-1", + }, + "tc_inner_2": { + Version: 1, + Status: agent.CheckpointStatusAwaitingApproval, + AgentName: "inner-agent-2", + }, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + Turns: 1, + }) + require.NoError(t, err) + + innerAgent1 := agent.New( + "inner-agent-1", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + innerAgent2 := agent.New( + "inner-agent-2", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{ + "test-agent": ag, + "inner-agent-1": innerAgent1, + "inner-agent-2": innerAgent2, + }, + } + + _, err = agent.Restore( + context.Background(), + store, + "run-nested", + registry, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "expected one inner checkpoint") + }, + ) + + t.Run( + "unknown agent name returns error", + func(t *testing.T) { + t.Parallel() + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-unknown", &agent.Checkpoint{ + Version: 1, + Status: agent.CheckpointStatusSuspended, + AgentName: "missing-agent", + }) + require.NoError(t, err) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{}, + } + + _, err = agent.Restore( + context.Background(), + store, + "run-unknown", + registry, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot resolve agent") + }, + ) + + t.Run( + "unknown checkpoint status returns error", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "test-agent", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + store := newMemoryCheckpointStore() + err := store.Save(context.Background(), "run-bad-status", &agent.Checkpoint{ + Version: 1, + Status: agent.CheckpointStatus("bogus"), + AgentName: "test-agent", + }) + require.NoError(t, err) + + registry := &simpleRegistry{ + agents: map[string]*agent.Agent{ + "test-agent": ag, + }, + } + + _, err = agent.Restore( + context.Background(), + store, + "run-bad-status", + registry, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "unknown checkpoint status") + }, + ) +}