From cc579b1ed48d51ae51441ef1930ffe8504da0479 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aur=C3=A9lien=20Sibiril?= <81782+aureliensibiril@users.noreply.github.com> Date: Fri, 24 Apr 2026 20:01:00 +0200 Subject: [PATCH] Parallelize PG checkpointer subtests MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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> --- pkg/coredata/pg_checkpointer_test.go | 124 ++++++++++++++++++++++++--- 1 file changed, 111 insertions(+), 13 deletions(-) diff --git a/pkg/coredata/pg_checkpointer_test.go b/pkg/coredata/pg_checkpointer_test.go index 322c7256c..0337741d1 100644 --- a/pkg/coredata/pg_checkpointer_test.go +++ b/pkg/coredata/pg_checkpointer_test.go @@ -23,6 +23,7 @@ import ( "go.probo.inc/probo/pkg/agent" "go.probo.inc/probo/pkg/agentruntest" "go.probo.inc/probo/pkg/coredata" + "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/llm" ) @@ -31,20 +32,20 @@ func TestPGCheckpointer(t *testing.T) { 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) + 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) assert.Nil(t, cp) }, @@ -53,6 +54,16 @@ func TestPGCheckpointer(t *testing.T) { t.Run( "save and load round-trip", 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.AgentStatusSuspended, AgentName: "test-agent", @@ -84,6 +95,23 @@ func TestPGCheckpointer(t *testing.T) { t.Run( "save overwrites previous checkpoint", 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{ Status: agent.AgentStatusSuspended, AgentName: "test-agent", @@ -96,8 +124,7 @@ func TestPGCheckpointer(t *testing.T) { Turns: 2, } - err := store.Save(ctx, runID, updated) - require.NoError(t, err) + require.NoError(t, store.Save(ctx, runID, updated)) loaded, err := store.Load(ctx, runID) 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( "save to nonexistent run returns error", 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{ Status: agent.AgentStatusSuspended, AgentName: "test-agent", } - err := store.Save(ctx, "nonexistent-run-id", cp) + err := store.Save(ctx, otherID.String(), cp) require.Error(t, err) assert.Contains(t, err.Error(), "not found") },