Add Restore function for agent checkpoint recovery
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>
This commit is contained in:
354
pkg/agent/restore.go
Normal file
354
pkg/agent/restore.go
Normal file
@@ -0,0 +1,354 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// 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
|
||||||
|
}
|
||||||
515
pkg/agent/restore_test.go
Normal file
515
pkg/agent/restore_test.go
Normal file
@@ -0,0 +1,515 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// 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")
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user