Rename CheckpointStatus to AgentStatus
The status values describe the agent state, not the checkpoint data state. Signed-off-by: Aurélien Sibiril <81782+aureliensibiril@users.noreply.github.com>
This commit is contained in:
@@ -29,7 +29,7 @@ import (
|
|||||||
// that may have been active (including handoff targets).
|
// that may have been active (including handoff targets).
|
||||||
func Restore(
|
func Restore(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
store CheckpointStore,
|
store Checkpointer,
|
||||||
runID string,
|
runID string,
|
||||||
registry AgentRegistry,
|
registry AgentRegistry,
|
||||||
) (*Result, error) {
|
) (*Result, error) {
|
||||||
@@ -40,10 +40,6 @@ func Restore(
|
|||||||
if cp == nil {
|
if cp == nil {
|
||||||
return nil, fmt.Errorf("cannot restore: no checkpoint for run %s", runID)
|
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)
|
agent, err := registry.Agent(cp.AgentName)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("cannot resolve agent %q: %w", cp.AgentName, err)
|
return nil, fmt.Errorf("cannot resolve agent %q: %w", cp.AgentName, err)
|
||||||
@@ -56,17 +52,17 @@ func restoreCheckpoint(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
agent *Agent,
|
agent *Agent,
|
||||||
cp *Checkpoint,
|
cp *Checkpoint,
|
||||||
store CheckpointStore,
|
store Checkpointer,
|
||||||
runID string,
|
runID string,
|
||||||
registry AgentRegistry,
|
registry AgentRegistry,
|
||||||
) (*Result, error) {
|
) (*Result, error) {
|
||||||
emitHook(agent, func(h RunHooks) { h.OnRunRestore(ctx, agent, cp) })
|
emitHook(agent, func(h RunHooks) { h.OnRunRestore(ctx, agent, cp) })
|
||||||
|
|
||||||
switch cp.Status {
|
switch cp.Status {
|
||||||
case CheckpointStatusSuspended:
|
case AgentStatusSuspended:
|
||||||
return restoreSuspended(ctx, agent, cp, store, runID, registry)
|
return restoreSuspended(ctx, agent, cp, store, runID, registry)
|
||||||
|
|
||||||
case CheckpointStatusAwaitingApproval:
|
case AgentStatusAwaitingApproval:
|
||||||
return restoreAwaitingApproval(ctx, agent, cp, store, runID, registry)
|
return restoreAwaitingApproval(ctx, agent, cp, store, runID, registry)
|
||||||
|
|
||||||
default:
|
default:
|
||||||
@@ -78,7 +74,7 @@ func restoreSuspended(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
agent *Agent,
|
agent *Agent,
|
||||||
cp *Checkpoint,
|
cp *Checkpoint,
|
||||||
store CheckpointStore,
|
store Checkpointer,
|
||||||
runID string,
|
runID string,
|
||||||
registry AgentRegistry,
|
registry AgentRegistry,
|
||||||
) (*Result, error) {
|
) (*Result, error) {
|
||||||
@@ -94,7 +90,7 @@ func continueFromMessages(
|
|||||||
agent *Agent,
|
agent *Agent,
|
||||||
messages []llm.Message,
|
messages []llm.Message,
|
||||||
cp *Checkpoint,
|
cp *Checkpoint,
|
||||||
store CheckpointStore,
|
store Checkpointer,
|
||||||
runID string,
|
runID string,
|
||||||
) (*Result, error) {
|
) (*Result, error) {
|
||||||
messagesCopy := make([]llm.Message, len(messages))
|
messagesCopy := make([]llm.Message, len(messages))
|
||||||
@@ -121,7 +117,7 @@ func continueFromMessages(
|
|||||||
skipSessionLoad: true,
|
skipSessionLoad: true,
|
||||||
initialUsage: cp.Usage,
|
initialUsage: cp.Usage,
|
||||||
initialTurns: cp.Turns,
|
initialTurns: cp.Turns,
|
||||||
checkpointStore: store,
|
checkpointer: store,
|
||||||
runID: runID,
|
runID: runID,
|
||||||
toolUsedInRun: cp.ToolUsedInRun,
|
toolUsedInRun: cp.ToolUsedInRun,
|
||||||
},
|
},
|
||||||
@@ -132,7 +128,7 @@ func restoreNestedSuspended(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
agent *Agent,
|
agent *Agent,
|
||||||
cp *Checkpoint,
|
cp *Checkpoint,
|
||||||
store CheckpointStore,
|
store Checkpointer,
|
||||||
runID string,
|
runID string,
|
||||||
registry AgentRegistry,
|
registry AgentRegistry,
|
||||||
) (*Result, error) {
|
) (*Result, error) {
|
||||||
@@ -277,7 +273,7 @@ func restoreAwaitingApproval(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
agent *Agent,
|
agent *Agent,
|
||||||
cp *Checkpoint,
|
cp *Checkpoint,
|
||||||
store CheckpointStore,
|
store Checkpointer,
|
||||||
runID string,
|
runID string,
|
||||||
registry AgentRegistry,
|
registry AgentRegistry,
|
||||||
) (*Result, error) {
|
) (*Result, error) {
|
||||||
@@ -340,7 +336,7 @@ func restoreAwaitingApproval(
|
|||||||
runOpts{
|
runOpts{
|
||||||
callLLM: blockingCallLLM,
|
callLLM: blockingCallLLM,
|
||||||
onEvent: noopEvent,
|
onEvent: noopEvent,
|
||||||
checkpointStore: store,
|
checkpointer: store,
|
||||||
runID: runID,
|
runID: runID,
|
||||||
toolUsedInRun: cp.ToolUsedInRun,
|
toolUsedInRun: cp.ToolUsedInRun,
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -26,18 +26,18 @@ import (
|
|||||||
"go.probo.inc/probo/pkg/llm"
|
"go.probo.inc/probo/pkg/llm"
|
||||||
)
|
)
|
||||||
|
|
||||||
type memoryCheckpointStore struct {
|
type memoryCheckpointer struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
checkpoints map[string]*agent.Checkpoint
|
checkpoints map[string]*agent.Checkpoint
|
||||||
}
|
}
|
||||||
|
|
||||||
func newMemoryCheckpointStore() *memoryCheckpointStore {
|
func newMemoryCheckpointer() *memoryCheckpointer {
|
||||||
return &memoryCheckpointStore{
|
return &memoryCheckpointer{
|
||||||
checkpoints: make(map[string]*agent.Checkpoint),
|
checkpoints: make(map[string]*agent.Checkpoint),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *memoryCheckpointStore) Save(_ context.Context, runID string, cp *agent.Checkpoint) error {
|
func (s *memoryCheckpointer) Save(_ context.Context, runID string, cp *agent.Checkpoint) error {
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
defer s.mu.Unlock()
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
@@ -46,7 +46,7 @@ func (s *memoryCheckpointStore) Save(_ context.Context, runID string, cp *agent.
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *memoryCheckpointStore) Load(_ context.Context, runID string) (*agent.Checkpoint, error) {
|
func (s *memoryCheckpointer) Load(_ context.Context, runID string) (*agent.Checkpoint, error) {
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
defer s.mu.Unlock()
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
@@ -59,14 +59,6 @@ func (s *memoryCheckpointStore) Load(_ context.Context, runID string) (*agent.Ch
|
|||||||
return &clone, nil
|
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 {
|
type simpleRegistry struct {
|
||||||
agents map[string]*agent.Agent
|
agents map[string]*agent.Agent
|
||||||
}
|
}
|
||||||
@@ -87,7 +79,7 @@ func TestRestore(t *testing.T) {
|
|||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
registry := &simpleRegistry{agents: map[string]*agent.Agent{}}
|
registry := &simpleRegistry{agents: map[string]*agent.Agent{}}
|
||||||
|
|
||||||
_, err := agent.Restore(
|
_, err := agent.Restore(
|
||||||
@@ -102,41 +94,6 @@ func TestRestore(t *testing.T) {
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
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(
|
t.Run(
|
||||||
"suspended checkpoint restores and completes",
|
"suspended checkpoint restores and completes",
|
||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
@@ -155,10 +112,9 @@ func TestRestore(t *testing.T) {
|
|||||||
agent.WithModel("test-model"),
|
agent.WithModel("test-model"),
|
||||||
)
|
)
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
err := store.Save(context.Background(), "run-suspended", &agent.Checkpoint{
|
err := store.Save(context.Background(), "run-suspended", &agent.Checkpoint{
|
||||||
Version: 1,
|
Status: agent.AgentStatusSuspended,
|
||||||
Status: agent.CheckpointStatusSuspended,
|
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
Messages: []llm.Message{
|
Messages: []llm.Message{
|
||||||
{
|
{
|
||||||
@@ -217,10 +173,9 @@ func TestRestore(t *testing.T) {
|
|||||||
}),
|
}),
|
||||||
)
|
)
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
err := store.Save(context.Background(), "run-approval", &agent.Checkpoint{
|
err := store.Save(context.Background(), "run-approval", &agent.Checkpoint{
|
||||||
Version: 1,
|
Status: agent.AgentStatusAwaitingApproval,
|
||||||
Status: agent.CheckpointStatusAwaitingApproval,
|
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
Messages: []llm.Message{
|
Messages: []llm.Message{
|
||||||
{
|
{
|
||||||
@@ -303,10 +258,9 @@ func TestRestore(t *testing.T) {
|
|||||||
}),
|
}),
|
||||||
)
|
)
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
err := store.Save(context.Background(), "run-approved", &agent.Checkpoint{
|
err := store.Save(context.Background(), "run-approved", &agent.Checkpoint{
|
||||||
Version: 1,
|
Status: agent.AgentStatusAwaitingApproval,
|
||||||
Status: agent.CheckpointStatusAwaitingApproval,
|
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
Messages: []llm.Message{
|
Messages: []llm.Message{
|
||||||
{
|
{
|
||||||
@@ -370,10 +324,9 @@ func TestRestore(t *testing.T) {
|
|||||||
agent.WithModel("test-model"),
|
agent.WithModel("test-model"),
|
||||||
)
|
)
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
err := store.Save(context.Background(), "run-nested", &agent.Checkpoint{
|
err := store.Save(context.Background(), "run-nested", &agent.Checkpoint{
|
||||||
Version: 1,
|
Status: agent.AgentStatusAwaitingApproval,
|
||||||
Status: agent.CheckpointStatusAwaitingApproval,
|
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
Messages: []llm.Message{
|
Messages: []llm.Message{
|
||||||
{
|
{
|
||||||
@@ -401,13 +354,11 @@ func TestRestore(t *testing.T) {
|
|||||||
},
|
},
|
||||||
InnerCheckpoints: map[string]*agent.Checkpoint{
|
InnerCheckpoints: map[string]*agent.Checkpoint{
|
||||||
"tc_inner_1": {
|
"tc_inner_1": {
|
||||||
Version: 1,
|
Status: agent.AgentStatusAwaitingApproval,
|
||||||
Status: agent.CheckpointStatusAwaitingApproval,
|
|
||||||
AgentName: "inner-agent-1",
|
AgentName: "inner-agent-1",
|
||||||
},
|
},
|
||||||
"tc_inner_2": {
|
"tc_inner_2": {
|
||||||
Version: 1,
|
Status: agent.AgentStatusAwaitingApproval,
|
||||||
Status: agent.CheckpointStatusAwaitingApproval,
|
|
||||||
AgentName: "inner-agent-2",
|
AgentName: "inner-agent-2",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
@@ -452,10 +403,9 @@ func TestRestore(t *testing.T) {
|
|||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
err := store.Save(context.Background(), "run-unknown", &agent.Checkpoint{
|
err := store.Save(context.Background(), "run-unknown", &agent.Checkpoint{
|
||||||
Version: 1,
|
Status: agent.AgentStatusSuspended,
|
||||||
Status: agent.CheckpointStatusSuspended,
|
|
||||||
AgentName: "missing-agent",
|
AgentName: "missing-agent",
|
||||||
})
|
})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -487,10 +437,9 @@ func TestRestore(t *testing.T) {
|
|||||||
agent.WithModel("test-model"),
|
agent.WithModel("test-model"),
|
||||||
)
|
)
|
||||||
|
|
||||||
store := newMemoryCheckpointStore()
|
store := newMemoryCheckpointer()
|
||||||
err := store.Save(context.Background(), "run-bad-status", &agent.Checkpoint{
|
err := store.Save(context.Background(), "run-bad-status", &agent.Checkpoint{
|
||||||
Version: 1,
|
Status: agent.AgentStatus("bogus"),
|
||||||
Status: agent.CheckpointStatus("bogus"),
|
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
})
|
})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ type (
|
|||||||
skipSessionLoad bool
|
skipSessionLoad bool
|
||||||
initialUsage llm.Usage
|
initialUsage llm.Usage
|
||||||
initialTurns int
|
initialTurns int
|
||||||
checkpointStore CheckpointStore
|
checkpointer Checkpointer
|
||||||
runID string
|
runID string
|
||||||
toolUsedInRun bool
|
toolUsedInRun bool
|
||||||
}
|
}
|
||||||
@@ -77,9 +77,9 @@ type (
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
func WithCheckpointStore(store CheckpointStore, runID string) RunOption {
|
func WithCheckpointer(cp Checkpointer, runID string) RunOption {
|
||||||
return func(o *runOpts) {
|
return func(o *runOpts) {
|
||||||
o.checkpointStore = store
|
o.checkpointer = cp
|
||||||
o.runID = runID
|
o.runID = runID
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -222,12 +222,11 @@ func (s *loopState) finishRun(ctx context.Context, result *Result, err error) (*
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *loopState) buildCheckpoint(status CheckpointStatus) *Checkpoint {
|
func (s *loopState) buildCheckpoint(status AgentStatus) *Checkpoint {
|
||||||
msgsCopy := make([]llm.Message, len(s.messages))
|
msgsCopy := make([]llm.Message, len(s.messages))
|
||||||
copy(msgsCopy, s.messages)
|
copy(msgsCopy, s.messages)
|
||||||
|
|
||||||
return &Checkpoint{
|
return &Checkpoint{
|
||||||
Version: CheckpointVersion,
|
|
||||||
Status: status,
|
Status: status,
|
||||||
AgentName: s.agent.name,
|
AgentName: s.agent.name,
|
||||||
Messages: msgsCopy,
|
Messages: msgsCopy,
|
||||||
@@ -381,11 +380,11 @@ func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Messag
|
|||||||
if ch := stopSignalFrom(ctx); ch != nil {
|
if ch := stopSignalFrom(ctx); ch != nil {
|
||||||
select {
|
select {
|
||||||
case <-ch:
|
case <-ch:
|
||||||
cp := s.buildCheckpoint(CheckpointStatusSuspended)
|
cp := s.buildCheckpoint(AgentStatusSuspended)
|
||||||
se := &SuspendedError{RunID: s.opts.runID}
|
se := &SuspendedError{RunID: s.opts.runID}
|
||||||
|
|
||||||
if s.opts.checkpointStore != nil && s.opts.runID != "" {
|
if s.opts.checkpointer != nil && s.opts.runID != "" {
|
||||||
if saveErr := s.opts.checkpointStore.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
if saveErr := s.opts.checkpointer.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
||||||
s.logger.ErrorCtx(ctx, "cannot save suspension checkpoint", log.Error(saveErr))
|
s.logger.ErrorCtx(ctx, "cannot save suspension checkpoint", log.Error(saveErr))
|
||||||
se.Checkpoint = cp
|
se.Checkpoint = cp
|
||||||
}
|
}
|
||||||
@@ -552,14 +551,14 @@ func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Messag
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if se, ok := errors.AsType[*SuspendedError](err); ok {
|
if se, ok := errors.AsType[*SuspendedError](err); ok {
|
||||||
outerCP := s.buildCheckpoint(CheckpointStatusSuspended)
|
outerCP := s.buildCheckpoint(AgentStatusSuspended)
|
||||||
if se.Checkpoint != nil {
|
if se.Checkpoint != nil {
|
||||||
outerCP.AllToolCalls = se.Checkpoint.AllToolCalls
|
outerCP.AllToolCalls = se.Checkpoint.AllToolCalls
|
||||||
outerCP.InnerCheckpoints = se.Checkpoint.InnerCheckpoints
|
outerCP.InnerCheckpoints = se.Checkpoint.InnerCheckpoints
|
||||||
outerCP.CompletedCalls = se.Checkpoint.CompletedCalls
|
outerCP.CompletedCalls = se.Checkpoint.CompletedCalls
|
||||||
}
|
}
|
||||||
if s.opts.checkpointStore != nil && s.opts.runID != "" {
|
if s.opts.checkpointer != nil && s.opts.runID != "" {
|
||||||
if saveErr := s.opts.checkpointStore.Save(ctx, s.opts.runID, outerCP); saveErr != nil {
|
if saveErr := s.opts.checkpointer.Save(ctx, s.opts.runID, outerCP); saveErr != nil {
|
||||||
s.logger.ErrorCtx(ctx, "cannot save checkpoint", log.Error(saveErr))
|
s.logger.ErrorCtx(ctx, "cannot save checkpoint", log.Error(saveErr))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -576,11 +575,11 @@ func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Messag
|
|||||||
msgsCopy := make([]llm.Message, len(s.messages))
|
msgsCopy := make([]llm.Message, len(s.messages))
|
||||||
copy(msgsCopy, s.messages)
|
copy(msgsCopy, s.messages)
|
||||||
|
|
||||||
if s.opts.checkpointStore != nil && s.opts.runID != "" {
|
if s.opts.checkpointer != nil && s.opts.runID != "" {
|
||||||
cp := s.buildCheckpoint(CheckpointStatusAwaitingApproval)
|
cp := s.buildCheckpoint(AgentStatusAwaitingApproval)
|
||||||
cp.PendingToolCalls = nae.allToolCalls
|
cp.PendingToolCalls = nae.allToolCalls
|
||||||
cp.PendingApprovals = nae.pendingApprovals
|
cp.PendingApprovals = nae.pendingApprovals
|
||||||
if saveErr := s.opts.checkpointStore.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
if saveErr := s.opts.checkpointer.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
||||||
s.logger.ErrorCtx(ctx, "cannot save approval checkpoint", log.Error(saveErr))
|
s.logger.ErrorCtx(ctx, "cannot save approval checkpoint", log.Error(saveErr))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -610,16 +609,15 @@ func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Messag
|
|||||||
msgsCopy := make([]llm.Message, len(s.messages))
|
msgsCopy := make([]llm.Message, len(s.messages))
|
||||||
copy(msgsCopy, s.messages)
|
copy(msgsCopy, s.messages)
|
||||||
|
|
||||||
if s.opts.checkpointStore != nil && s.opts.runID != "" {
|
if s.opts.checkpointer != nil && s.opts.runID != "" {
|
||||||
cp := s.buildCheckpoint(CheckpointStatusAwaitingApproval)
|
cp := s.buildCheckpoint(AgentStatusAwaitingApproval)
|
||||||
cp.PendingToolCalls = nie.inner.ToolCalls
|
cp.PendingToolCalls = nie.inner.ToolCalls
|
||||||
cp.PendingApprovals = nie.inner.PendingApprovals
|
cp.PendingApprovals = nie.inner.PendingApprovals
|
||||||
cp.AllToolCalls = nie.allToolCalls
|
cp.AllToolCalls = nie.allToolCalls
|
||||||
cp.CompletedCalls = nie.completedCalls
|
cp.CompletedCalls = nie.completedCalls
|
||||||
cp.InnerCheckpoints = map[string]*Checkpoint{
|
cp.InnerCheckpoints = map[string]*Checkpoint{
|
||||||
nie.toolCallID: {
|
nie.toolCallID: {
|
||||||
Version: CheckpointVersion,
|
Status: AgentStatusAwaitingApproval,
|
||||||
Status: CheckpointStatusAwaitingApproval,
|
|
||||||
AgentName: nie.inner.Agent.name,
|
AgentName: nie.inner.Agent.name,
|
||||||
Messages: nie.inner.Messages,
|
Messages: nie.inner.Messages,
|
||||||
Usage: nie.inner.Usage,
|
Usage: nie.inner.Usage,
|
||||||
@@ -628,7 +626,7 @@ func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Messag
|
|||||||
PendingApprovals: nie.inner.PendingApprovals,
|
PendingApprovals: nie.inner.PendingApprovals,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
if saveErr := s.opts.checkpointStore.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
if saveErr := s.opts.checkpointer.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
||||||
s.logger.ErrorCtx(ctx, "cannot save nested approval checkpoint", log.Error(saveErr))
|
s.logger.ErrorCtx(ctx, "cannot save nested approval checkpoint", log.Error(saveErr))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -694,9 +692,9 @@ func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Messag
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Save incremental checkpoint after completed tool-call turn.
|
// Save incremental checkpoint after completed tool-call turn.
|
||||||
if s.opts.checkpointStore != nil && s.opts.runID != "" {
|
if s.opts.checkpointer != nil && s.opts.runID != "" {
|
||||||
cp := s.buildCheckpoint(CheckpointStatusSuspended)
|
cp := s.buildCheckpoint(AgentStatusSuspended)
|
||||||
if saveErr := s.opts.checkpointStore.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
if saveErr := s.opts.checkpointer.Save(ctx, s.opts.runID, cp); saveErr != nil {
|
||||||
s.logger.ErrorCtx(ctx, "cannot save checkpoint", log.Error(saveErr))
|
s.logger.ErrorCtx(ctx, "cannot save checkpoint", log.Error(saveErr))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1008,8 +1006,7 @@ func executeParallel(
|
|||||||
|
|
||||||
outerSE := &SuspendedError{
|
outerSE := &SuspendedError{
|
||||||
Checkpoint: &Checkpoint{
|
Checkpoint: &Checkpoint{
|
||||||
Version: CheckpointVersion,
|
Status: AgentStatusSuspended,
|
||||||
Status: CheckpointStatusSuspended,
|
|
||||||
AllToolCalls: toolCalls,
|
AllToolCalls: toolCalls,
|
||||||
InnerCheckpoints: innerCheckpoints,
|
InnerCheckpoints: innerCheckpoints,
|
||||||
CompletedCalls: completed,
|
CompletedCalls: completed,
|
||||||
@@ -1414,7 +1411,7 @@ func resumeWithOpts(ctx context.Context, interrupted *InterruptedError, input Re
|
|||||||
skipSessionLoad: true,
|
skipSessionLoad: true,
|
||||||
initialUsage: interrupted.Usage,
|
initialUsage: interrupted.Usage,
|
||||||
initialTurns: interrupted.Turns,
|
initialTurns: interrupted.Turns,
|
||||||
checkpointStore: ro.checkpointStore,
|
checkpointer: ro.checkpointer,
|
||||||
runID: ro.runID,
|
runID: ro.runID,
|
||||||
toolUsedInRun: ro.toolUsedInRun,
|
toolUsedInRun: ro.toolUsedInRun,
|
||||||
},
|
},
|
||||||
@@ -1497,7 +1494,7 @@ func resumeNested(ctx context.Context, interrupted *InterruptedError, input Resu
|
|||||||
skipSessionLoad: true,
|
skipSessionLoad: true,
|
||||||
initialUsage: outer.usage,
|
initialUsage: outer.usage,
|
||||||
initialTurns: outer.turns,
|
initialTurns: outer.turns,
|
||||||
checkpointStore: ro.checkpointStore,
|
checkpointer: ro.checkpointer,
|
||||||
runID: ro.runID,
|
runID: ro.runID,
|
||||||
toolUsedInRun: ro.toolUsedInRun,
|
toolUsedInRun: ro.toolUsedInRun,
|
||||||
},
|
},
|
||||||
|
|||||||
870
pkg/agentruntest/agent_run_supervisor_test.go
Normal file
870
pkg/agentruntest/agent_run_supervisor_test.go
Normal file
@@ -0,0 +1,870 @@
|
|||||||
|
// 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 agentruntest_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"os/signal"
|
||||||
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
"syscall"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.gearno.de/kit/log"
|
||||||
|
"go.gearno.de/kit/pg"
|
||||||
|
"go.probo.inc/probo/pkg/agent"
|
||||||
|
"go.probo.inc/probo/pkg/agentruntest"
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/llm"
|
||||||
|
"go.probo.inc/probo/pkg/probo"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Test helpers
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func testLogger() *log.Logger {
|
||||||
|
return log.NewLogger(log.WithFormat(log.FormatPretty))
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Mock LLM provider
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
type mockProvider struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
responses []*llm.ChatCompletionResponse
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
if m.calls >= len(m.responses) {
|
||||||
|
return nil, errors.New("no more mock responses")
|
||||||
|
}
|
||||||
|
resp := m.responses[m.calls]
|
||||||
|
m.calls++
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
|
||||||
|
return nil, errors.New("not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestClient(provider llm.Provider) *llm.Client {
|
||||||
|
return llm.NewClient(provider, "test")
|
||||||
|
}
|
||||||
|
|
||||||
|
func stopResponse(text string) *llm.ChatCompletionResponse {
|
||||||
|
return &llm.ChatCompletionResponse{
|
||||||
|
Model: "test-model",
|
||||||
|
Message: llm.Message{
|
||||||
|
Role: llm.RoleAssistant,
|
||||||
|
Parts: []llm.Part{llm.TextPart{Text: text}},
|
||||||
|
},
|
||||||
|
Usage: llm.Usage{InputTokens: 10, OutputTokens: 5},
|
||||||
|
FinishReason: llm.FinishReasonStop,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func toolCallResponse(toolCalls ...llm.ToolCall) *llm.ChatCompletionResponse {
|
||||||
|
return &llm.ChatCompletionResponse{
|
||||||
|
Model: "test-model",
|
||||||
|
Message: llm.Message{
|
||||||
|
Role: llm.RoleAssistant,
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
},
|
||||||
|
Usage: llm.Usage{InputTokens: 10, OutputTokens: 5},
|
||||||
|
FinishReason: llm.FinishReasonToolCalls,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Simple agent registry
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Test 3: Supervisor picks up a PENDING run and completes it
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func TestAgentRunSupervisor_PicksUpAndCompletes(t *testing.T) {
|
||||||
|
client := agentruntest.PGClient(t)
|
||||||
|
store := coredata.NewPGCheckpointer(client)
|
||||||
|
|
||||||
|
provider := &mockProvider{
|
||||||
|
responses: []*llm.ChatCompletionResponse{
|
||||||
|
stopResponse("Done."),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
ag := agent.New(
|
||||||
|
"echo-agent",
|
||||||
|
newTestClient(provider),
|
||||||
|
agent.WithModel("test-model"),
|
||||||
|
agent.WithInstructions("Reply with done."),
|
||||||
|
)
|
||||||
|
|
||||||
|
registry := &simpleRegistry{
|
||||||
|
agents: map[string]*agent.Agent{"echo-agent": ag},
|
||||||
|
}
|
||||||
|
|
||||||
|
run := agentruntest.InsertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"echo-agent",
|
||||||
|
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "go"}}}},
|
||||||
|
)
|
||||||
|
|
||||||
|
supervisor := probo.NewAgentRunSupervisor(
|
||||||
|
client,
|
||||||
|
store,
|
||||||
|
registry,
|
||||||
|
testLogger(),
|
||||||
|
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
||||||
|
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
go supervisor.Run(ctx)
|
||||||
|
|
||||||
|
// Poll until the run is completed.
|
||||||
|
require.Eventually(
|
||||||
|
t,
|
||||||
|
func() bool {
|
||||||
|
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
||||||
|
return err == nil && r.Status == coredata.AgentRunStatusCompleted
|
||||||
|
},
|
||||||
|
10*time.Second,
|
||||||
|
200*time.Millisecond,
|
||||||
|
"run should reach COMPLETED status",
|
||||||
|
)
|
||||||
|
|
||||||
|
completed := agentruntest.LoadAgentRun(t, client, run.ID)
|
||||||
|
assert.Equal(t, coredata.AgentRunStatusCompleted, completed.Status)
|
||||||
|
assert.NotNil(t, completed.Result)
|
||||||
|
assert.Nil(t, completed.Checkpoint, "checkpoint should be cleared after completion")
|
||||||
|
assert.Nil(t, completed.ErrorMessage)
|
||||||
|
assert.False(t, completed.StopRequested)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Test 4: Supervisor stop/resume cycle with checkpoint
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func TestAgentRunSupervisor_StopAndResume(t *testing.T) {
|
||||||
|
client := agentruntest.PGClient(t)
|
||||||
|
store := coredata.NewPGCheckpointer(client)
|
||||||
|
|
||||||
|
// The tool blocks until signaled, giving us time to set stop_requested.
|
||||||
|
toolReady := make(chan struct{})
|
||||||
|
toolRelease := make(chan struct{})
|
||||||
|
|
||||||
|
slowTool := agent.FunctionTool[struct{}](
|
||||||
|
"slow_work",
|
||||||
|
"Does slow work",
|
||||||
|
func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
|
||||||
|
close(toolReady)
|
||||||
|
<-toolRelease
|
||||||
|
return agent.ToolResult{Content: "work done"}, nil
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Provider sequence:
|
||||||
|
// Call 1: request tool call (first execution)
|
||||||
|
// Call 2: final stop response (after restoration)
|
||||||
|
provider := &mockProvider{
|
||||||
|
responses: []*llm.ChatCompletionResponse{
|
||||||
|
// First execution: LLM asks to call the tool.
|
||||||
|
toolCallResponse(llm.ToolCall{
|
||||||
|
ID: "tc_1",
|
||||||
|
Function: llm.FunctionCall{Name: "slow_work", Arguments: `{}`},
|
||||||
|
}),
|
||||||
|
// After resume: the incremental checkpoint saved after tool completion
|
||||||
|
// means restore continues with these messages; LLM returns final answer.
|
||||||
|
stopResponse("All done after resume."),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
ag := agent.New(
|
||||||
|
"worker-agent",
|
||||||
|
newTestClient(provider),
|
||||||
|
agent.WithModel("test-model"),
|
||||||
|
agent.WithTools(slowTool),
|
||||||
|
)
|
||||||
|
|
||||||
|
registry := &simpleRegistry{
|
||||||
|
agents: map[string]*agent.Agent{"worker-agent": ag},
|
||||||
|
}
|
||||||
|
|
||||||
|
run := agentruntest.InsertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"worker-agent",
|
||||||
|
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "do work"}}}},
|
||||||
|
)
|
||||||
|
|
||||||
|
supervisor := probo.NewAgentRunSupervisor(
|
||||||
|
client,
|
||||||
|
store,
|
||||||
|
registry,
|
||||||
|
testLogger(),
|
||||||
|
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
||||||
|
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
||||||
|
)
|
||||||
|
|
||||||
|
// --- Phase 1: Start and let the supervisor pick up the run ---
|
||||||
|
|
||||||
|
ctx1, cancel1 := context.WithTimeout(context.Background(), 15*time.Second)
|
||||||
|
defer cancel1()
|
||||||
|
|
||||||
|
go supervisor.Run(ctx1)
|
||||||
|
|
||||||
|
// Wait for the tool to start executing — this confirms the supervisor
|
||||||
|
// claimed the run and the agent called the tool.
|
||||||
|
select {
|
||||||
|
case <-toolReady:
|
||||||
|
case <-ctx1.Done():
|
||||||
|
t.Fatal("timed out waiting for tool to start")
|
||||||
|
}
|
||||||
|
|
||||||
|
// The run should now be RUNNING.
|
||||||
|
running := agentruntest.LoadAgentRun(t, client, run.ID)
|
||||||
|
assert.Equal(t, coredata.AgentRunStatusRunning, running.Status)
|
||||||
|
|
||||||
|
// Set stop_requested in the database WHILE the tool is still blocked.
|
||||||
|
// The supervisor polls for this on each tick and signals the run's
|
||||||
|
// stop channel.
|
||||||
|
err := client.WithConn(
|
||||||
|
context.Background(),
|
||||||
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
_, err := conn.Exec(
|
||||||
|
ctx,
|
||||||
|
"UPDATE agent_runs SET stop_requested = true WHERE id = $1",
|
||||||
|
run.ID.String(),
|
||||||
|
)
|
||||||
|
return err
|
||||||
|
},
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Give the supervisor at least one tick to poll stop requests and
|
||||||
|
// close the run's stop channel before the tool finishes.
|
||||||
|
time.Sleep(1 * time.Second)
|
||||||
|
|
||||||
|
// Now release the tool. After completion the coreLoop saves an
|
||||||
|
// incremental checkpoint and checks the stop signal at the next
|
||||||
|
// turn boundary — it should already be closed.
|
||||||
|
close(toolRelease)
|
||||||
|
|
||||||
|
// Wait for the checkpoint to appear. The supervisor leaves the row
|
||||||
|
// in RUNNING because SuspendedError triggers the "leaving for stale
|
||||||
|
// recovery" path.
|
||||||
|
require.Eventually(
|
||||||
|
t,
|
||||||
|
func() bool {
|
||||||
|
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
||||||
|
return err == nil && r.Checkpoint != nil
|
||||||
|
},
|
||||||
|
10*time.Second,
|
||||||
|
200*time.Millisecond,
|
||||||
|
"checkpoint should be saved after stop",
|
||||||
|
)
|
||||||
|
|
||||||
|
// Stop the first supervisor.
|
||||||
|
cancel1()
|
||||||
|
time.Sleep(500 * time.Millisecond)
|
||||||
|
|
||||||
|
// Verify checkpoint content.
|
||||||
|
cp, err := store.Load(context.Background(), run.ID.String())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, cp, "checkpoint must exist after suspension")
|
||||||
|
assert.Equal(t, agent.AgentStatusSuspended, cp.Status)
|
||||||
|
assert.Equal(t, "worker-agent", cp.AgentName)
|
||||||
|
assert.True(t, len(cp.Messages) > 0, "checkpoint should contain messages")
|
||||||
|
|
||||||
|
// --- Phase 2: Simulate resume by resetting to PENDING ---
|
||||||
|
|
||||||
|
err = client.WithConn(
|
||||||
|
context.Background(),
|
||||||
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
_, err := conn.Exec(
|
||||||
|
ctx,
|
||||||
|
`UPDATE agent_runs
|
||||||
|
SET status = 'PENDING',
|
||||||
|
stop_requested = false,
|
||||||
|
started_at = NULL,
|
||||||
|
lease_owner = NULL,
|
||||||
|
lease_expires_at = NULL,
|
||||||
|
updated_at = now()
|
||||||
|
WHERE id = $1`,
|
||||||
|
run.ID.String(),
|
||||||
|
)
|
||||||
|
return err
|
||||||
|
},
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Start a fresh supervisor to pick up the resumed run.
|
||||||
|
supervisor2 := probo.NewAgentRunSupervisor(
|
||||||
|
client,
|
||||||
|
store,
|
||||||
|
registry,
|
||||||
|
testLogger(),
|
||||||
|
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
||||||
|
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx2, cancel2 := context.WithTimeout(context.Background(), 15*time.Second)
|
||||||
|
defer cancel2()
|
||||||
|
|
||||||
|
go supervisor2.Run(ctx2)
|
||||||
|
|
||||||
|
// The resumed run should load the checkpoint, call Restore, get the
|
||||||
|
// second LLM response (stopResponse), and complete.
|
||||||
|
require.Eventually(
|
||||||
|
t,
|
||||||
|
func() bool {
|
||||||
|
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
||||||
|
return err == nil && r.Status == coredata.AgentRunStatusCompleted
|
||||||
|
},
|
||||||
|
10*time.Second,
|
||||||
|
200*time.Millisecond,
|
||||||
|
"run should reach COMPLETED after resume",
|
||||||
|
)
|
||||||
|
|
||||||
|
completed := agentruntest.LoadAgentRun(t, client, run.ID)
|
||||||
|
assert.Equal(t, coredata.AgentRunStatusCompleted, completed.Status)
|
||||||
|
assert.NotNil(t, completed.Result)
|
||||||
|
assert.Nil(t, completed.Checkpoint, "checkpoint should be cleared after completion")
|
||||||
|
assert.Nil(t, completed.ErrorMessage)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Test 5: SIGTERM battle test — realistic multi-turn security audit
|
||||||
|
// with parallel tool calls, long-running operations, thinking turns,
|
||||||
|
// and multiple kill/resume cycles.
|
||||||
|
//
|
||||||
|
// Simulated workflow (10 tool-call turns + 1 final response):
|
||||||
|
//
|
||||||
|
// Turn 0: [think] scan_repos (single, 800ms)
|
||||||
|
// Turn 1: [think] fetch_config ×3 (parallel, 300-500ms)
|
||||||
|
// Turn 2: [think] analyze (single, 1000ms)
|
||||||
|
// Turn 3: [think] check ×3 (parallel, 400-800ms)
|
||||||
|
// Turn 4: [think] deep_analysis (single, 1500ms — long running)
|
||||||
|
// Turn 5: [think] generate ×2 (parallel, 500-600ms)
|
||||||
|
// Turn 6: [think] cve_lookup (single, 700ms)
|
||||||
|
// Turn 7: [think] compile (single, 600ms)
|
||||||
|
// Turn 8: [think] validate + format (parallel, 300-400ms)
|
||||||
|
// Turn 9: [think] publish (single, 400ms)
|
||||||
|
// Turn 10: final response — "Security audit complete..."
|
||||||
|
//
|
||||||
|
// SIGTERM is sent 3 times at different points, each time interrupting
|
||||||
|
// during tool execution (sometimes single, sometimes parallel).
|
||||||
|
// After each kill the checkpoint is verified to show progressive
|
||||||
|
// accumulation. A final in-process resume runs the remaining turns
|
||||||
|
// to completion.
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
// workInput is the shared parameter type for all battle-test tools.
|
||||||
|
type workInput struct {
|
||||||
|
Task string `json:"task"`
|
||||||
|
DurationMs int `json:"duration_ms"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// battleTestResponses returns the full LLM response sequence for a
|
||||||
|
// simulated security-audit agent. Each tool-call turn includes
|
||||||
|
// thinking text so the checkpoint messages are realistic.
|
||||||
|
func battleTestResponses() []*llm.ChatCompletionResponse {
|
||||||
|
tc := func(id, name, task string, ms int) llm.ToolCall {
|
||||||
|
return llm.ToolCall{
|
||||||
|
ID: id,
|
||||||
|
Function: llm.FunctionCall{
|
||||||
|
Name: name,
|
||||||
|
Arguments: fmt.Sprintf(`{"task":%q,"duration_ms":%d}`, task, ms),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
think := func(text string, calls ...llm.ToolCall) *llm.ChatCompletionResponse {
|
||||||
|
return &llm.ChatCompletionResponse{
|
||||||
|
Model: "test-model",
|
||||||
|
Message: llm.Message{
|
||||||
|
Role: llm.RoleAssistant,
|
||||||
|
Parts: []llm.Part{llm.TextPart{Text: text}},
|
||||||
|
ToolCalls: calls,
|
||||||
|
},
|
||||||
|
Usage: llm.Usage{InputTokens: 50, OutputTokens: 30},
|
||||||
|
FinishReason: llm.FinishReasonToolCalls,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return []*llm.ChatCompletionResponse{
|
||||||
|
// Turn 0 — single long scan
|
||||||
|
think(
|
||||||
|
"I'll begin the security audit by scanning all repositories to identify codebases, dependency manifests, and access-control configurations.",
|
||||||
|
tc("tc_0_1", "scan", "scan_repos", 800),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 1 — 3 parallel fetches
|
||||||
|
think(
|
||||||
|
"Found 3 repositories: api-gateway, auth-service, data-pipeline. Fetching their configurations in parallel to save time.",
|
||||||
|
tc("tc_1_1", "fetch", "fetch_api_config", 300),
|
||||||
|
tc("tc_1_2", "fetch", "fetch_auth_config", 400),
|
||||||
|
tc("tc_1_3", "fetch", "fetch_data_config", 500),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 2 — single analysis
|
||||||
|
think(
|
||||||
|
"All configurations retrieved. Running a comprehensive vulnerability analysis against the OWASP Top-10 checklist.",
|
||||||
|
tc("tc_2_1", "analyze", "analyze_configs", 1000),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 3 — 3 parallel security checks
|
||||||
|
think(
|
||||||
|
"Analysis flagged several areas of concern. Running dependency audit, secret scanning, and IAM permission checks in parallel.",
|
||||||
|
tc("tc_3_1", "check", "check_dependencies", 600),
|
||||||
|
tc("tc_3_2", "check", "check_secrets", 400),
|
||||||
|
tc("tc_3_3", "check", "check_permissions", 800),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 4 — single very long deep-dive
|
||||||
|
think(
|
||||||
|
"Multiple issues found: 3 outdated dependencies with known CVEs, 2 overly permissive IAM roles. Performing a deep analysis on the critical findings to determine exploitability and blast radius.",
|
||||||
|
tc("tc_4_1", "analyze", "deep_analysis", 1500),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 5 — 2 parallel report sections
|
||||||
|
think(
|
||||||
|
"Deep analysis complete. auth-service uses deprecated TLS 1.1 and data-pipeline stores PII unencrypted. Generating the executive summary and detailed findings sections in parallel.",
|
||||||
|
tc("tc_5_1", "generate", "generate_summary", 500),
|
||||||
|
tc("tc_5_2", "generate", "generate_findings", 600),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 6 — single CVE lookup
|
||||||
|
think(
|
||||||
|
"Report sections drafted. Cross-referencing all findings against the NVD and GitHub Advisory databases for known CVE identifiers.",
|
||||||
|
tc("tc_6_1", "lookup", "cve_lookup", 700),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 7 — single compile
|
||||||
|
think(
|
||||||
|
"CVE-2026-1234 matches the auth-service TLS vulnerability (CVSS 9.1). Compiling all sections, references, and remediation steps into the final report.",
|
||||||
|
tc("tc_7_1", "compile", "compile_report", 600),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 8 — 2 parallel validation + formatting
|
||||||
|
think(
|
||||||
|
"Draft report assembled (12 pages). Running structural validation and PDF formatting concurrently.",
|
||||||
|
tc("tc_8_1", "validate", "validate_report", 400),
|
||||||
|
tc("tc_8_2", "format", "format_pdf", 300),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 9 — single publish
|
||||||
|
think(
|
||||||
|
"Validation passed, PDF formatted. Publishing the finalized audit report to the internal portal.",
|
||||||
|
tc("tc_9_1", "publish", "publish_report", 400),
|
||||||
|
),
|
||||||
|
|
||||||
|
// Turn 10 — final text response
|
||||||
|
{
|
||||||
|
Model: "test-model",
|
||||||
|
Message: llm.Message{
|
||||||
|
Role: llm.RoleAssistant,
|
||||||
|
Parts: []llm.Part{llm.TextPart{Text: "Security audit complete.\n\nFindings:\n- 3 critical (auth-service TLS 1.1, unencrypted PII, CVE-2026-1234)\n- 5 medium (outdated deps, permissive IAM)\n- 4 low (missing rate-limiting, verbose logging)\n\nFull report: https://audits.internal/report-2026-04"}},
|
||||||
|
},
|
||||||
|
Usage: llm.Usage{InputTokens: 50, OutputTokens: 30},
|
||||||
|
FinishReason: llm.FinishReasonStop,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// makeBattleTools creates the tool set for the battle test. Every tool
|
||||||
|
// shares the same handler that sleeps for the requested duration and
|
||||||
|
// records progress to a shared file.
|
||||||
|
func makeBattleTools(progressFile string) []agent.Tool {
|
||||||
|
var mu sync.Mutex
|
||||||
|
|
||||||
|
handler := func(_ context.Context, input workInput) (agent.ToolResult, error) {
|
||||||
|
// Simulate real work.
|
||||||
|
time.Sleep(time.Duration(input.DurationMs) * time.Millisecond)
|
||||||
|
|
||||||
|
// Record completion — written AFTER the sleep so the parent's
|
||||||
|
// step count reflects truly-finished work.
|
||||||
|
mu.Lock()
|
||||||
|
f, err := os.OpenFile(progressFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
|
||||||
|
if err != nil {
|
||||||
|
mu.Unlock()
|
||||||
|
return agent.ToolResult{}, err
|
||||||
|
}
|
||||||
|
fmt.Fprintln(f, input.Task)
|
||||||
|
f.Close()
|
||||||
|
mu.Unlock()
|
||||||
|
|
||||||
|
return agent.ToolResult{Content: fmt.Sprintf("completed: %s", input.Task)}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
names := []struct{ name, desc string }{
|
||||||
|
{"scan", "Scan repositories for audit targets"},
|
||||||
|
{"fetch", "Fetch configuration or source files"},
|
||||||
|
{"analyze", "Run vulnerability analysis"},
|
||||||
|
{"check", "Execute a specific security check"},
|
||||||
|
{"generate", "Generate a report section"},
|
||||||
|
{"lookup", "Query external vulnerability databases"},
|
||||||
|
{"compile", "Compile report sections into final document"},
|
||||||
|
{"validate", "Validate report structure"},
|
||||||
|
{"format", "Apply output formatting"},
|
||||||
|
{"publish", "Publish report to internal portal"},
|
||||||
|
}
|
||||||
|
|
||||||
|
tools := make([]agent.Tool, len(names))
|
||||||
|
for i, n := range names {
|
||||||
|
tools[i] = agent.FunctionTool[workInput](n.name, n.desc, handler)
|
||||||
|
}
|
||||||
|
|
||||||
|
return tools
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentRunSupervisor_SIGTERM(t *testing.T) {
|
||||||
|
// ---- Subprocess mode ----
|
||||||
|
if os.Getenv("TEST_SIGTERM_SUBPROCESS") == "1" {
|
||||||
|
runSIGTERMSubprocess()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Parent mode ----
|
||||||
|
client := agentruntest.PGClient(t)
|
||||||
|
store := coredata.NewPGCheckpointer(client)
|
||||||
|
|
||||||
|
run := agentruntest.InsertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"battle-agent",
|
||||||
|
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Run a full security audit on all repositories."}}}},
|
||||||
|
)
|
||||||
|
|
||||||
|
progressFile := filepath.Join(t.TempDir(), "progress")
|
||||||
|
|
||||||
|
// ---- Helpers ----
|
||||||
|
|
||||||
|
countSteps := func() int {
|
||||||
|
data, err := os.ReadFile(progressFile)
|
||||||
|
if err != nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
n := 0
|
||||||
|
for _, b := range data {
|
||||||
|
if b == '\n' {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
startSubprocess := func(skipResponses int) *exec.Cmd {
|
||||||
|
cmd := exec.Command(
|
||||||
|
os.Args[0],
|
||||||
|
"-test.run=^TestAgentRunSupervisor_SIGTERM$",
|
||||||
|
"-test.v",
|
||||||
|
)
|
||||||
|
cmd.Env = append(os.Environ(),
|
||||||
|
"TEST_SIGTERM_SUBPROCESS=1",
|
||||||
|
"TEST_SIGTERM_PROGRESS_FILE="+progressFile,
|
||||||
|
"TEST_SIGTERM_SKIP_RESPONSES="+strconv.Itoa(skipResponses),
|
||||||
|
)
|
||||||
|
cmd.Stdout = os.Stdout
|
||||||
|
cmd.Stderr = os.Stderr
|
||||||
|
require.NoError(t, cmd.Start())
|
||||||
|
return cmd
|
||||||
|
}
|
||||||
|
|
||||||
|
killAndWait := func(cmd *exec.Cmd) {
|
||||||
|
require.NoError(t, cmd.Process.Signal(syscall.SIGTERM))
|
||||||
|
err := cmd.Wait()
|
||||||
|
if err != nil {
|
||||||
|
var exitErr *exec.ExitError
|
||||||
|
if errors.As(err, &exitErr) {
|
||||||
|
t.Logf("subprocess exited: %v", exitErr)
|
||||||
|
} else {
|
||||||
|
t.Fatalf("subprocess error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
resetToPending := func() {
|
||||||
|
err := client.WithConn(
|
||||||
|
context.Background(),
|
||||||
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
_, err := conn.Exec(ctx, `
|
||||||
|
UPDATE agent_runs
|
||||||
|
SET status = 'PENDING',
|
||||||
|
stop_requested = false,
|
||||||
|
started_at = NULL,
|
||||||
|
lease_owner = NULL,
|
||||||
|
lease_expires_at = NULL,
|
||||||
|
updated_at = now()
|
||||||
|
WHERE id = $1`,
|
||||||
|
run.ID.String(),
|
||||||
|
)
|
||||||
|
return err
|
||||||
|
},
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
waitForSteps := func(target int) {
|
||||||
|
require.Eventually(
|
||||||
|
t,
|
||||||
|
func() bool { return countSteps() >= target },
|
||||||
|
30*time.Second,
|
||||||
|
100*time.Millisecond,
|
||||||
|
fmt.Sprintf("expected at least %d completed tool executions", target),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
verifyCheckpoint := func(phase int) *agent.Checkpoint {
|
||||||
|
cp, err := store.Load(context.Background(), run.ID.String())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, cp, "phase %d: checkpoint must exist", phase)
|
||||||
|
assert.Equal(t, agent.AgentStatusSuspended, cp.Status, "phase %d", phase)
|
||||||
|
assert.Equal(t, "battle-agent", cp.AgentName, "phase %d", phase)
|
||||||
|
assert.Greater(t, len(cp.Messages), 1, "phase %d: checkpoint should have messages", phase)
|
||||||
|
assert.Greater(t, cp.Turns, 0, "phase %d: checkpoint should have turns", phase)
|
||||||
|
t.Logf(
|
||||||
|
" checkpoint: %d messages, %d turns, usage=%+v",
|
||||||
|
len(cp.Messages), cp.Turns, cp.Usage,
|
||||||
|
)
|
||||||
|
return cp
|
||||||
|
}
|
||||||
|
|
||||||
|
// ============================================================
|
||||||
|
// Phase 1: SIGTERM during the parallel fetch (turn 1)
|
||||||
|
// Steps so far: turn0=1(scan) + turn1=3(fetch×3) = 4
|
||||||
|
// ============================================================
|
||||||
|
t.Log("=== Phase 1: SIGTERM after scan + parallel fetch (4 steps) ===")
|
||||||
|
cmd1 := startSubprocess(0)
|
||||||
|
waitForSteps(4)
|
||||||
|
killAndWait(cmd1)
|
||||||
|
|
||||||
|
steps1 := countSteps()
|
||||||
|
t.Logf(" %d tool executions completed", steps1)
|
||||||
|
require.GreaterOrEqual(t, steps1, 4)
|
||||||
|
|
||||||
|
cp1 := verifyCheckpoint(1)
|
||||||
|
require.GreaterOrEqual(t, cp1.Turns, 2, "should have completed at least turns 0-1")
|
||||||
|
|
||||||
|
resetToPending()
|
||||||
|
|
||||||
|
// ============================================================
|
||||||
|
// Phase 2: SIGTERM during the parallel security checks (turn 3)
|
||||||
|
// New steps: turn2=1(analyze) + turn3=3(check×3) = 4
|
||||||
|
// ============================================================
|
||||||
|
t.Log("=== Phase 2: SIGTERM after analyze + parallel checks (4 more steps) ===")
|
||||||
|
cmd2 := startSubprocess(cp1.Turns)
|
||||||
|
waitForSteps(steps1 + 4)
|
||||||
|
killAndWait(cmd2)
|
||||||
|
|
||||||
|
steps2 := countSteps()
|
||||||
|
t.Logf(" %d tool executions completed (total)", steps2)
|
||||||
|
require.GreaterOrEqual(t, steps2, steps1+4)
|
||||||
|
|
||||||
|
cp2 := verifyCheckpoint(2)
|
||||||
|
assert.Greater(t, cp2.Turns, cp1.Turns, "turns should grow")
|
||||||
|
assert.Greater(t, len(cp2.Messages), len(cp1.Messages), "messages should grow")
|
||||||
|
|
||||||
|
resetToPending()
|
||||||
|
|
||||||
|
// ============================================================
|
||||||
|
// Phase 3: SIGTERM during deep_analysis (turn 4, long-running)
|
||||||
|
// or after generate ×2 (turn 5)
|
||||||
|
// New steps: turn4=1(deep) + turn5=2(generate×2) = 3
|
||||||
|
// ============================================================
|
||||||
|
t.Log("=== Phase 3: SIGTERM during long-running deep analysis (3 more steps) ===")
|
||||||
|
cmd3 := startSubprocess(cp2.Turns)
|
||||||
|
waitForSteps(steps2 + 3)
|
||||||
|
killAndWait(cmd3)
|
||||||
|
|
||||||
|
steps3 := countSteps()
|
||||||
|
t.Logf(" %d tool executions completed (total)", steps3)
|
||||||
|
require.GreaterOrEqual(t, steps3, steps2+3)
|
||||||
|
|
||||||
|
cp3 := verifyCheckpoint(3)
|
||||||
|
assert.Greater(t, cp3.Turns, cp2.Turns, "turns should grow again")
|
||||||
|
assert.Greater(t, len(cp3.Messages), len(cp2.Messages), "messages should grow again")
|
||||||
|
assert.Greater(t, cp3.Usage.InputTokens, 0, "usage should accumulate")
|
||||||
|
assert.Greater(t, cp3.Usage.OutputTokens, 0, "usage should accumulate")
|
||||||
|
|
||||||
|
t.Logf(
|
||||||
|
" after 3 SIGTERM cycles: %d steps, %d turns, %d messages, usage=%+v",
|
||||||
|
steps3, cp3.Turns, len(cp3.Messages), cp3.Usage,
|
||||||
|
)
|
||||||
|
|
||||||
|
resetToPending()
|
||||||
|
|
||||||
|
// ============================================================
|
||||||
|
// Phase 4: final in-process resume — run remaining turns to
|
||||||
|
// completion (lookup, compile, validate+format, publish, done)
|
||||||
|
// ============================================================
|
||||||
|
t.Log("=== Phase 4: in-process resume to completion ===")
|
||||||
|
|
||||||
|
remaining := battleTestResponses()[cp3.Turns:]
|
||||||
|
t.Logf(" %d LLM responses remaining (turns %d–10)", len(remaining), cp3.Turns)
|
||||||
|
|
||||||
|
tools := makeBattleTools(progressFile)
|
||||||
|
|
||||||
|
resumeAgent := agent.New(
|
||||||
|
"battle-agent",
|
||||||
|
newTestClient(&mockProvider{responses: remaining}),
|
||||||
|
agent.WithModel("test-model"),
|
||||||
|
agent.WithTools(tools...),
|
||||||
|
agent.WithMaxTurns(25),
|
||||||
|
)
|
||||||
|
|
||||||
|
supervisor := probo.NewAgentRunSupervisor(
|
||||||
|
client,
|
||||||
|
store,
|
||||||
|
&simpleRegistry{agents: map[string]*agent.Agent{"battle-agent": resumeAgent}},
|
||||||
|
testLogger(),
|
||||||
|
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
||||||
|
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
go supervisor.Run(ctx)
|
||||||
|
|
||||||
|
require.Eventually(
|
||||||
|
t,
|
||||||
|
func() bool {
|
||||||
|
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
||||||
|
return err == nil && r.Status == coredata.AgentRunStatusCompleted
|
||||||
|
},
|
||||||
|
25*time.Second,
|
||||||
|
200*time.Millisecond,
|
||||||
|
"run should complete after final resume",
|
||||||
|
)
|
||||||
|
|
||||||
|
stepsFinal := countSteps()
|
||||||
|
final := agentruntest.LoadAgentRun(t, client, run.ID)
|
||||||
|
assert.Equal(t, coredata.AgentRunStatusCompleted, final.Status)
|
||||||
|
assert.NotNil(t, final.Result)
|
||||||
|
assert.Nil(t, final.Checkpoint, "checkpoint should be cleared")
|
||||||
|
assert.Nil(t, final.ErrorMessage)
|
||||||
|
assert.Contains(t, string(final.Result), "Security audit complete")
|
||||||
|
|
||||||
|
t.Logf(
|
||||||
|
" battle test done: %d total tool executions across 3 SIGTERM cycles + final resume",
|
||||||
|
stepsFinal,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// runSIGTERMSubprocess is the child-process entry point. It sets up a
|
||||||
|
// supervisor with the full security-audit agent (10 distinct tools,
|
||||||
|
// thinking text, parallel calls, varying durations) and handles
|
||||||
|
// SIGTERM via signal.NotifyContext — identical to production probod.
|
||||||
|
func runSIGTERMSubprocess() {
|
||||||
|
progressFile := os.Getenv("TEST_SIGTERM_PROGRESS_FILE")
|
||||||
|
skip, _ := strconv.Atoi(os.Getenv("TEST_SIGTERM_SKIP_RESPONSES"))
|
||||||
|
|
||||||
|
addr := os.Getenv("PROBO_TEST_PG_ADDR")
|
||||||
|
if addr == "" {
|
||||||
|
addr = "localhost:5432"
|
||||||
|
}
|
||||||
|
user := os.Getenv("PROBO_TEST_PG_USER")
|
||||||
|
if user == "" {
|
||||||
|
user = "probod"
|
||||||
|
}
|
||||||
|
password := os.Getenv("PROBO_TEST_PG_PASSWORD")
|
||||||
|
if password == "" {
|
||||||
|
password = "probod"
|
||||||
|
}
|
||||||
|
database := os.Getenv("PROBO_TEST_PG_DATABASE")
|
||||||
|
if database == "" {
|
||||||
|
database = "probod_test"
|
||||||
|
}
|
||||||
|
|
||||||
|
pgClient, err := pg.NewClient(
|
||||||
|
pg.WithAddr(addr),
|
||||||
|
pg.WithUser(user),
|
||||||
|
pg.WithPassword(password),
|
||||||
|
pg.WithDatabase(database),
|
||||||
|
pg.WithPoolSize(5),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "subprocess: cannot create pg client: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
defer pgClient.Close()
|
||||||
|
|
||||||
|
store := coredata.NewPGCheckpointer(pgClient)
|
||||||
|
tools := makeBattleTools(progressFile)
|
||||||
|
|
||||||
|
responses := battleTestResponses()
|
||||||
|
if skip > 0 && skip < len(responses) {
|
||||||
|
responses = responses[skip:]
|
||||||
|
}
|
||||||
|
|
||||||
|
ag := agent.New(
|
||||||
|
"battle-agent",
|
||||||
|
newTestClient(&mockProvider{responses: responses}),
|
||||||
|
agent.WithModel("test-model"),
|
||||||
|
agent.WithTools(tools...),
|
||||||
|
agent.WithMaxTurns(25),
|
||||||
|
)
|
||||||
|
|
||||||
|
supervisor := probo.NewAgentRunSupervisor(
|
||||||
|
pgClient,
|
||||||
|
store,
|
||||||
|
&simpleRegistry{agents: map[string]*agent.Agent{"battle-agent": ag}},
|
||||||
|
log.NewLogger(log.WithFormat(log.FormatPretty)),
|
||||||
|
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
||||||
|
probo.WithAgentRunSupervisorLeaseDuration(5*time.Second),
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGTERM)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
|
err = supervisor.Run(ctx)
|
||||||
|
if err != nil && !errors.Is(err, context.Canceled) {
|
||||||
|
fmt.Fprintf(os.Stderr, "subprocess: supervisor error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
os.Exit(0)
|
||||||
|
}
|
||||||
@@ -419,7 +419,7 @@ RETURNING
|
|||||||
|
|
||||||
// ClearCheckpoint is the explicit path for removing persisted checkpoint
|
// ClearCheckpoint is the explicit path for removing persisted checkpoint
|
||||||
// data. AgentRun.Update intentionally does not write checkpoint so status
|
// data. AgentRun.Update intentionally does not write checkpoint so status
|
||||||
// commits cannot erase a checkpoint saved by PGCheckpointStore.Save.
|
// commits cannot erase a checkpoint saved by PGCheckpointer.Save.
|
||||||
func (e *AgentRun) ClearCheckpoint(
|
func (e *AgentRun) ClearCheckpoint(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
tx pg.Tx,
|
tx pg.Tx,
|
||||||
@@ -651,20 +651,20 @@ func LoadRunningStopRequestedIDs(ctx context.Context, conn pg.Querier) ([]string
|
|||||||
return ids, nil
|
return ids, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// PGCheckpointStore implements agent.CheckpointStore backed by the
|
// PGCheckpointer implements agent.Checkpointer backed by the
|
||||||
// agent_runs table checkpoint column. It is supervisor-internal and
|
// agent_runs table checkpoint column. It is supervisor-internal and
|
||||||
// intentionally uses raw run IDs with no tenant scope; public service/API
|
// intentionally uses raw run IDs with no tenant scope; public service/API
|
||||||
// methods must load AgentRun through scoped coredata methods before invoking
|
// methods must load AgentRun through scoped coredata methods before invoking
|
||||||
// lifecycle transitions.
|
// lifecycle transitions.
|
||||||
type PGCheckpointStore struct {
|
type PGCheckpointer struct {
|
||||||
pg *pg.Client
|
pg *pg.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewPGCheckpointStore(pgClient *pg.Client) *PGCheckpointStore {
|
func NewPGCheckpointer(pgClient *pg.Client) *PGCheckpointer {
|
||||||
return &PGCheckpointStore{pg: pgClient}
|
return &PGCheckpointer{pg: pgClient}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *PGCheckpointStore) Save(ctx context.Context, runID string, cp *agent.Checkpoint) error {
|
func (s *PGCheckpointer) Save(ctx context.Context, runID string, cp *agent.Checkpoint) error {
|
||||||
data, err := marshalAgentCheckpoint(cp)
|
data, err := marshalAgentCheckpoint(cp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -694,7 +694,7 @@ func (s *PGCheckpointStore) Save(ctx context.Context, runID string, cp *agent.Ch
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *PGCheckpointStore) Load(ctx context.Context, runID string) (*agent.Checkpoint, error) {
|
func (s *PGCheckpointer) Load(ctx context.Context, runID string) (*agent.Checkpoint, error) {
|
||||||
var cp *agent.Checkpoint
|
var cp *agent.Checkpoint
|
||||||
|
|
||||||
err := s.pg.WithConn(
|
err := s.pg.WithConn(
|
||||||
@@ -730,10 +730,6 @@ func (s *PGCheckpointStore) Load(ctx context.Context, runID string) (*agent.Chec
|
|||||||
return fmt.Errorf("cannot unmarshal checkpoint: %w", err)
|
return fmt.Errorf("cannot unmarshal checkpoint: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if cp.Version != agent.CheckpointVersion {
|
|
||||||
return fmt.Errorf("cannot load checkpoint: unsupported checkpoint version %d", cp.Version)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -741,42 +737,12 @@ func (s *PGCheckpointStore) Load(ctx context.Context, runID string) (*agent.Chec
|
|||||||
return cp, err
|
return cp, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *PGCheckpointStore) Delete(ctx context.Context, runID string) error {
|
|
||||||
return s.pg.WithConn(
|
|
||||||
ctx,
|
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
q := `UPDATE agent_runs SET checkpoint = NULL, updated_at = now() WHERE id = @id;`
|
|
||||||
|
|
||||||
args := pgx.StrictNamedArgs{"id": runID}
|
|
||||||
|
|
||||||
_, err := conn.Exec(ctx, q, args)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("cannot delete checkpoint: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Delete is intentionally idempotent: callers use it for cleanup after
|
|
||||||
// completion, and a concurrently-cleared checkpoint is already the
|
|
||||||
// desired state.
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
func marshalAgentCheckpoint(cp *agent.Checkpoint) ([]byte, error) {
|
func marshalAgentCheckpoint(cp *agent.Checkpoint) ([]byte, error) {
|
||||||
if cp == nil {
|
if cp == nil {
|
||||||
return nil, fmt.Errorf("cannot marshal checkpoint: checkpoint is required")
|
return nil, fmt.Errorf("cannot marshal checkpoint: checkpoint is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
next := *cp
|
data, err := json.Marshal(cp)
|
||||||
if next.Version == 0 {
|
|
||||||
next.Version = agent.CheckpointVersion
|
|
||||||
}
|
|
||||||
|
|
||||||
if next.Version != agent.CheckpointVersion {
|
|
||||||
return nil, fmt.Errorf("cannot marshal checkpoint: unsupported checkpoint version %d", next.Version)
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := json.Marshal(&next)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("cannot marshal checkpoint: %w", err)
|
return nil, fmt.Errorf("cannot marshal checkpoint: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
122
pkg/coredata/pg_checkpointer_test.go
Normal file
122
pkg/coredata/pg_checkpointer_test.go
Normal file
@@ -0,0 +1,122 @@
|
|||||||
|
// 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 coredata_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.probo.inc/probo/pkg/agent"
|
||||||
|
"go.probo.inc/probo/pkg/agentruntest"
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/llm"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPGCheckpointer(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
client := agentruntest.PGClient(t)
|
||||||
|
store := coredata.NewPGCheckpointer(client)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
run := agentruntest.InsertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"test-agent",
|
||||||
|
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "hello"}}}},
|
||||||
|
)
|
||||||
|
runID := run.ID.String()
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"load returns nil when no checkpoint exists",
|
||||||
|
func(t *testing.T) {
|
||||||
|
cp, err := store.Load(ctx, runID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Nil(t, cp)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"save and load round-trip",
|
||||||
|
func(t *testing.T) {
|
||||||
|
original := &agent.Checkpoint{
|
||||||
|
Status: agent.AgentStatusSuspended,
|
||||||
|
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..."}}},
|
||||||
|
},
|
||||||
|
Usage: llm.Usage{InputTokens: 20, OutputTokens: 10},
|
||||||
|
Turns: 1,
|
||||||
|
ToolUsedInRun: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
err := store.Save(ctx, runID, original)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
loaded, err := store.Load(ctx, runID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, loaded)
|
||||||
|
|
||||||
|
assert.Equal(t, original.Status, loaded.Status)
|
||||||
|
assert.Equal(t, original.AgentName, loaded.AgentName)
|
||||||
|
assert.Equal(t, original.Usage, loaded.Usage)
|
||||||
|
assert.Equal(t, original.Turns, loaded.Turns)
|
||||||
|
assert.Equal(t, original.ToolUsedInRun, loaded.ToolUsedInRun)
|
||||||
|
assert.Len(t, loaded.Messages, 2)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"save overwrites previous checkpoint",
|
||||||
|
func(t *testing.T) {
|
||||||
|
updated := &agent.Checkpoint{
|
||||||
|
Status: agent.AgentStatusSuspended,
|
||||||
|
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..."}}},
|
||||||
|
{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "continue"}}},
|
||||||
|
},
|
||||||
|
Usage: llm.Usage{InputTokens: 30, OutputTokens: 15},
|
||||||
|
Turns: 2,
|
||||||
|
}
|
||||||
|
|
||||||
|
err := store.Save(ctx, runID, updated)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
loaded, err := store.Load(ctx, runID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, loaded)
|
||||||
|
assert.Equal(t, 2, loaded.Turns)
|
||||||
|
assert.Len(t, loaded.Messages, 3)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"save to nonexistent run returns error",
|
||||||
|
func(t *testing.T) {
|
||||||
|
cp := &agent.Checkpoint{
|
||||||
|
Status: agent.AgentStatusSuspended,
|
||||||
|
AgentName: "test-agent",
|
||||||
|
}
|
||||||
|
err := store.Save(ctx, "nonexistent-run-id", cp)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "not found")
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -33,7 +33,7 @@ import (
|
|||||||
type (
|
type (
|
||||||
AgentRunSupervisor struct {
|
AgentRunSupervisor struct {
|
||||||
pg *pg.Client
|
pg *pg.Client
|
||||||
store *coredata.PGCheckpointStore
|
store *coredata.PGCheckpointer
|
||||||
registry agent.AgentRegistry
|
registry agent.AgentRegistry
|
||||||
logger *log.Logger
|
logger *log.Logger
|
||||||
interval time.Duration
|
interval time.Duration
|
||||||
@@ -89,7 +89,7 @@ func WithAgentRunSupervisorMaxConcurrency(n int) AgentRunSupervisorOption {
|
|||||||
|
|
||||||
func NewAgentRunSupervisor(
|
func NewAgentRunSupervisor(
|
||||||
pgClient *pg.Client,
|
pgClient *pg.Client,
|
||||||
store *coredata.PGCheckpointStore,
|
store *coredata.PGCheckpointer,
|
||||||
registry agent.AgentRegistry,
|
registry agent.AgentRegistry,
|
||||||
logger *log.Logger,
|
logger *log.Logger,
|
||||||
opts ...AgentRunSupervisorOption,
|
opts ...AgentRunSupervisorOption,
|
||||||
@@ -286,7 +286,7 @@ func (s *AgentRunSupervisor) executeRun(ctx context.Context, run *coredata.Agent
|
|||||||
result, runErr = a.RunWithOpts(
|
result, runErr = a.RunWithOpts(
|
||||||
ctx,
|
ctx,
|
||||||
inputMsgs,
|
inputMsgs,
|
||||||
agent.WithCheckpointStore(s.store, runID),
|
agent.WithCheckpointer(s.store, runID),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user