Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions internal/command/acp.go
Original file line number Diff line number Diff line change
Expand Up @@ -917,11 +917,11 @@ func (a *acpAgent) Prompt(ctx context.Context, params acp.PromptRequest) (acp.Pr
// Reset per-turn approval-reviewer state (denial circuit breaker) at the
// start of each user turn.
sess.approvalState.OnTurnStart()
resp := runner.Run(promptCtx, sess.ag, history, sess.h, sess.rec, sess.todoStore, sess.env.GoalStore, sess.tracer, sess.tokenUsage)
result := runner.Run(promptCtx, sess.ag, history, sess.h, sess.rec, sess.todoStore, sess.env.GoalStore, sess.tracer, sess.tokenUsage)

sess.mu.Lock()
if resp != "" {
sess.history = append(sess.history, &schema.Message{Role: schema.Assistant, Content: resp})
if len(result.Messages) > 0 {
sess.history = append(sess.history, result.Messages...)
}
sess.cancel = nil
sess.mu.Unlock()
Expand Down
33 changes: 17 additions & 16 deletions internal/command/interactive.go
Original file line number Diff line number Diff line change
Expand Up @@ -615,18 +615,19 @@ func (s *interactiveState) handlePrompt(userPrompt string) {
s.history = append(s.history, schema.UserMessage(userPrompt))
s.history = agent.DrainBgNotifications(s.bgManager, s.history)
s.approvalState.OnTurnStart() // reset the per-turn reviewer denial breaker
resp := runner.Run(runCtx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
if resp != "" {
s.history = append(s.history, &schema.Message{Role: schema.Assistant, Content: resp})
result := runner.Run(runCtx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
if len(result.Messages) > 0 {
s.history = append(s.history, result.Messages...)
}
s.history = agent.SyncSummarization(s.summCapture, s.history, s.rec)
s.handlePlanCompletion(resp)
s.handlePlanCompletion(result)
}

func (s *interactiveState) handlePlanCompletion(resp string) {
if s.agentMode != tui.ModePlanning || resp == "" {
func (s *interactiveState) handlePlanCompletion(planResult runner.RunResult) {
if s.agentMode != tui.ModePlanning || planResult.Response == "" || planResult.Err != nil {
return
}
resp := planResult.Response

s.planStore.Submit("Plan", resp)
config.Logger().Printf("[plan] plan submitted for review (%d chars)", len(resp))
Expand Down Expand Up @@ -657,12 +658,12 @@ func (s *interactiveState) handlePlanCompletion(resp string) {
s.rec.RecordUser(revisePrompt)
}
s.history = append(s.history, schema.UserMessage(revisePrompt))
newResp := runner.Run(s.runCtx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
if newResp != "" {
s.history = append(s.history, &schema.Message{Role: schema.Assistant, Content: newResp})
result := runner.Run(s.runCtx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
if len(result.Messages) > 0 {
s.history = append(s.history, result.Messages...)
}
s.history = agent.SyncSummarization(s.summCapture, s.history, s.rec)
s.handlePlanCompletion(newResp)
s.handlePlanCompletion(result)
return
}

Expand All @@ -687,9 +688,9 @@ func (s *interactiveState) handlePlanCompletion(resp string) {
s.rec.RecordUser(execPrompt)
}
s.history = append(s.history, schema.UserMessage(execPrompt))
execResp := runner.Run(s.ctx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
if execResp != "" {
s.history = append(s.history, &schema.Message{Role: schema.Assistant, Content: execResp})
result := runner.Run(s.ctx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
if len(result.Messages) > 0 {
s.history = append(s.history, result.Messages...)
}
s.history = agent.SyncSummarization(s.summCapture, s.history, s.rec)

Expand Down Expand Up @@ -995,12 +996,12 @@ func (s *interactiveState) runEventLoop(initialHistory []adk.Message, initialRes
s.runCtx = runCtx
s.agentRunning.Store(true)
s.history = append(s.history, schema.UserMessage(prompt))
resp := runner.Run(runCtx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
result := runner.Run(runCtx, s.ag, s.history, s.h, s.rec, s.env.TodoStore, s.env.GoalStore, s.langfuseTracer, s.agentTokenUsage)
runCancel()
s.runCtx = nil
s.agentRunning.Store(false)
if resp != "" {
s.history = append(s.history, &schema.Message{Role: schema.Assistant, Content: resp})
if len(result.Messages) > 0 {
s.history = append(s.history, result.Messages...)
}
s.history = agent.SyncSummarization(s.summCapture, s.history, s.rec)
}
Expand Down
28 changes: 28 additions & 0 deletions internal/command/interactive_plan_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
package command

import (
"context"
"testing"
"time"

"github.com/cnjack/jcode/internal/runner"
"github.com/cnjack/jcode/internal/tui"
)

func TestHandlePlanCompletionSkipsInterruptedRun(t *testing.T) {
state := &interactiveState{agentMode: tui.ModePlanning}
returned := make(chan struct{})
go func() {
state.handlePlanCompletion(runner.RunResult{
Response: "partial plan",
Err: context.Canceled,
})
close(returned)
}()

select {
case <-returned:
case <-time.After(time.Second):
t.Fatal("interrupted plan entered the approval flow")
}
}
156 changes: 151 additions & 5 deletions internal/runner/cancel_backfill_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,43 @@ func (m *cancelModel) WithTools([]*schema.ToolInfo) (einomodel.ToolCallingChatMo
return m, nil
}

// partialStreamCancelModel emits a visible text prefix and an incomplete tool
// call, then returns the context error after the handler cancels the turn.
type partialStreamCancelModel struct{}

func (m *partialStreamCancelModel) WithTools([]*schema.ToolInfo) (einomodel.ToolCallingChatModel, error) {
return m, nil
}

func (m *partialStreamCancelModel) Generate(context.Context, []*schema.Message, ...einomodel.Option) (*schema.Message, error) {
return nil, errors.New("Generate is not used: streaming is enabled")
}

func (m *partialStreamCancelModel) Stream(ctx context.Context, _ []*schema.Message, _ ...einomodel.Option) (*schema.StreamReader[*schema.Message], error) {
reader, writer := schema.Pipe[*schema.Message](1)
go func() {
defer writer.Close()
index := 0
if writer.Send(&schema.Message{
Role: schema.Assistant,
Content: "partial",
ToolCalls: []schema.ToolCall{{
Index: &index,
ID: "call-partial-1",
Function: schema.FunctionCall{
Name: "block",
Arguments: `{"note":"`,
},
}},
}, nil) {
return
}
<-ctx.Done()
writer.Send(nil, ctx.Err())
}()
return reader, nil
}

func (m *cancelModel) Generate(context.Context, []*schema.Message, ...einomodel.Option) (*schema.Message, error) {
return nil, errors.New("Generate is not used: streaming is enabled")
}
Expand Down Expand Up @@ -78,6 +115,23 @@ type cancelRecordingHandler struct {
doneErr error
}

type cancelOnTextHandler struct {
stubHandler
cancel context.CancelFunc
toolCalls int
toolResults int
doneErr error
}

func (h *cancelOnTextHandler) OnAgentText(string) { h.cancel() }
func (h *cancelOnTextHandler) OnToolCall(handler.ToolCallEvent) {
h.toolCalls++
}
func (h *cancelOnTextHandler) OnToolResult(handler.ToolResultEvent) {
h.toolResults++
}
func (h *cancelOnTextHandler) OnAgentDone(err error) { h.doneErr = err }

func (h *cancelRecordingHandler) OnToolCall(handler.ToolCallEvent) { h.cancel() }

func (h *cancelRecordingHandler) OnToolResult(ev handler.ToolResultEvent) {
Expand Down Expand Up @@ -127,15 +181,20 @@ func TestRunInnerCancellationBackfillsToolResult(t *testing.T) {
}
h := &cancelRecordingHandler{cancel: cancel, done: make(chan struct{})}

finished := make(chan bool, 1)
type runOutcome struct {
result RunResult
done bool
}
finished := make(chan runOutcome, 1)
go func() {
_, done := runInner(ctx, ag, []adk.Message{schema.UserMessage("go")}, h, rec)
finished <- done
result, done := runInner(ctx, ag, []adk.Message{schema.UserMessage("go")}, h, rec)
finished <- runOutcome{result: result, done: done}
}()

var outcome runOutcome
select {
case done := <-finished:
if !done {
case outcome = <-finished:
if !outcome.done {
t.Errorf("runInner done = false, want true after cancellation")
}
case <-time.After(10 * time.Second):
Expand All @@ -152,6 +211,9 @@ func TestRunInnerCancellationBackfillsToolResult(t *testing.T) {
if !errors.Is(h.doneErr, context.Canceled) {
t.Errorf("OnAgentDone err = %v, want context.Canceled", h.doneErr)
}
if !errors.Is(outcome.result.Err, context.Canceled) {
t.Errorf("RunResult.Err = %v, want context.Canceled", outcome.result.Err)
}

// The announced call got exactly one result (drain backfill or a folded
// framework result — either satisfies the invariant).
Expand All @@ -161,6 +223,15 @@ func TestRunInnerCancellationBackfillsToolResult(t *testing.T) {
if h.results[0].ToolCallID != "call-block-1" {
t.Errorf("result ToolCallID = %q, want call-block-1", h.results[0].ToolCallID)
}
if len(outcome.result.Messages) != 2 {
t.Fatalf("live history messages = %d, want assistant tool call + backfill result", len(outcome.result.Messages))
}
if outcome.result.Messages[0].Role != schema.Assistant || len(outcome.result.Messages[0].ToolCalls) != 1 {
t.Fatalf("live history first message = %#v, want assistant tool call", outcome.result.Messages[0])
}
if outcome.result.Messages[1].Role != schema.Tool || outcome.result.Messages[1].ToolCallID != "call-block-1" {
t.Fatalf("live history second message = %#v, want matching backfill result", outcome.result.Messages[1])
}

// The session on disk pairs the recorded call with a recorded result.
entries, err := session.LoadSession(rec.UUID())
Expand All @@ -184,6 +255,7 @@ func TestRunInnerCancellationBackfillsToolResult(t *testing.T) {
}

state := session.ReconstructState(entries)
assertMessagesEqual(t, state.History, outcome.result.Messages)
for i, m := range state.History {
if m.Role != schema.Assistant || len(m.ToolCalls) == 0 {
continue
Expand All @@ -199,3 +271,77 @@ func TestRunInnerCancellationBackfillsToolResult(t *testing.T) {
}
}
}

func TestRunInnerCancellationDropsIncompleteStreamingToolCall(t *testing.T) {
t.Setenv("HOME", t.TempDir())
ctx, cancel := context.WithCancel(context.Background())
defer cancel()

ag, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{
Name: "partial-cancel-test",
Description: "partial-cancel-test",
Instruction: "test",
Model: &partialStreamCancelModel{},
ToolsConfig: adk.ToolsConfig{ToolsNodeConfig: compose.ToolsNodeConfig{
Tools: []tool.BaseTool{newBlockingTool()},
}},
MaxIterations: 5,
})
if err != nil {
t.Fatalf("create agent: %v", err)
}

rec, err := session.NewRecorder("partial-cancel-test", "test", "test")
if err != nil {
t.Fatalf("create recorder: %v", err)
}
h := &cancelOnTextHandler{cancel: cancel}
type runOutcome struct {
result RunResult
done bool
}
finished := make(chan runOutcome, 1)
go func() {
result, done := runInner(ctx, ag, []adk.Message{schema.UserMessage("go")}, h, rec)
finished <- runOutcome{result: result, done: done}
}()

var result RunResult
var done bool
select {
case outcome := <-finished:
result, done = outcome.result, outcome.done
case <-time.After(10 * time.Second):
t.Fatal("runInner did not return after stream cancellation")
}

if !done {
t.Fatal("runInner done = false, want true after stream cancellation")
}
if !errors.Is(h.doneErr, context.Canceled) {
t.Fatalf("OnAgentDone err = %v, want context.Canceled", h.doneErr)
}
if !errors.Is(result.Err, context.Canceled) {
t.Fatalf("RunResult.Err = %v, want context.Canceled", result.Err)
}
if h.toolCalls != 0 || h.toolResults != 0 {
t.Fatalf("handler saw calls/results = %d/%d, want 0/0 for incomplete call", h.toolCalls, h.toolResults)
}
if result.Response != "partial" || len(result.Messages) != 1 {
t.Fatalf("result = %#v, want one partial text message", result)
}
if result.Messages[0].Role != schema.Assistant || result.Messages[0].Content != "partial" || len(result.Messages[0].ToolCalls) != 0 {
t.Fatalf("result message = %#v, want text-only assistant", result.Messages[0])
}

entries, err := session.LoadSession(rec.UUID())
if err != nil {
t.Fatalf("load session: %v", err)
}
for _, entry := range entries {
if entry.Type == session.EntryToolCall || entry.Type == session.EntryToolResult {
t.Fatalf("session persisted incomplete tool entry: %#v", entry)
}
}
assertMessagesEqual(t, session.ReconstructState(entries).History, result.Messages)
}
Loading
Loading