Parallelize PG checkpointer subtests
Each subtest now inserts its own PENDING run and runs under t.Parallel(); shared state across subtests was the only reason they had to stay sequential. Also adds a round-trip test that exercises the approval-state fields (PendingToolCalls, PendingApprovals, ApprovalInput, AllToolCalls, InnerCheckpoints, CompletedCalls) to catch regressions where Save/Load drops nested or approval payloads. The nonexistent-run case now uses a valid GID in the same tenant so it reaches the row-not-found branch instead of short-circuiting on the tenant-scope guard. Signed-off-by: Aurélien Sibiril <81782+aureliensibiril@users.noreply.github.com>
This commit is contained in:
@@ -23,6 +23,7 @@ import (
|
|||||||
"go.probo.inc/probo/pkg/agent"
|
"go.probo.inc/probo/pkg/agent"
|
||||||
"go.probo.inc/probo/pkg/agentruntest"
|
"go.probo.inc/probo/pkg/agentruntest"
|
||||||
"go.probo.inc/probo/pkg/coredata"
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/gid"
|
||||||
"go.probo.inc/probo/pkg/llm"
|
"go.probo.inc/probo/pkg/llm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -31,20 +32,20 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
|
|
||||||
client := agentruntest.PGClient(t)
|
client := agentruntest.PGClient(t)
|
||||||
store := coredata.NewPGCheckpointer(client)
|
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(
|
t.Run(
|
||||||
"load returns nil when no checkpoint exists",
|
"load returns nil when no checkpoint exists",
|
||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
cp, err := store.Load(ctx, runID)
|
t.Parallel()
|
||||||
|
ctx := context.Background()
|
||||||
|
run := agentruntest.InsertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"test-agent",
|
||||||
|
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "hello"}}}},
|
||||||
|
)
|
||||||
|
|
||||||
|
cp, err := store.Load(ctx, run.ID.String())
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Nil(t, cp)
|
assert.Nil(t, cp)
|
||||||
},
|
},
|
||||||
@@ -53,6 +54,16 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Run(
|
t.Run(
|
||||||
"save and load round-trip",
|
"save and load round-trip",
|
||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
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()
|
||||||
|
|
||||||
original := &agent.Checkpoint{
|
original := &agent.Checkpoint{
|
||||||
Status: agent.AgentStatusSuspended,
|
Status: agent.AgentStatusSuspended,
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
@@ -84,6 +95,23 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Run(
|
t.Run(
|
||||||
"save overwrites previous checkpoint",
|
"save overwrites previous checkpoint",
|
||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
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()
|
||||||
|
|
||||||
|
first := &agent.Checkpoint{
|
||||||
|
Status: agent.AgentStatusSuspended,
|
||||||
|
AgentName: "test-agent",
|
||||||
|
Turns: 1,
|
||||||
|
}
|
||||||
|
require.NoError(t, store.Save(ctx, runID, first))
|
||||||
|
|
||||||
updated := &agent.Checkpoint{
|
updated := &agent.Checkpoint{
|
||||||
Status: agent.AgentStatusSuspended,
|
Status: agent.AgentStatusSuspended,
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
@@ -96,8 +124,7 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
Turns: 2,
|
Turns: 2,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := store.Save(ctx, runID, updated)
|
require.NoError(t, store.Save(ctx, runID, updated))
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
loaded, err := store.Load(ctx, runID)
|
loaded, err := store.Load(ctx, runID)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -107,14 +134,85 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"save and load preserves approval state",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
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()
|
||||||
|
|
||||||
|
original := &agent.Checkpoint{
|
||||||
|
Status: agent.AgentStatusAwaitingApproval,
|
||||||
|
AgentName: "test-agent",
|
||||||
|
Turns: 3,
|
||||||
|
PendingToolCalls: []llm.ToolCall{
|
||||||
|
{ID: "call-1", Function: llm.FunctionCall{Name: "send_email", Arguments: `{"to":"a@b.c"}`}},
|
||||||
|
},
|
||||||
|
PendingApprovals: []llm.ToolCall{
|
||||||
|
{ID: "call-1", Function: llm.FunctionCall{Name: "send_email", Arguments: `{"to":"a@b.c"}`}},
|
||||||
|
},
|
||||||
|
ApprovalInput: map[string]agent.ApprovalResult{
|
||||||
|
"call-2": {Approved: true},
|
||||||
|
},
|
||||||
|
AllToolCalls: []llm.ToolCall{
|
||||||
|
{ID: "call-1", Function: llm.FunctionCall{Name: "send_email"}},
|
||||||
|
{ID: "call-2", Function: llm.FunctionCall{Name: "log_event"}},
|
||||||
|
},
|
||||||
|
InnerCheckpoints: map[string]*agent.Checkpoint{
|
||||||
|
"call-3": {
|
||||||
|
Status: agent.AgentStatusSuspended,
|
||||||
|
AgentName: "inner-agent",
|
||||||
|
Turns: 1,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
CompletedCalls: []agent.CompletedCall{
|
||||||
|
{ToolCallID: "call-2", Result: agent.ToolResult{Content: "ok"}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, store.Save(ctx, runID, original))
|
||||||
|
|
||||||
|
loaded, err := store.Load(ctx, runID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, loaded)
|
||||||
|
assert.Len(t, loaded.PendingToolCalls, 1)
|
||||||
|
assert.Len(t, loaded.PendingApprovals, 1)
|
||||||
|
assert.True(t, loaded.ApprovalInput["call-2"].Approved)
|
||||||
|
assert.Len(t, loaded.AllToolCalls, 2)
|
||||||
|
require.Contains(t, loaded.InnerCheckpoints, "call-3")
|
||||||
|
assert.Equal(t, "inner-agent", loaded.InnerCheckpoints["call-3"].AgentName)
|
||||||
|
require.Len(t, loaded.CompletedCalls, 1)
|
||||||
|
assert.Equal(t, "ok", loaded.CompletedCalls[0].Result.Content)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
t.Run(
|
t.Run(
|
||||||
"save to nonexistent run returns error",
|
"save to nonexistent run returns error",
|
||||||
func(t *testing.T) {
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ctx := context.Background()
|
||||||
|
run := agentruntest.InsertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"test-agent",
|
||||||
|
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "hello"}}}},
|
||||||
|
)
|
||||||
|
// Build a syntactically valid GID for the same tenant but a
|
||||||
|
// different (unknown) entity so the tenant-scope check does
|
||||||
|
// not short-circuit with a parse error.
|
||||||
|
otherID := gid.New(run.ID.TenantID(), coredata.AgentRunEntityType)
|
||||||
|
|
||||||
cp := &agent.Checkpoint{
|
cp := &agent.Checkpoint{
|
||||||
Status: agent.AgentStatusSuspended,
|
Status: agent.AgentStatusSuspended,
|
||||||
AgentName: "test-agent",
|
AgentName: "test-agent",
|
||||||
}
|
}
|
||||||
err := store.Save(ctx, "nonexistent-run-id", cp)
|
err := store.Save(ctx, otherID.String(), cp)
|
||||||
require.Error(t, err)
|
require.Error(t, err)
|
||||||
assert.Contains(t, err.Error(), "not found")
|
assert.Contains(t, err.Error(), "not found")
|
||||||
},
|
},
|
||||||
|
|||||||
Reference in New Issue
Block a user