The source headers, LICENSE files, and license metadata had drifted apart. Align the entire project to MIT: - Convert every source-file header to the MIT text across all comment styles (Go, TS, TSX, JS, MJS, SQL, CSS, GraphQL, shell), including SPDX-License-Identifier tags - Set the root and cookie-banner LICENSE files to the MIT text with a "MIT License" title line - Switch the package.json license fields, Docker image label, and cookie-banner README to MIT - Update docs and the genmodels header generator accordingly - Normalize copyright lines to a single format (Copyright (c) <year(s)> Probo Inc <hello@probo.com>.): unify the hello@getprobo.com and hello@probo.inc emails to hello@probo.com and the comma-separated years to a hyphenated range Genuine third-party references are intentionally left untouched: the Lucide icon attributions (Lucide is ISC) and the trivy dependency license allowlist. Signed-off-by: Sacha Al Himdani <sacha@probo.com>
267 lines
7.3 KiB
Go
267 lines
7.3 KiB
Go
// Copyright (c) 2026 Probo Inc <hello@probo.com>.
|
|
//
|
|
// Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
// of this software and associated documentation files (the "Software"), to deal
|
|
// in the Software without restriction, including without limitation the rights
|
|
// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
// copies of the Software, and to permit persons to whom the Software is
|
|
// furnished to do so, subject to the following conditions:
|
|
//
|
|
// The above copyright notice and this permission notice shall be included in
|
|
// all copies or substantial portions of the Software.
|
|
//
|
|
// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
// SOFTWARE.
|
|
|
|
package agent_test
|
|
|
|
import (
|
|
"context"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
"go.probo.inc/probo/pkg/agent"
|
|
"go.probo.inc/probo/pkg/llm"
|
|
)
|
|
|
|
// blockingProvider holds the ChatCompletion call until a release
|
|
// channel fires, so a test can race ctx cancellation against an
|
|
// in-flight LLM call.
|
|
type blockingProvider struct {
|
|
ready chan struct{}
|
|
release chan struct{}
|
|
response *llm.ChatCompletionResponse
|
|
|
|
mu sync.Mutex
|
|
calls int
|
|
ctxAtEnd error
|
|
}
|
|
|
|
func (p *blockingProvider) ChatCompletion(ctx context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
|
|
p.mu.Lock()
|
|
p.calls++
|
|
first := p.calls == 1
|
|
p.mu.Unlock()
|
|
|
|
if first {
|
|
close(p.ready)
|
|
<-p.release
|
|
p.mu.Lock()
|
|
p.ctxAtEnd = ctx.Err()
|
|
p.mu.Unlock()
|
|
}
|
|
|
|
return p.response, nil
|
|
}
|
|
|
|
func (p *blockingProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
|
|
return nil, assert.AnError
|
|
}
|
|
|
|
func TestRun_CtxCancelGracefulSuspend(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
t.Run(
|
|
"cancel before first turn suspends with empty checkpoint",
|
|
func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
provider := &mockProvider{
|
|
responses: []*llm.ChatCompletionResponse{
|
|
stopResponse("never called"),
|
|
},
|
|
}
|
|
|
|
ag := agent.New(
|
|
"assistant",
|
|
newTestClient(provider),
|
|
agent.WithModel("test-model"),
|
|
)
|
|
|
|
store := newMemoryCheckpointer()
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel()
|
|
|
|
_, err := ag.Run(
|
|
ctx,
|
|
[]llm.Message{userMessage("hi")},
|
|
agent.WithCheckpointer(store, "run-cancel"),
|
|
)
|
|
|
|
var se *agent.SuspendedError
|
|
require.ErrorAs(t, err, &se)
|
|
assert.Equal(t, 0, provider.calls, "LLM must not be invoked when ctx was already cancelled at entry")
|
|
|
|
// When a checkpointer is configured, the persistent store
|
|
// is the source of truth — the error itself doesn't carry a
|
|
// Checkpoint. Load from the store to verify.
|
|
cp, loadErr := store.Load(context.Background(), "run-cancel")
|
|
require.NoError(t, loadErr)
|
|
require.NotNil(t, cp, "checkpoint should be persisted before SuspendedError surfaces")
|
|
assert.Equal(t, agent.AgentStatusSuspended, cp.Status)
|
|
},
|
|
)
|
|
|
|
t.Run(
|
|
"cancel mid-run preserves the just-completed turn",
|
|
func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
|
|
provider := &mockProvider{
|
|
responses: []*llm.ChatCompletionResponse{
|
|
// First turn completes a tool call; the tool body
|
|
// then cancels ctx so the next turn-boundary check
|
|
// in coreLoop observes the cancellation.
|
|
{
|
|
Message: llm.Message{
|
|
Role: llm.RoleAssistant,
|
|
ToolCalls: []llm.ToolCall{{
|
|
ID: "tc_1",
|
|
Function: llm.FunctionCall{Name: "noop", Arguments: `{}`},
|
|
}},
|
|
},
|
|
FinishReason: llm.FinishReasonToolCalls,
|
|
},
|
|
stopResponse("never reached"),
|
|
},
|
|
}
|
|
|
|
noopTool := agent.FunctionTool[struct{}](
|
|
"noop",
|
|
"no-op",
|
|
func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
|
|
cancel()
|
|
return agent.ToolResult{Content: "ok"}, nil
|
|
},
|
|
)
|
|
|
|
ag := agent.New(
|
|
"assistant",
|
|
newTestClient(provider),
|
|
agent.WithModel("test-model"),
|
|
agent.WithTools(noopTool),
|
|
)
|
|
|
|
store := newMemoryCheckpointer()
|
|
|
|
_, err := ag.Run(
|
|
ctx,
|
|
[]llm.Message{userMessage("hi")},
|
|
agent.WithCheckpointer(store, "run-mid"),
|
|
)
|
|
|
|
var se *agent.SuspendedError
|
|
require.ErrorAs(t, err, &se)
|
|
assert.Equal(t, 1, provider.calls, "second LLM call must not fire after cancel")
|
|
|
|
cp, loadErr := store.Load(context.Background(), "run-mid")
|
|
require.NoError(t, loadErr)
|
|
require.NotNil(t, cp)
|
|
assert.Equal(t, agent.AgentStatusSuspended, cp.Status)
|
|
// The first LLM call completed; its output and the tool
|
|
// reply must be in the checkpointed messages so a Restore
|
|
// can resume from the next turn.
|
|
assert.Equal(t, 1, cp.Turns)
|
|
assert.GreaterOrEqual(t, len(cp.Messages), 3, "user + assistant tool-call + tool result")
|
|
},
|
|
)
|
|
|
|
t.Run(
|
|
"cancel during in-flight LLM call shields the call",
|
|
func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
// First response is a tool call so the loop iterates back
|
|
// to its turn-boundary cancel check after the LLM returns.
|
|
provider := &blockingProvider{
|
|
ready: make(chan struct{}),
|
|
release: make(chan struct{}),
|
|
response: &llm.ChatCompletionResponse{
|
|
Message: llm.Message{
|
|
Role: llm.RoleAssistant,
|
|
ToolCalls: []llm.ToolCall{{
|
|
ID: "tc_inflight",
|
|
Function: llm.FunctionCall{Name: "noop", Arguments: `{}`},
|
|
}},
|
|
},
|
|
FinishReason: llm.FinishReasonToolCalls,
|
|
},
|
|
}
|
|
|
|
noopTool := agent.FunctionTool[struct{}](
|
|
"noop",
|
|
"no-op",
|
|
func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
|
|
return agent.ToolResult{Content: "ok"}, nil
|
|
},
|
|
)
|
|
|
|
ag := agent.New(
|
|
"assistant",
|
|
newTestClient(provider),
|
|
agent.WithModel("test-model"),
|
|
agent.WithTools(noopTool),
|
|
)
|
|
|
|
store := newMemoryCheckpointer()
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
|
|
done := make(chan error, 1)
|
|
|
|
go func() {
|
|
_, err := ag.Run(
|
|
ctx,
|
|
[]llm.Message{userMessage("hi")},
|
|
agent.WithCheckpointer(store, "run-inflight"),
|
|
)
|
|
done <- err
|
|
}()
|
|
|
|
// Wait until the provider is parked inside ChatCompletion,
|
|
// then cancel ctx while the call is still in flight.
|
|
select {
|
|
case <-provider.ready:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("LLM call never started")
|
|
}
|
|
|
|
cancel()
|
|
close(provider.release)
|
|
|
|
var err error
|
|
select {
|
|
case err = <-done:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("agent.Run did not return after release")
|
|
}
|
|
|
|
var se *agent.SuspendedError
|
|
require.ErrorAs(t, err, &se)
|
|
|
|
provider.mu.Lock()
|
|
assert.NoError(t, provider.ctxAtEnd, "ctx passed to LLM must remain non-cancellable so the call completes")
|
|
assert.Equal(t, 1, provider.calls, "second LLM call must not fire after cancel")
|
|
provider.mu.Unlock()
|
|
|
|
cp, loadErr := store.Load(context.Background(), "run-inflight")
|
|
require.NoError(t, loadErr)
|
|
require.NotNil(t, cp)
|
|
assert.Equal(t, agent.AgentStatusSuspended, cp.Status)
|
|
},
|
|
)
|
|
}
|