From b5eb9ded108a130e6c4465a8d401bf5be73987fb Mon Sep 17 00:00:00 2001 From: Haitao Pan Date: Tue, 28 Apr 2026 20:22:44 +0800 Subject: [PATCH] feat(acp): implement task tracking and artifact recording - Added TaskState and TaskKind tracking in SessionOrchestrator - Implemented ArtifactRecord to capture remote workspace hints and artifacts - Enhanced OpenCode adapter with session ID sanitization and validation - Added unit tests for remote workspace hints and invalid session handling --- internal/acp/orchestrator.go | 102 ++++++++++++++++++++++++ internal/acp/routing_test.go | 99 +++++++++++++++++++++++ internal/acp/rpc_handler.go | 9 +++ internal/acp/types.go | 52 ++++++++++++ internal/opencodeadapter/http_client.go | 22 ++--- internal/opencodeadapter/server.go | 4 + internal/opencodeadapter/server_test.go | 27 +++++++ 7 files changed, 305 insertions(+), 10 deletions(-) diff --git a/internal/acp/orchestrator.go b/internal/acp/orchestrator.go index 0452ce8..783deea 100644 --- a/internal/acp/orchestrator.go +++ b/internal/acp/orchestrator.go @@ -41,6 +41,21 @@ func (o *SessionOrchestrator) Process(ctx context.Context, method string, params sess.target = res.TargetID sess.provider = res.ProviderID sess.mode = res.TargetID + sess.control.ControlPlaneSessionID = sessionID + sess.control.ThreadID = threadID + sess.control.RequestedWorkingDir = strings.TrimSpace(shared.StringArg(params, "workingDirectory", "")) + sess.control.RemoteWorkingDirHint = strings.TrimSpace(shared.StringArg(params, "remoteWorkingDirectoryHint", "")) + sess.control.UpdatedAt = time.Now() + sess.task = QueuedTask{ + SessionID: sessionID, + ThreadID: threadID, + TurnID: turnID, + Provider: res.ProviderID, + Target: res.TargetID, + State: TaskStateRunning, + Kind: taskKindFromParams(params, res), + UpdatedAt: time.Now(), + } prompt := strings.TrimSpace(shared.StringArg(params, "taskPrompt", "")) if prompt != "" { sess.history = append(sess.history, "USER: "+prompt) @@ -86,6 +101,10 @@ func (o *SessionOrchestrator) Process(ctx context.Context, method string, params err = fmt.Errorf("unsupported session method: %s", method) } if err != nil { + sess.mu.Lock() + sess.task.State = TaskStateFailed + sess.task.UpdatedAt = time.Now() + sess.mu.Unlock() return nil, &shared.RPCError{Code: -32002, Message: "EXECUTION_FAILED: " + err.Error()} } @@ -205,6 +224,8 @@ func (o *SessionOrchestrator) normalizeResult(sess *session, result map[string]a if output != "" { sess.history = append(sess.history, "ASSISTANT: "+output) } + sess.task.State = TaskStateCompleted + sess.task.UpdatedAt = time.Now() sess.mu.Unlock() result["turnId"] = turnID @@ -231,6 +252,26 @@ func (o *SessionOrchestrator) normalizeResult(sess *session, result map[string]a result["error"] = "provider returned no displayable output" result["message"] = "provider returned no displayable output" } + if !parseBool(result["success"]) { + sess.mu.Lock() + sess.task.State = TaskStateFailed + sess.task.UpdatedAt = time.Now() + sess.mu.Unlock() + } + + artifactRecord := buildArtifactRecord(sess, result, output) + if artifactRecord.RemoteWorkingDirectory != "" { + result["remoteWorkingDirectory"] = artifactRecord.RemoteWorkingDirectory + } + if artifactRecord.RemoteWorkspaceRefKind != "" { + result["remoteWorkspaceRefKind"] = artifactRecord.RemoteWorkspaceRefKind + } + if artifactRecord.ResultSummary != "" && strings.TrimSpace(shared.StringArg(result, "resultSummary", "")) == "" { + result["resultSummary"] = artifactRecord.ResultSummary + } + if len(artifactRecord.Artifacts) > 0 { + result["artifacts"] = artifactRecord.Artifacts + } workingDirectory := shared.StringArg(params, "workingDirectory", "") routingParams := shared.AsMap(params["routing"]) @@ -248,6 +289,67 @@ func (o *SessionOrchestrator) normalizeResult(sess *session, result map[string]a return result } +func taskKindFromParams(params map[string]any, routing RoutingResult) TaskKind { + if parseBool(params["multiAgent"]) { + return TaskKindMultiAgent + } + if routing.TargetID == "gateway" { + return TaskKindGateway + } + return TaskKindSingleAgent +} + +func buildArtifactRecord(sess *session, result map[string]any, output string) ArtifactRecord { + record := ArtifactRecord{ + SessionID: sess.sessionID, + ThreadID: sess.threadID, + ResultSummary: strings.TrimSpace(output), + UpdatedAt: time.Now(), + } + if record.ResultSummary == "" { + record.ResultSummary = strings.TrimSpace(shared.StringArg(result, "resultSummary", "")) + } + if record.ResultSummary == "" { + record.ResultSummary = strings.TrimSpace(shared.StringArg(result, "summary", "")) + } + record.RemoteWorkingDirectory = strings.TrimSpace(shared.StringArg(result, "remoteWorkingDirectory", "")) + if record.RemoteWorkingDirectory == "" { + record.RemoteWorkingDirectory = strings.TrimSpace(sess.control.RemoteWorkingDirHint) + } + record.RemoteWorkspaceRefKind = strings.TrimSpace(shared.StringArg(result, "remoteWorkspaceRefKind", "")) + if record.RemoteWorkspaceRefKind == "" && record.RemoteWorkingDirectory != "" { + record.RemoteWorkspaceRefKind = "remotePath" + } + record.Artifacts = extractArtifactPayloads(result) + sess.mu.Lock() + sess.artifacts = record + sess.control.UpdatedAt = record.UpdatedAt + sess.mu.Unlock() + return record +} + +func extractArtifactPayloads(result map[string]any) []map[string]any { + rawArtifacts := result["artifacts"] + items, ok := rawArtifacts.([]any) + if !ok { + if typed, ok := rawArtifacts.([]map[string]any); ok { + copied := make([]map[string]any, 0, len(typed)) + for _, item := range typed { + copied = append(copied, item) + } + return copied + } + return nil + } + artifacts := make([]map[string]any, 0, len(items)) + for _, item := range items { + if mapped := shared.AsMap(item); len(mapped) > 0 { + artifacts = append(artifacts, mapped) + } + } + return artifacts +} + func (s *Server) getOrCreateSession(sessionID, threadID string) *session { s.mu.Lock() defer s.mu.Unlock() diff --git a/internal/acp/routing_test.go b/internal/acp/routing_test.go index 1d03c7d..dbefbbb 100644 --- a/internal/acp/routing_test.go +++ b/internal/acp/routing_test.go @@ -525,6 +525,105 @@ func TestExecuteSessionTaskAutoRoutingUsesBridgeProductionProviderOrder(t *testi } } +func TestExecuteSessionTaskKeepsRemoteWorkspaceHintOutOfLocalCWD(t *testing.T) { + workspaceDir := filepath.Join(t.TempDir(), "workspace") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatalf("create workspace: %v", err) + } + + server := NewServer() + providerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/acp/rpc" { + http.NotFound(w, r) + return + } + defer func() { _ = r.Body.Close() }() + var request map[string]any + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Fatalf("decode request: %v", err) + } + method := strings.TrimSpace(shared.StringArg(request, "method", "")) + result := map[string]any{ + "success": true, + "output": "hello", + "summary": "hello", + "message": "hello", + "remoteWorkingDirectory": "/owners/local/user/demo/threads/main", + "remoteWorkspaceRefKind": "remotePath", + "artifacts": []map[string]any{ + { + "relativePath": "notes/hello.txt", + "content": "hello artifact", + "contentType": "text/plain", + }, + }, + } + if method == "thread/start" { + result = map[string]any{"id": "codex-thread-1"} + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "jsonrpc": "2.0", + "id": request["id"], + "result": result, + }) + })) + defer providerServer.Close() + setTestBridgeProvider(server, syncedProvider{ + ProviderID: "codex", + Label: "Codex", + Endpoint: providerServer.URL, + Enabled: true, + }) + + response, rpcErr := server.executeSessionTask(task{ + req: shared.RPCRequest{ + Method: "session.start", + Params: map[string]any{ + "sessionId": "session-remote-hint", + "threadId": "thread-remote-hint", + "taskPrompt": "say hello", + "workingDirectory": workspaceDir, + "remoteWorkingDirectoryHint": "/owners/local/user/demo/threads/main", + "routing": map[string]any{ + "routingMode": "explicit", + "explicitExecutionTarget": "singleAgent", + "explicitProviderId": "codex", + }, + }, + }, + }) + if rpcErr != nil { + t.Fatalf("expected success, got rpc error: %v", rpcErr) + } + if got := response["remoteWorkingDirectory"]; got != "/owners/local/user/demo/threads/main" { + t.Fatalf("expected remote working directory in response, got %#v", response) + } + if got := response["remoteWorkspaceRefKind"]; got != "remotePath" { + t.Fatalf("expected remote workspace kind in response, got %#v", response) + } + if _, ok := response["artifacts"].([]map[string]any); !ok { + if _, ok := response["artifacts"].([]any); !ok { + t.Fatalf("expected artifacts payload, got %#v", response["artifacts"]) + } + } + sess := server.sessions["session-remote-hint"] + if sess == nil { + t.Fatal("expected session state to be retained") + } + if sess.control.RequestedWorkingDir != workspaceDir { + t.Fatalf("expected local requested cwd %q, got %q", workspaceDir, sess.control.RequestedWorkingDir) + } + if sess.control.RemoteWorkingDirHint != "/owners/local/user/demo/threads/main" { + t.Fatalf("expected remote hint retained, got %#v", sess.control) + } + if sess.task.Kind != TaskKindSingleAgent || sess.task.State != TaskStateCompleted { + t.Fatalf("expected completed single-agent task, got %#v", sess.task) + } + if sess.artifacts.RemoteWorkingDirectory != "/owners/local/user/demo/threads/main" { + t.Fatalf("expected artifact record to keep remote directory, got %#v", sess.artifacts) + } +} + func TestExecuteSessionTaskRequiresRouting(t *testing.T) { server := NewServer() _, rpcErr := server.executeSessionTask(task{ diff --git a/internal/acp/rpc_handler.go b/internal/acp/rpc_handler.go index b37189b..5b0a4f7 100644 --- a/internal/acp/rpc_handler.go +++ b/internal/acp/rpc_handler.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "strings" + "time" "xworkmate-bridge/internal/shared" ) @@ -68,6 +69,10 @@ func (s *Server) cancelSession(ctx context.Context, sessionID string) { sess, ok := s.sessions[sessionID] s.mu.RUnlock() if ok && sess != nil && sess.compat != nil { + sess.mu.Lock() + sess.task.State = TaskStateCancelled + sess.task.UpdatedAt = time.Now() + sess.mu.Unlock() _ = sess.compat.CancelSession(ctx, sessionID) } } @@ -78,6 +83,10 @@ func (s *Server) closeSession(ctx context.Context, sessionID string) { delete(s.sessions, sessionID) s.mu.Unlock() if ok && sess != nil && sess.compat != nil { + sess.mu.Lock() + sess.task.State = TaskStateCancelled + sess.task.UpdatedAt = time.Now() + sess.mu.Unlock() _ = sess.compat.CloseSession(ctx, sessionID) } } diff --git a/internal/acp/types.go b/internal/acp/types.go index 85e9fc6..9eaa846 100644 --- a/internal/acp/types.go +++ b/internal/acp/types.go @@ -2,11 +2,60 @@ package acp import ( "sync" + "time" "xworkmate-bridge/internal/gatewayruntime" "xworkmate-bridge/internal/memory" ) +type TaskState string + +const ( + TaskStateQueued TaskState = "queued" + TaskStateRunning TaskState = "running" + TaskStateCompleted TaskState = "completed" + TaskStateFailed TaskState = "failed" + TaskStateCancelled TaskState = "cancelled" +) + +type TaskKind string + +const ( + TaskKindSingleAgent TaskKind = "single-agent" + TaskKindGateway TaskKind = "gateway" + TaskKindMultiAgent TaskKind = "multi-agent" +) + +type ControlPlaneSession struct { + ControlPlaneSessionID string + ThreadID string + ProviderSessionID string + RequestedWorkingDir string + RemoteWorkingDirHint string + UpdatedAt time.Time +} + +type QueuedTask struct { + SessionID string + ThreadID string + TurnID string + Provider string + Target string + State TaskState + Kind TaskKind + UpdatedAt time.Time +} + +type ArtifactRecord struct { + SessionID string + ThreadID string + ResultSummary string + Artifacts []map[string]any + RemoteWorkingDirectory string + RemoteWorkspaceRefKind string + UpdatedAt time.Time +} + type session struct { sessionID string threadID string @@ -16,6 +65,9 @@ type session struct { compat ProviderCompat mu sync.Mutex history []string + control ControlPlaneSession + task QueuedTask + artifacts ArtifactRecord } type Server struct { diff --git a/internal/opencodeadapter/http_client.go b/internal/opencodeadapter/http_client.go index 0527656..2766072 100644 --- a/internal/opencodeadapter/http_client.go +++ b/internal/opencodeadapter/http_client.go @@ -85,13 +85,7 @@ func (c *opencodeHTTPClient) Call(method string, params map[string]any) (map[str } func sharedStringArg(params map[string]any, key, fallback string) string { - if params == nil { - return fallback - } - if value := strings.TrimSpace(fmt.Sprint(params[key])); value != "" { - return value - } - return fallback + return shared.StringArg(params, key, fallback) } func (c *opencodeHTTPClient) CreateSession(title string) (string, error) { @@ -247,10 +241,10 @@ func (c *opencodeHTTPClient) postSessionMessage(sessionID, prompt string, params func extractOpenCodeSessionID(value any) string { switch v := value.(type) { case string: - return strings.TrimSpace(v) + return sanitizeOpenCodeProviderSessionID(v) case map[string]any: - for _, key := range []string{"sessionId", "session_id", "id"} { - if sessionID := strings.TrimSpace(fmt.Sprint(v[key])); sessionID != "" { + for _, key := range []string{"sessionId", "sessionID", "session_id", "id"} { + if sessionID := sanitizeOpenCodeProviderSessionID(shared.StringArg(v, key, "")); sessionID != "" { return sessionID } } @@ -258,6 +252,14 @@ func extractOpenCodeSessionID(value any) string { return "" } +func sanitizeOpenCodeProviderSessionID(raw string) string { + sessionID := strings.TrimSpace(raw) + if sessionID == "" || sessionID == "" { + return "" + } + return sessionID +} + func extractOpenCodeText(value any) string { switch v := value.(type) { case string: diff --git a/internal/opencodeadapter/server.go b/internal/opencodeadapter/server.go index e9dffa3..7f9ffb5 100644 --- a/internal/opencodeadapter/server.go +++ b/internal/opencodeadapter/server.go @@ -259,6 +259,10 @@ func (s *Server) handleSessionRequest(method string, params map[string]any) map[ if err != nil { return map[string]any{"success": false, "provider": s.providerID, "mode": "single-agent", "error": err.Error()} } + upstreamSessionID = sanitizeOpenCodeProviderSessionID(upstreamSessionID) + if upstreamSessionID == "" { + return map[string]any{"success": false, "provider": s.providerID, "mode": "single-agent", "error": "opencode create session returned no session id"} + } state.upstreamSessionID = upstreamSessionID s.setSession(sessionID, state) } diff --git a/internal/opencodeadapter/server_test.go b/internal/opencodeadapter/server_test.go index f24fd09..e41e244 100644 --- a/internal/opencodeadapter/server_test.go +++ b/internal/opencodeadapter/server_test.go @@ -14,6 +14,10 @@ type stubOpenCodeClient struct { sendParams map[string]any } +type invalidSessionOpenCodeClient struct { + stubOpenCodeClient +} + func (s *stubOpenCodeClient) Initialize() (initializeResult, error) { s.initializeCalled++ return initializeResult{ProtocolVersion: 1}, nil @@ -37,6 +41,11 @@ func (s *stubOpenCodeClient) SendMessage(sessionID, prompt string, params map[st func (s *stubOpenCodeClient) Close() error { return nil } +func (s *invalidSessionOpenCodeClient) CreateSession(title string) (string, error) { + s.createTitle = title + return "", nil +} + func TestHandleSessionStartUsesCreateSessionAndSendMessage(t *testing.T) { client := &stubOpenCodeClient{} server := NewServer(client) @@ -79,6 +88,24 @@ func TestHandleSessionMessageReusesSession(t *testing.T) { } } +func TestHandleSessionStartRejectsInvalidProviderSessionID(t *testing.T) { + client := &invalidSessionOpenCodeClient{} + server := NewServer(client) + result := server.handleRequest(sharedRequest("session.start", map[string]any{ + "sessionId": "thread-1", + "taskPrompt": "hello", + })) + if success, _ := result["success"].(bool); success { + t.Fatalf("expected invalid upstream session to fail, got %#v", result) + } + if got := result["error"]; got != "opencode create session returned no session id" { + t.Fatalf("expected missing session id error, got %#v", result) + } + if client.sendSessionID != "" { + t.Fatalf("expected no follow-up message call, got session id %q", client.sendSessionID) + } +} + func sharedRequest(method string, params map[string]any) shared.RPCRequest { return shared.RPCRequest{ Method: method,