From d2aa69cc96c47db8dc3cf08fe4757d48fea66513 Mon Sep 17 00:00:00 2001 From: x x Date: Mon, 14 Sep 2026 18:50:36 +0800 Subject: [PATCH] =?UTF-8?q?fix(acp):=20=E4=BF=AE=E5=A4=8D=20prompt=20?= =?UTF-8?q?=E4=B8=8E=20steering=20=E7=9A=84=E5=90=AF=E5=8A=A8=E7=AB=9E?= =?UTF-8?q?=E6=80=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/acp/connection.go | 7 ++ internal/acp/prompts.go | 17 +++- internal/acp/steering_compatibility_test.go | 87 +++++++++++++++++++++ 3 files changed, 107 insertions(+), 4 deletions(-) diff --git a/internal/acp/connection.go b/internal/acp/connection.go index c73ba057..6304a8e1 100644 --- a/internal/acp/connection.go +++ b/internal/acp/connection.go @@ -71,6 +71,10 @@ func NewConnection(reader io.ReadCloser, writer io.WriteCloser, requestHandler R } func (c *Connection) Request(ctx context.Context, method string, params any, result any) error { + return c.request(ctx, method, params, result, nil) +} + +func (c *Connection) request(ctx context.Context, method string, params any, result any, dispatched func()) error { if c == nil { return errors.New("ACP connection is nil") } @@ -102,6 +106,9 @@ func (c *Connection) Request(ctx context.Context, method string, params any, res c.removePending(key) return err } + if dispatched != nil { + dispatched() + } select { case reply := <-response: diff --git a/internal/acp/prompts.go b/internal/acp/prompts.go index 83bb5805..c8385df5 100644 --- a/internal/acp/prompts.go +++ b/internal/acp/prompts.go @@ -5,6 +5,7 @@ import ( "errors" "log/slog" "strings" + "sync" "time" ) @@ -148,7 +149,14 @@ func (m *Manager) StartPromptBlocks(ctx context.Context, sessionID string, block m.finishRun(run, RunFailed, "", err) return PromptStartResult{}, err } - go m.runPrompt(runCtx, run, record, blocks) + // StartPrompt 返回 started 后,调用方可能立即 Cancel/Steer。必须先保证 + // session/prompt 已写入 ACP 连接,否则 cancel notification 可能抢在 prompt + // 前面到达 Adapter,被当成“当前没有 turn”直接消费,随后原 prompt 永久等待。 + dispatched := make(chan struct{}) + var dispatchOnce sync.Once + markDispatched := func() { dispatchOnce.Do(func() { close(dispatched) }) } + go m.runPrompt(runCtx, run, record, blocks, markDispatched) + <-dispatched return PromptStartResult{RunID: run.ID, SessionID: sessionID, Status: RunRunning, Disposition: "started", StartedAt: run.StartedAt}, nil } @@ -428,7 +436,8 @@ func (m *Manager) markSessionInterrupted(record SessionRecord, reason string) { } } -func (m *Manager) runPrompt(ctx context.Context, run *Run, record SessionRecord, blocks []ContentBlock) { +func (m *Manager) runPrompt(ctx context.Context, run *Run, record SessionRecord, blocks []ContentBlock, markDispatched func()) { + defer markDispatched() m.mu.RLock() process := m.process m.mu.RUnlock() @@ -439,10 +448,10 @@ func (m *Manager) runPrompt(ctx context.Context, run *Run, record SessionRecord, var response struct { StopReason string `json:"stopReason"` } - err := process.connection.Request(ctx, "session/prompt", map[string]any{ + err := process.connection.request(ctx, "session/prompt", map[string]any{ "sessionId": record.RemoteSessionID, "prompt": blocks, - }, &response) + }, &response, markDispatched) if err != nil { if errors.Is(ctx.Err(), context.Canceled) { m.finishRun(run, RunCancelled, "cancelled", nil) diff --git a/internal/acp/steering_compatibility_test.go b/internal/acp/steering_compatibility_test.go index 5c9a92c1..3e81f83d 100644 --- a/internal/acp/steering_compatibility_test.go +++ b/internal/acp/steering_compatibility_test.go @@ -2,9 +2,96 @@ package acp import ( "context" + "io" + "sync" "testing" + "time" ) +type gatedACPWriter struct { + io.WriteCloser + started chan struct{} + release chan struct{} + once sync.Once +} + +func (w *gatedACPWriter) Write(data []byte) (int, error) { + w.once.Do(func() { close(w.started) }) + <-w.release + return w.WriteCloser.Write(data) +} + +func TestStartPromptWaitsForPromptDispatchBeforeSteering(t *testing.T) { + workspace := t.TempDir() + manager, err := newTestManagerWithAgent(t.TempDir(), workspace, claudeAgentACPName, "0.64.2", "claude_steer_fallback") + if err != nil { + t.Fatal(err) + } + defer func() { _ = manager.Close() }() + + created, err := manager.NewSession(context.Background(), workspace, nil) + if err != nil { + t.Fatal(err) + } + + manager.mu.RLock() + connection := manager.process.connection + manager.mu.RUnlock() + connection.writeMu.Lock() + gate := &gatedACPWriter{ + WriteCloser: connection.writer, + started: make(chan struct{}), + release: make(chan struct{}), + } + connection.writer = gate + connection.writeMu.Unlock() + + type startResult struct { + result PromptStartResult + err error + } + started := make(chan startResult, 1) + go func() { + result, startErr := manager.StartPrompt(context.Background(), created.Session.ID, "original") + started <- startResult{result: result, err: startErr} + }() + + select { + case <-gate.started: + case <-time.After(5 * time.Second): + close(gate.release) + t.Fatal("session/prompt write did not start") + } + select { + case result := <-started: + close(gate.release) + t.Fatalf("StartPrompt returned before session/prompt was dispatched: result=%#v err=%v", result.result, result.err) + default: + } + close(gate.release) + + var original PromptStartResult + select { + case result := <-started: + if result.err != nil { + t.Fatal(result.err) + } + original = result.result + case <-time.After(5 * time.Second): + t.Fatal("StartPrompt did not return after session/prompt dispatch") + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + steering, err := manager.Steer(ctx, created.Session.ID, "STEERED") + if err != nil { + t.Fatal(err) + } + if steering["cancelledRunId"] != original.RunID { + t.Fatalf("steering cancelled run = %#v, want %q", steering["cancelledRunId"], original.RunID) + } +} + func TestClaudeSteeringCompatibilityVersionBoundary(t *testing.T) { tests := []struct { name string