From bd2d7a9d2694f72eecaa2db0a9dd1b1b49780296 Mon Sep 17 00:00:00 2001 From: Haitao Pan Date: Sun, 29 Mar 2026 14:51:06 +0800 Subject: [PATCH] refactor(go-core): carve batch0 internal packages --- go/go_core/internal/acp/server.go | 636 +++++++++++++++++++ go/go_core/internal/shared/helpers.go | 81 +++ go/go_core/internal/shared/rpc.go | 108 ++++ go/go_core/internal/shared/tools.go | 397 ++++++++++++ go/go_core/internal/toolbridge/runner.go | 173 +++++ go/go_core/main.go | 777 +---------------------- go/go_core/main_tools.go | 482 +------------- 7 files changed, 1417 insertions(+), 1237 deletions(-) create mode 100644 go/go_core/internal/acp/server.go create mode 100644 go/go_core/internal/shared/helpers.go create mode 100644 go/go_core/internal/shared/rpc.go create mode 100644 go/go_core/internal/shared/tools.go create mode 100644 go/go_core/internal/toolbridge/runner.go diff --git a/go/go_core/internal/acp/server.go b/go/go_core/internal/acp/server.go new file mode 100644 index 00000000..e534b33e --- /dev/null +++ b/go/go_core/internal/acp/server.go @@ -0,0 +1,636 @@ +package acp + +import ( + "context" + "encoding/json" + "errors" + "flag" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "github.com/gorilla/websocket" + + "xworkmate/go_core/internal/shared" +) + +type session struct { + sessionID string + threadID string + mode string + provider string + history []string + seq int + cancel context.CancelFunc + closed bool +} + +type task struct { + req shared.RPCRequest + notify func(map[string]any) + done chan taskResult +} + +type taskResult struct { + response map[string]any + err *shared.RPCError +} + +type Server struct { + mu sync.Mutex + sessions map[string]*session + queues map[string]chan task +} + +var wsUpgrader = websocket.Upgrader{ + ReadBufferSize: 16 * 1024, + WriteBufferSize: 16 * 1024, + CheckOrigin: func(*http.Request) bool { + return true + }, +} + +func Serve(args []string) error { + flags := flag.NewFlagSet("serve", flag.ExitOnError) + listen := flags.String( + "listen", + shared.EnvOrDefault("ACP_LISTEN_ADDR", "127.0.0.1:8787"), + "ACP listen address", + ) + _ = flags.Parse(args) + + server := NewServer() + mux := http.NewServeMux() + mux.HandleFunc("/acp", server.HandleWebSocket) + mux.HandleFunc("/acp/rpc", server.HandleRPC) + + httpServer := &http.Server{ + Addr: strings.TrimSpace(*listen), + Handler: mux, + ReadTimeout: 30 * time.Second, + WriteTimeout: 5 * time.Minute, + IdleTimeout: 2 * time.Minute, + } + + if err := httpServer.ListenAndServe(); err != nil && + !errors.Is(err, http.ErrServerClosed) { + return fmt.Errorf("ACP server failed: %w", err) + } + return nil +} + +func NewServer() *Server { + return &Server{ + sessions: make(map[string]*session), + queues: make(map[string]chan task), + } +} + +func (s *Server) HandleWebSocket(w http.ResponseWriter, r *http.Request) { + conn, err := wsUpgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + + var writeMu sync.Mutex + notify := func(message map[string]any) { + writeMu.Lock() + defer writeMu.Unlock() + _ = conn.WriteJSON(message) + } + + for { + _, payload, err := conn.ReadMessage() + if err != nil { + return + } + request, err := shared.DecodeRPCRequest(payload) + if err != nil { + notify(shared.ErrorEnvelope(nil, -32700, err.Error())) + continue + } + response, rpcErr := s.handleRequest(request, notify) + if request.ID == nil { + continue + } + if rpcErr != nil { + notify(shared.ErrorEnvelope(request.ID, rpcErr.Code, rpcErr.Message)) + continue + } + notify(shared.ResultEnvelope(request.ID, response)) + } +} + +func (s *Server) HandleRPC(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + payload, err := io.ReadAll(r.Body) + if err != nil { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte("invalid body")) + return + } + request, err := shared.DecodeRPCRequest(payload) + if err != nil { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(err.Error())) + return + } + + accept := strings.ToLower(r.Header.Get("Accept")) + stream := strings.Contains(accept, "text/event-stream") + if stream { + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + } + + flusher, _ := w.(http.Flusher) + writeNotification := func(message map[string]any) { + if !stream { + return + } + shared.WriteSSE(w, message) + if flusher != nil { + flusher.Flush() + } + } + + response, rpcErr := s.handleRequest(request, writeNotification) + if request.ID == nil { + if stream { + _, _ = w.Write([]byte("data: [DONE]\n\n")) + } + return + } + if rpcErr != nil { + envelope := shared.ErrorEnvelope(request.ID, rpcErr.Code, rpcErr.Message) + if stream { + shared.WriteSSE(w, envelope) + if flusher != nil { + flusher.Flush() + } + return + } + w.WriteHeader(http.StatusOK) + _ = json.NewEncoder(w).Encode(envelope) + return + } + if stream { + shared.WriteSSE(w, shared.ResultEnvelope(request.ID, response)) + if flusher != nil { + flusher.Flush() + } + return + } + w.WriteHeader(http.StatusOK) + _ = json.NewEncoder(w).Encode(shared.ResultEnvelope(request.ID, response)) +} + +func (s *Server) handleRequest( + request shared.RPCRequest, + notify func(map[string]any), +) (map[string]any, *shared.RPCError) { + method := strings.TrimSpace(request.Method) + switch method { + case "acp.capabilities": + providers := shared.DetectACPProviders() + singleAgent := len(providers) > 0 + multiAgent := shared.BoolArg( + shared.EnvOrDefault("ACP_MULTI_AGENT_ENABLED", "true"), + true, + ) + result := map[string]any{ + "singleAgent": singleAgent, + "multiAgent": multiAgent, + "providers": providers, + "capabilities": map[string]any{ + "single_agent": singleAgent, + "multi_agent": multiAgent, + "providers": providers, + }, + } + return result, nil + case "session.start", "session.message": + params := request.Params + sessionID := strings.TrimSpace(shared.StringArg(params, "sessionId", "")) + if sessionID == "" { + return nil, &shared.RPCError{ + Code: -32602, + Message: "sessionId is required", + } + } + threadID := strings.TrimSpace( + shared.StringArg(params, "threadId", sessionID), + ) + if threadID == "" { + threadID = sessionID + } + if method == "session.start" { + s.resetSession(sessionID, threadID) + } + result, rpcErr := s.enqueue(threadID, task{ + req: request, + notify: notify, + done: make(chan taskResult, 1), + }) + if rpcErr != nil { + return nil, rpcErr + } + return result, nil + case "session.cancel": + params := request.Params + sessionID := strings.TrimSpace(shared.StringArg(params, "sessionId", "")) + if sessionID == "" { + return nil, &shared.RPCError{ + Code: -32602, + Message: "sessionId is required", + } + } + cancelled := s.cancelSession(sessionID) + return map[string]any{"accepted": true, "cancelled": cancelled}, nil + case "session.close": + params := request.Params + sessionID := strings.TrimSpace(shared.StringArg(params, "sessionId", "")) + if sessionID == "" { + return nil, &shared.RPCError{ + Code: -32602, + Message: "sessionId is required", + } + } + closed := s.closeSession(sessionID) + return map[string]any{"accepted": true, "closed": closed}, nil + default: + return nil, &shared.RPCError{ + Code: -32601, + Message: fmt.Sprintf("unknown method: %s", method), + } + } +} + +func (s *Server) enqueue(threadID string, task task) (map[string]any, *shared.RPCError) { + queue := s.ensureQueue(threadID) + queue <- task + result := <-task.done + return result.response, result.err +} + +func (s *Server) ensureQueue(threadID string) chan task { + s.mu.Lock() + defer s.mu.Unlock() + queue, ok := s.queues[threadID] + if ok { + return queue + } + queue = make(chan task, 32) + s.queues[threadID] = queue + go s.runQueue(queue) + return queue +} + +func (s *Server) runQueue(queue chan task) { + for task := range queue { + response, err := s.executeSessionTask(task) + task.done <- taskResult{response: response, err: err} + } +} + +func (s *Server) executeSessionTask(task task) (map[string]any, *shared.RPCError) { + params := task.req.Params + sessionID := strings.TrimSpace(shared.StringArg(params, "sessionId", "")) + threadID := strings.TrimSpace(shared.StringArg(params, "threadId", sessionID)) + mode := strings.TrimSpace(shared.StringArg(params, "mode", "single-agent")) + provider := strings.TrimSpace(shared.StringArg(params, "provider", "")) + if mode == "single-agent" && provider == "" { + provider = "codex" + } + + session := s.getOrCreateSession(sessionID, threadID) + session.mode = mode + if provider != "" { + session.provider = provider + } + + prompt := strings.TrimSpace(shared.StringArg(params, "taskPrompt", "")) + if prompt != "" { + session.history = append(session.history, prompt) + } + turnID := fmt.Sprintf("turn-%d", time.Now().UnixNano()) + + ctx, cancel := context.WithCancel(context.Background()) + s.setSessionCancel(sessionID, cancel) + defer s.clearSessionCancel(sessionID) + + notify := task.notify + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "event": "started", + "message": "session started", + "pending": true, + "error": false, + }) + + if mode == "multi-agent" { + result := s.runMultiAgent(ctx, session, params, turnID, notify) + if result.err != nil { + return nil, result.err + } + return result.response, nil + } + + result := s.runSingleAgent(ctx, session, params, turnID, notify) + if result.err != nil { + return nil, result.err + } + return result.response, nil +} + +func (s *Server) runSingleAgent( + ctx context.Context, + session *session, + params map[string]any, + turnID string, + notify func(map[string]any), +) taskResult { + provider := session.provider + if provider == "" { + provider = strings.TrimSpace(shared.StringArg(params, "provider", "codex")) + } + workingDirectory := strings.TrimSpace( + shared.StringArg(params, "workingDirectory", ""), + ) + model := strings.TrimSpace(shared.StringArg(params, "model", "")) + prompt := strings.TrimSpace(shared.StringArg(params, "taskPrompt", "")) + prompt = shared.AugmentPromptWithAttachments(prompt, params) + + output, err := shared.RunProviderCommand( + ctx, + provider, + model, + prompt, + workingDirectory, + ) + if err != nil { + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "event": "completed", + "message": err.Error(), + "pending": false, + "error": true, + }) + return taskResult{ + response: map[string]any{ + "success": false, + "error": err.Error(), + "turnId": turnID, + "mode": "single-agent", + "provider": provider, + }, + } + } + + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "delta", + "delta": output, + "pending": false, + "error": false, + }) + + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "event": "completed", + "message": "single-agent completed", + "pending": false, + "error": false, + }) + + return taskResult{ + response: map[string]any{ + "success": true, + "output": output, + "turnId": turnID, + "mode": "single-agent", + "provider": provider, + }, + } +} + +func (s *Server) runMultiAgent( + ctx context.Context, + session *session, + params map[string]any, + turnID string, + notify func(map[string]any), +) taskResult { + prompt := shared.ComposeHistoryPrompt(session.history) + if prompt == "" { + prompt = strings.TrimSpace(shared.StringArg(params, "taskPrompt", "")) + } + prompt = shared.AugmentPromptWithAttachments(prompt, params) + + baseURL := shared.NormalizeBaseURL( + shared.StringArg(params, "aiGatewayBaseUrl", ""), + ) + apiKey := strings.TrimSpace(shared.StringArg(params, "aiGatewayApiKey", "")) + model := strings.TrimSpace( + shared.StringArg( + params, + "model", + shared.EnvOrDefault("ACP_MULTI_AGENT_MODEL", "gpt-4o"), + ), + ) + if model == "" { + model = "gpt-4o" + } + + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "step", + "mode": "multi-agent", + "title": "Planner", + "message": "Preparing multi-agent run", + "pending": false, + "error": false, + "role": "architect", + "iteration": 1, + "score": 0, + }) + + if apiKey == "" { + errMsg := "aiGatewayApiKey is required for multi-agent mode" + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "mode": "multi-agent", + "message": errMsg, + "pending": false, + "error": true, + }) + return taskResult{ + response: map[string]any{ + "success": false, + "error": errMsg, + "turnId": turnID, + "mode": "multi-agent", + }, + } + } + + messages := []map[string]string{ + { + "role": "system", + "content": "You are a multi-agent coordinator. Return concise actionable output.", + }, + {"role": "user", "content": prompt}, + } + output, err := shared.CallOpenAICompatibleCtx( + ctx, + baseURL, + apiKey, + model, + messages, + ) + if err != nil { + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "mode": "multi-agent", + "message": err.Error(), + "pending": false, + "error": true, + }) + return taskResult{ + response: map[string]any{ + "success": false, + "error": err.Error(), + "turnId": turnID, + "mode": "multi-agent", + }, + } + } + + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "step", + "mode": "multi-agent", + "title": "Reviewer", + "message": output, + "pending": false, + "error": false, + "role": "tester", + "iteration": 1, + "score": 9, + }) + + return taskResult{ + response: map[string]any{ + "success": true, + "summary": output, + "finalScore": 9, + "iterations": 1, + "turnId": turnID, + "mode": "multi-agent", + }, + } +} + +func (s *Server) emitSessionUpdate( + session *session, + notify func(map[string]any), + turnID string, + payload map[string]any, +) { + if notify == nil { + return + } + s.mu.Lock() + session.seq++ + seq := session.seq + s.mu.Unlock() + params := map[string]any{ + "sessionId": session.sessionID, + "threadId": session.threadID, + "turnId": turnID, + "seq": seq, + } + for key, value := range payload { + params[key] = value + } + notify(shared.NotificationEnvelope("session.update", params)) +} + +func (s *Server) getOrCreateSession(sessionID, threadID string) *session { + s.mu.Lock() + defer s.mu.Unlock() + if session, ok := s.sessions[sessionID]; ok { + if threadID != "" { + session.threadID = threadID + } + session.closed = false + return session + } + session := &session{sessionID: sessionID, threadID: threadID} + s.sessions[sessionID] = session + return session +} + +func (s *Server) resetSession(sessionID, threadID string) { + s.mu.Lock() + defer s.mu.Unlock() + s.sessions[sessionID] = &session{ + sessionID: sessionID, + threadID: threadID, + history: []string{}, + } +} + +func (s *Server) setSessionCancel(sessionID string, cancel context.CancelFunc) { + s.mu.Lock() + defer s.mu.Unlock() + if session, ok := s.sessions[sessionID]; ok { + session.cancel = cancel + } +} + +func (s *Server) clearSessionCancel(sessionID string) { + s.mu.Lock() + defer s.mu.Unlock() + if session, ok := s.sessions[sessionID]; ok { + session.cancel = nil + } +} + +func (s *Server) cancelSession(sessionID string) bool { + s.mu.Lock() + session, ok := s.sessions[sessionID] + if !ok { + s.mu.Unlock() + return false + } + cancel := session.cancel + s.mu.Unlock() + if cancel != nil { + cancel() + return true + } + return false +} + +func (s *Server) closeSession(sessionID string) bool { + s.mu.Lock() + session, ok := s.sessions[sessionID] + if !ok { + s.mu.Unlock() + return false + } + cancel := session.cancel + session.closed = true + delete(s.sessions, sessionID) + s.mu.Unlock() + if cancel != nil { + cancel() + } + return true +} diff --git a/go/go_core/internal/shared/helpers.go b/go/go_core/internal/shared/helpers.go new file mode 100644 index 00000000..fc5fc0f5 --- /dev/null +++ b/go/go_core/internal/shared/helpers.go @@ -0,0 +1,81 @@ +package shared + +import ( + "fmt" + "os" + "strings" +) + +func NormalizeBaseURL(raw string) string { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "https://api.openai.com/v1" + } + if strings.HasSuffix(trimmed, "/v1") { + return trimmed + } + return strings.TrimRight(trimmed, "/") + "/v1" +} + +func EnvOrDefault(key, fallback string) string { + value := strings.TrimSpace(os.Getenv(key)) + if value == "" { + return fallback + } + return value +} + +func StringArg(arguments map[string]any, key, fallback string) string { + if arguments == nil { + return fallback + } + value, ok := arguments[key] + if !ok { + return fallback + } + text := strings.TrimSpace(fmt.Sprint(value)) + if text == "" || text == "" { + return fallback + } + return text +} + +func ListArg(arguments map[string]any, key string) []any { + if arguments == nil { + return nil + } + raw, ok := arguments[key] + if !ok || raw == nil { + return nil + } + if values, ok := raw.([]any); ok { + return values + } + if values, ok := raw.([]interface{}); ok { + return values + } + return nil +} + +func IntArg(raw string, fallback int) int { + var parsed int + if _, err := fmt.Sscanf(raw, "%d", &parsed); err != nil || parsed <= 0 { + return fallback + } + return parsed +} + +func BoolArg(raw string, fallback bool) bool { + trimmed := strings.TrimSpace(strings.ToLower(raw)) + if trimmed == "" { + return fallback + } + switch trimmed { + case "1", "true", "yes", "on": + return true + case "0", "false", "no", "off": + return false + default: + return fallback + } +} diff --git a/go/go_core/internal/shared/rpc.go b/go/go_core/internal/shared/rpc.go new file mode 100644 index 00000000..a6ab29d3 --- /dev/null +++ b/go/go_core/internal/shared/rpc.go @@ -0,0 +1,108 @@ +package shared + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "strings" +) + +type RPCRequest struct { + JSONRPC string `json:"jsonrpc,omitempty"` + ID any `json:"id,omitempty"` + Method string `json:"method,omitempty"` + Params map[string]any `json:"params,omitempty"` +} + +type RPCError struct { + Code int `json:"code"` + Message string `json:"message"` +} + +type ToolCallParams struct { + Name string `json:"name"` + Arguments map[string]any `json:"arguments"` +} + +func DecodeRPCRequest(payload []byte) (RPCRequest, error) { + var request RPCRequest + if err := json.Unmarshal(payload, &request); err != nil { + return RPCRequest{}, fmt.Errorf("invalid json: %w", err) + } + if strings.TrimSpace(request.Method) == "" { + return RPCRequest{}, errors.New("missing method") + } + if request.Params == nil { + request.Params = map[string]any{} + } + return request, nil +} + +func WriteSSE(w http.ResponseWriter, payload map[string]any) { + encoded, _ := json.Marshal(payload) + _, _ = fmt.Fprintf(w, "data: %s\n\n", encoded) +} + +func ResultEnvelope(id any, result map[string]any) map[string]any { + return map[string]any{ + "jsonrpc": "2.0", + "id": id, + "result": result, + } +} + +func ErrorEnvelope(id any, code int, message string) map[string]any { + return map[string]any{ + "jsonrpc": "2.0", + "id": id, + "error": map[string]any{ + "code": code, + "message": message, + }, + } +} + +func NotificationEnvelope(method string, params map[string]any) map[string]any { + return map[string]any{ + "jsonrpc": "2.0", + "method": method, + "params": params, + } +} + +func ErrorResponse(id any, code int, message string) map[string]any { + return map[string]any{ + "jsonrpc": "2.0", + "id": id, + "error": map[string]any{ + "code": code, + "message": message, + }, + } +} + +func ToolTextResult(id any, content string) map[string]any { + return map[string]any{ + "jsonrpc": "2.0", + "id": id, + "result": map[string]any{ + "content": []map[string]any{ + {"type": "text", "text": content}, + }, + }, + } +} + +func ToolErrorResult(id any, err error) map[string]any { + return map[string]any{ + "jsonrpc": "2.0", + "id": id, + "result": map[string]any{ + "content": []map[string]any{ + {"type": "text", "text": fmt.Sprintf("Error: %v", err)}, + }, + "isError": true, + }, + } +} diff --git a/go/go_core/internal/shared/tools.go b/go/go_core/internal/shared/tools.go new file mode 100644 index 00000000..af954e13 --- /dev/null +++ b/go/go_core/internal/shared/tools.go @@ -0,0 +1,397 @@ +package shared + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "os/exec" + "sort" + "strings" + "time" +) + +func DetectACPProviders() []string { + candidates := []struct { + provider string + envKey string + binary string + }{ + {provider: "codex", envKey: "ACP_CODEX_BIN", binary: "codex"}, + {provider: "opencode", envKey: "ACP_OPENCODE_BIN", binary: "opencode"}, + {provider: "claude", envKey: "ACP_CLAUDE_BIN", binary: "claude"}, + {provider: "gemini", envKey: "ACP_GEMINI_BIN", binary: "gemini"}, + } + providers := make([]string, 0, len(candidates)) + for _, candidate := range candidates { + binary := strings.TrimSpace(EnvOrDefault(candidate.envKey, candidate.binary)) + if binary == "" { + continue + } + if _, err := exec.LookPath(binary); err == nil { + providers = append(providers, candidate.provider) + } + } + sort.Strings(providers) + return providers +} + +func RunProviderCommand( + ctx context.Context, + provider, + model, + prompt, + workingDirectory string, +) (string, error) { + command, args := ResolveProviderCommand( + provider, + model, + prompt, + workingDirectory, + ) + if command == "" { + return "", fmt.Errorf("unsupported provider: %s", provider) + } + cmd := exec.CommandContext(ctx, command, args...) + if strings.TrimSpace(workingDirectory) != "" { + cmd.Dir = strings.TrimSpace(workingDirectory) + } + var stdout bytes.Buffer + var stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + if errors.Is(ctx.Err(), context.Canceled) { + return "", errors.New("run canceled") + } + message := strings.TrimSpace(stderr.String()) + if message == "" { + message = err.Error() + } + return "", fmt.Errorf("%s run failed: %s", provider, message) + } + output := strings.TrimSpace(stdout.String()) + if output == "" { + output = strings.TrimSpace(stderr.String()) + } + if output == "" { + return "", fmt.Errorf("%s returned empty output", provider) + } + return output, nil +} + +func ResolveProviderCommand( + provider, + model, + prompt, + cwd string, +) (string, []string) { + switch strings.TrimSpace(strings.ToLower(provider)) { + case "codex": + binary := strings.TrimSpace(EnvOrDefault("ACP_CODEX_BIN", "codex")) + args := []string{"exec", "--skip-git-repo-check", "--color", "never"} + if strings.TrimSpace(cwd) != "" { + args = append(args, "-C", strings.TrimSpace(cwd)) + } + if strings.TrimSpace(model) != "" { + args = append(args, "-m", strings.TrimSpace(model)) + } + args = append(args, prompt) + return binary, args + case "opencode": + binary := strings.TrimSpace(EnvOrDefault("ACP_OPENCODE_BIN", "opencode")) + args := []string{"run", "--format", "default"} + if strings.TrimSpace(cwd) != "" { + args = append(args, "--dir", strings.TrimSpace(cwd)) + } + if strings.TrimSpace(model) != "" { + args = append(args, "-m", strings.TrimSpace(model)) + } + args = append(args, prompt) + return binary, args + case "claude": + binary := strings.TrimSpace(EnvOrDefault("ACP_CLAUDE_BIN", "claude")) + if strings.TrimSpace(model) == "" { + return binary, []string{"-p", prompt} + } + return binary, []string{ + "--model", + strings.TrimSpace(model), + "-p", + prompt, + } + case "gemini": + binary := strings.TrimSpace(EnvOrDefault("ACP_GEMINI_BIN", "gemini")) + if strings.TrimSpace(model) == "" { + return binary, []string{"-p", prompt} + } + return binary, []string{ + "--model", + strings.TrimSpace(model), + "-p", + prompt, + } + default: + return "", nil + } +} + +func AugmentPromptWithAttachments(prompt string, params map[string]any) string { + attachmentsRaw := ListArg(params, "attachments") + if len(attachmentsRaw) == 0 { + return prompt + } + lines := make([]string, 0, len(attachmentsRaw)) + for _, raw := range attachmentsRaw { + entry, ok := raw.(map[string]any) + if !ok { + continue + } + name := strings.TrimSpace(StringArg(entry, "name", "attachment")) + path := strings.TrimSpace(StringArg(entry, "path", "")) + if path == "" { + continue + } + lines = append(lines, fmt.Sprintf("- %s: %s", name, path)) + } + if len(lines) == 0 { + return prompt + } + var builder strings.Builder + builder.WriteString("User-selected local attachments:\n") + builder.WriteString(strings.Join(lines, "\n")) + builder.WriteString("\n\n") + builder.WriteString(prompt) + return builder.String() +} + +func ComposeHistoryPrompt(history []string) string { + if len(history) == 0 { + return "" + } + var builder strings.Builder + for index, turn := range history { + builder.WriteString(fmt.Sprintf("## User Turn %d\n", index+1)) + builder.WriteString(turn) + builder.WriteString("\n\n") + } + return strings.TrimSpace(builder.String()) +} + +func CallOpenAICompatibleCtx( + ctx context.Context, + baseURL, + apiKey, + model string, + messages []map[string]string, +) (string, error) { + payload := map[string]any{ + "model": model, + "messages": messages, + "max_tokens": 4096, + "stream": false, + } + body, _ := json.Marshal(payload) + request, err := http.NewRequestWithContext( + ctx, + http.MethodPost, + strings.TrimRight(baseURL, "/")+"/chat/completions", + bytes.NewReader(body), + ) + if err != nil { + return "", err + } + request.Header.Set("Content-Type", "application/json") + request.Header.Set("Authorization", "Bearer "+apiKey) + + client := &http.Client{Timeout: 120 * time.Second} + response, err := client.Do(request) + if err != nil { + return "", err + } + defer response.Body.Close() + responseBody, err := io.ReadAll(response.Body) + if err != nil { + return "", err + } + if response.StatusCode < 200 || response.StatusCode >= 300 { + return "", fmt.Errorf( + "api error %d: %s", + response.StatusCode, + strings.TrimSpace(string(responseBody)), + ) + } + + var decoded map[string]any + if err := json.Unmarshal(responseBody, &decoded); err != nil { + return "", err + } + choices, _ := decoded["choices"].([]any) + if len(choices) == 0 { + return "", errors.New("missing choices in response") + } + choice, _ := choices[0].(map[string]any) + message, _ := choice["message"].(map[string]any) + content := strings.TrimSpace(fmt.Sprint(message["content"])) + if content == "" || content == "" { + return "", errors.New("empty response content") + } + return content, nil +} + +func HandleChatTool(arguments map[string]any) (string, error) { + apiKey := strings.TrimSpace(EnvOrDefault("LLM_API_KEY", "")) + if apiKey == "" { + return "", errors.New("LLM_API_KEY environment variable not set") + } + baseURL := NormalizeBaseURL( + EnvOrDefault("LLM_BASE_URL", "https://api.openai.com/v1"), + ) + model := StringArg(arguments, "model", EnvOrDefault("LLM_MODEL", "gpt-4o")) + prompt := strings.TrimSpace(StringArg(arguments, "prompt", "")) + if prompt == "" { + return "", errors.New("prompt is required") + } + system := strings.TrimSpace(StringArg(arguments, "system", "")) + + messages := make([]map[string]string, 0, 2) + if system != "" { + messages = append(messages, map[string]string{ + "role": "system", + "content": system, + }) + } + messages = append(messages, map[string]string{ + "role": "user", + "content": prompt, + }) + return CallOpenAICompatible(baseURL, apiKey, model, messages) +} + +func HandleClaudeReviewTool(arguments map[string]any) (string, error) { + prompt := strings.TrimSpace(StringArg(arguments, "prompt", "")) + if prompt == "" { + return "", errors.New("prompt is required") + } + model := strings.TrimSpace( + StringArg(arguments, "model", EnvOrDefault("CLAUDE_REVIEW_MODEL", "")), + ) + system := strings.TrimSpace( + StringArg(arguments, "system", EnvOrDefault("CLAUDE_REVIEW_SYSTEM", "")), + ) + tools := strings.TrimSpace( + StringArg(arguments, "tools", EnvOrDefault("CLAUDE_REVIEW_TOOLS", "")), + ) + timeout := IntArg(EnvOrDefault("CLAUDE_REVIEW_TIMEOUT_SEC", "600"), 600) + return RunClaudeReview( + prompt, + model, + system, + tools, + time.Duration(timeout)*time.Second, + ) +} + +func CallOpenAICompatible( + baseURL, + apiKey, + model string, + messages []map[string]string, +) (string, error) { + return CallOpenAICompatibleCtx( + context.Background(), + baseURL, + apiKey, + model, + messages, + ) +} + +func RunClaudeReview( + prompt, + model, + system, + tools string, + timeout time.Duration, +) (string, error) { + claudeBin := strings.TrimSpace(EnvOrDefault("CLAUDE_BIN", "claude")) + resolved, err := exec.LookPath(claudeBin) + if err != nil { + return "", fmt.Errorf("Claude CLI not found: %s", claudeBin) + } + + args := []string{ + "-p", + prompt, + "--output-format", + "json", + "--permission-mode", + "plan", + } + if model != "" { + args = append(args, "--model", model) + } + if system != "" { + args = append(args, "--system-prompt", system) + } + if tools != "" { + args = append(args, "--tools", tools) + } + + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + cmd := exec.CommandContext(ctx, resolved, args...) + cmd.Stdin = nil + var stdout bytes.Buffer + var stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + if errors.Is(ctx.Err(), context.DeadlineExceeded) { + return "", fmt.Errorf("Claude review timed out after %s", timeout) + } + message := strings.TrimSpace(stderr.String()) + if message == "" { + message = err.Error() + } + return "", fmt.Errorf("Claude review failed: %s", message) + } + + payload, err := ParseClaudeJSON(stdout.String()) + if err != nil { + message := strings.TrimSpace(stderr.String()) + if message != "" { + return "", fmt.Errorf("%v. stderr: %s", err, message) + } + return "", err + } + if isError, _ := payload["is_error"].(bool); isError { + return "", fmt.Errorf("%v", payload["result"]) + } + response := strings.TrimSpace(fmt.Sprint(payload["result"])) + if response == "" || response == "" { + return "", errors.New("Claude review returned empty output") + } + return response, nil +} + +func ParseClaudeJSON(raw string) (map[string]any, error) { + lines := strings.Split(raw, "\n") + for i := len(lines) - 1; i >= 0; i-- { + candidate := strings.TrimSpace(lines[i]) + if candidate == "" { + continue + } + var payload map[string]any + if err := json.Unmarshal([]byte(candidate), &payload); err == nil { + return payload, nil + } + } + return nil, errors.New("Claude CLI did not return JSON output") +} diff --git a/go/go_core/internal/toolbridge/runner.go b/go/go_core/internal/toolbridge/runner.go new file mode 100644 index 00000000..f7ccdac1 --- /dev/null +++ b/go/go_core/internal/toolbridge/runner.go @@ -0,0 +1,173 @@ +package toolbridge + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "io" + "strings" + + "xworkmate/go_core/internal/shared" +) + +func Run(input io.Reader, output io.Writer) { + reader := bufio.NewReader(input) + for { + payload, err := readMessage(reader) + if err != nil { + if errors.Is(err, io.EOF) { + return + } + writeError(output, nil, -32700, err.Error()) + continue + } + if len(strings.TrimSpace(string(payload))) == 0 { + continue + } + + request, err := shared.DecodeRPCRequest(payload) + if err != nil { + writeError(output, nil, -32700, err.Error()) + continue + } + + response := handleRequest(request) + if response != nil { + writeMessage(output, response) + } + } +} + +func readMessage(reader *bufio.Reader) ([]byte, error) { + line, err := reader.ReadString('\n') + if err != nil { + return nil, err + } + line = strings.TrimSpace(line) + if line == "" { + return nil, nil + } + if strings.HasPrefix(strings.ToLower(line), "content-length:") { + var contentLength int + if _, err := fmt.Sscanf(line, "Content-Length: %d", &contentLength); err != nil { + if _, err2 := fmt.Sscanf(line, "content-length: %d", &contentLength); err2 != nil { + return nil, fmt.Errorf("invalid content-length header") + } + } + for { + headerLine, err := reader.ReadString('\n') + if err != nil { + return nil, err + } + if strings.TrimSpace(headerLine) == "" { + break + } + } + body := make([]byte, contentLength) + if _, err := io.ReadFull(reader, body); err != nil { + return nil, err + } + return body, nil + } + return []byte(line), nil +} + +func writeMessage(output io.Writer, message map[string]any) { + payload, _ := json.Marshal(message) + _, _ = output.Write(append(payload, '\n')) +} + +func writeError(output io.Writer, id any, code int, message string) { + writeMessage(output, shared.ErrorEnvelope(id, code, message)) +} + +func handleRequest(request shared.RPCRequest) map[string]any { + if request.ID == nil { + return nil + } + + switch request.Method { + case "initialize": + return shared.ResultEnvelope(request.ID, map[string]any{ + "protocolVersion": "2024-11-05", + "capabilities": map[string]any{ + "tools": map[string]any{}, + }, + "serverInfo": map[string]any{ + "name": "xworkmate-go-core", + "version": "0.2.0", + }, + }) + case "ping": + return shared.ResultEnvelope(request.ID, map[string]any{}) + case "tools/list": + return shared.ResultEnvelope(request.ID, map[string]any{ + "tools": []map[string]any{ + { + "name": "chat", + "description": "OpenAI-compatible reviewer chat bridge", + "inputSchema": map[string]any{ + "type": "object", + "properties": map[string]any{ + "prompt": map[string]any{"type": "string"}, + "model": map[string]any{"type": "string"}, + "system": map[string]any{"type": "string"}, + }, + "required": []string{"prompt"}, + }, + }, + { + "name": "claude_review", + "description": "Review-only bridge over Claude CLI", + "inputSchema": map[string]any{ + "type": "object", + "properties": map[string]any{ + "prompt": map[string]any{"type": "string"}, + "model": map[string]any{"type": "string"}, + "system": map[string]any{"type": "string"}, + "tools": map[string]any{"type": "string"}, + }, + "required": []string{"prompt"}, + }, + }, + }, + }) + case "tools/call": + var params shared.ToolCallParams + raw, _ := json.Marshal(request.Params) + if err := json.Unmarshal(raw, ¶ms); err != nil { + return shared.ErrorResponse( + request.ID, + -32602, + fmt.Sprintf("invalid tool params: %v", err), + ) + } + switch params.Name { + case "chat": + content, err := shared.HandleChatTool(params.Arguments) + if err != nil { + return shared.ToolErrorResult(request.ID, err) + } + return shared.ToolTextResult(request.ID, content) + case "claude_review": + content, err := shared.HandleClaudeReviewTool(params.Arguments) + if err != nil { + return shared.ToolErrorResult(request.ID, err) + } + return shared.ToolTextResult(request.ID, content) + default: + return shared.ErrorResponse( + request.ID, + -32601, + fmt.Sprintf("unknown tool: %s", params.Name), + ) + } + default: + return shared.ErrorResponse( + request.ID, + -32601, + fmt.Sprintf("unknown method: %s", request.Method), + ) + } +} diff --git a/go/go_core/main.go b/go/go_core/main.go index 018df084..fdf4b210 100644 --- a/go/go_core/main.go +++ b/go/go_core/main.go @@ -1,784 +1,21 @@ package main import ( - "bufio" - "bytes" - "context" - "encoding/json" - "errors" - "flag" "fmt" - "io" - "net/http" "os" - "strings" - "sync" - "time" - "github.com/gorilla/websocket" + "xworkmate/go_core/internal/acp" + "xworkmate/go_core/internal/toolbridge" ) -type rpcRequest struct { - JSONRPC string `json:"jsonrpc,omitempty"` - ID any `json:"id,omitempty"` - Method string `json:"method,omitempty"` - Params map[string]any `json:"params,omitempty"` -} - -type rpcError struct { - Code int `json:"code"` - Message string `json:"message"` -} - -type toolCallParams struct { - Name string `json:"name"` - Arguments map[string]any `json:"arguments"` -} - -type acpSession struct { - sessionID string - threadID string - mode string - provider string - history []string - seq int - cancel context.CancelFunc - closed bool -} - -type acpTask struct { - req rpcRequest - notify func(map[string]any) - done chan acpTaskResult -} - -type acpTaskResult struct { - response map[string]any - err *rpcError -} - -type acpServer struct { - mu sync.Mutex - sessions map[string]*acpSession - queues map[string]chan acpTask -} - -var wsUpgrader = websocket.Upgrader{ - ReadBufferSize: 16 * 1024, - WriteBufferSize: 16 * 1024, - CheckOrigin: func(*http.Request) bool { - return true - }, -} - func main() { if len(os.Args) > 1 && os.Args[1] == "serve" { - serveACP() - return - } - runToolBridge() -} - -func serveACP() { - flags := flag.NewFlagSet("serve", flag.ExitOnError) - listen := flags.String( - "listen", - envOrDefault("ACP_LISTEN_ADDR", "127.0.0.1:8787"), - "ACP listen address", - ) - _ = flags.Parse(os.Args[2:]) - - server := newACPServer() - mux := http.NewServeMux() - mux.HandleFunc("/acp", server.handleWebSocket) - mux.HandleFunc("/acp/rpc", server.handleRPC) - - httpServer := &http.Server{ - Addr: strings.TrimSpace(*listen), - Handler: mux, - ReadTimeout: 30 * time.Second, - WriteTimeout: 5 * time.Minute, - IdleTimeout: 2 * time.Minute, - } - - if err := httpServer.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { - fmt.Fprintf(os.Stderr, "ACP server failed: %v\n", err) - os.Exit(1) - } -} - -func runToolBridge() { - reader := bufio.NewReader(os.Stdin) - for { - payload, err := readMessage(reader) - if err != nil { - if errors.Is(err, io.EOF) { - return - } - writeError(nil, -32700, err.Error()) - continue - } - if len(bytes.TrimSpace(payload)) == 0 { - continue - } - - var request rpcRequest - if err := json.Unmarshal(payload, &request); err != nil { - writeError(nil, -32700, fmt.Sprintf("invalid json: %v", err)) - continue - } - - response := handleToolBridgeRequest(request) - if response != nil { - writeMessage(response) - } - } -} - -func readMessage(reader *bufio.Reader) ([]byte, error) { - line, err := reader.ReadString('\n') - if err != nil { - return nil, err - } - line = strings.TrimSpace(line) - if line == "" { - return nil, nil - } - if strings.HasPrefix(strings.ToLower(line), "content-length:") { - var contentLength int - if _, err := fmt.Sscanf(line, "Content-Length: %d", &contentLength); err != nil { - if _, err2 := fmt.Sscanf(line, "content-length: %d", &contentLength); err2 != nil { - return nil, fmt.Errorf("invalid content-length header") - } - } - for { - headerLine, err := reader.ReadString('\n') - if err != nil { - return nil, err - } - if strings.TrimSpace(headerLine) == "" { - break - } - } - body := make([]byte, contentLength) - if _, err := io.ReadFull(reader, body); err != nil { - return nil, err - } - return body, nil - } - return []byte(line), nil -} - -func writeMessage(message map[string]any) { - payload, _ := json.Marshal(message) - _, _ = os.Stdout.Write(append(payload, '\n')) -} - -func writeError(id any, code int, message string) { - writeMessage(map[string]any{ - "jsonrpc": "2.0", - "id": id, - "error": map[string]any{ - "code": code, - "message": message, - }, - }) -} - -func handleToolBridgeRequest(request rpcRequest) map[string]any { - if request.ID == nil { - return nil - } - - switch request.Method { - case "initialize": - return map[string]any{ - "jsonrpc": "2.0", - "id": request.ID, - "result": map[string]any{ - "protocolVersion": "2024-11-05", - "capabilities": map[string]any{ - "tools": map[string]any{}, - }, - "serverInfo": map[string]any{ - "name": "xworkmate-go-core", - "version": "0.2.0", - }, - }, - } - case "ping": - return map[string]any{ - "jsonrpc": "2.0", - "id": request.ID, - "result": map[string]any{}, - } - case "tools/list": - return map[string]any{ - "jsonrpc": "2.0", - "id": request.ID, - "result": map[string]any{ - "tools": []map[string]any{ - { - "name": "chat", - "description": "OpenAI-compatible reviewer chat bridge", - "inputSchema": map[string]any{ - "type": "object", - "properties": map[string]any{ - "prompt": map[string]any{"type": "string"}, - "model": map[string]any{"type": "string"}, - "system": map[string]any{"type": "string"}, - }, - "required": []string{"prompt"}, - }, - }, - { - "name": "claude_review", - "description": "Review-only bridge over Claude CLI", - "inputSchema": map[string]any{ - "type": "object", - "properties": map[string]any{ - "prompt": map[string]any{"type": "string"}, - "model": map[string]any{"type": "string"}, - "system": map[string]any{"type": "string"}, - "tools": map[string]any{"type": "string"}, - }, - "required": []string{"prompt"}, - }, - }, - }, - }, - } - case "tools/call": - var params toolCallParams - raw, _ := json.Marshal(request.Params) - if err := json.Unmarshal(raw, ¶ms); err != nil { - return errorResponse(request.ID, -32602, fmt.Sprintf("invalid tool params: %v", err)) - } - switch params.Name { - case "chat": - content, err := handleChatTool(params.Arguments) - if err != nil { - return toolErrorResult(request.ID, err) - } - return toolTextResult(request.ID, content) - case "claude_review": - content, err := handleClaudeReviewTool(params.Arguments) - if err != nil { - return toolErrorResult(request.ID, err) - } - return toolTextResult(request.ID, content) - default: - return errorResponse(request.ID, -32601, fmt.Sprintf("unknown tool: %s", params.Name)) - } - default: - return errorResponse(request.ID, -32601, fmt.Sprintf("unknown method: %s", request.Method)) - } -} - -func newACPServer() *acpServer { - return &acpServer{ - sessions: make(map[string]*acpSession), - queues: make(map[string]chan acpTask), - } -} - -func (s *acpServer) handleWebSocket(w http.ResponseWriter, r *http.Request) { - conn, err := wsUpgrader.Upgrade(w, r, nil) - if err != nil { - return - } - defer conn.Close() - - var writeMu sync.Mutex - notify := func(message map[string]any) { - writeMu.Lock() - defer writeMu.Unlock() - _ = conn.WriteJSON(message) - } - - for { - _, payload, err := conn.ReadMessage() - if err != nil { - return - } - request, err := decodeRpcRequest(payload) - if err != nil { - notify(errorEnvelope(nil, -32700, err.Error())) - continue - } - response, rpcErr := s.handleACPRequest(request, notify) - if request.ID == nil { - continue - } - if rpcErr != nil { - notify(errorEnvelope(request.ID, rpcErr.Code, rpcErr.Message)) - continue - } - notify(resultEnvelope(request.ID, response)) - } -} - -func (s *acpServer) handleRPC(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - w.WriteHeader(http.StatusMethodNotAllowed) - return - } - payload, err := io.ReadAll(r.Body) - if err != nil { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte("invalid body")) - return - } - request, err := decodeRpcRequest(payload) - if err != nil { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(err.Error())) - return - } - - accept := strings.ToLower(r.Header.Get("Accept")) - stream := strings.Contains(accept, "text/event-stream") - if stream { - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") - } - - flusher, _ := w.(http.Flusher) - writeNotification := func(message map[string]any) { - if !stream { - return - } - writeSSE(w, message) - if flusher != nil { - flusher.Flush() - } - } - - response, rpcErr := s.handleACPRequest(request, writeNotification) - if request.ID == nil { - if stream { - _, _ = w.Write([]byte("data: [DONE]\n\n")) + if err := acp.Serve(os.Args[2:]); err != nil { + fmt.Fprintf(os.Stderr, "%v\n", err) + os.Exit(1) } return } - if rpcErr != nil { - envelope := errorEnvelope(request.ID, rpcErr.Code, rpcErr.Message) - if stream { - writeSSE(w, envelope) - if flusher != nil { - flusher.Flush() - } - return - } - w.WriteHeader(http.StatusOK) - _ = json.NewEncoder(w).Encode(envelope) - return - } - if stream { - writeSSE(w, resultEnvelope(request.ID, response)) - if flusher != nil { - flusher.Flush() - } - return - } - w.WriteHeader(http.StatusOK) - _ = json.NewEncoder(w).Encode(resultEnvelope(request.ID, response)) -} - -func (s *acpServer) handleACPRequest(request rpcRequest, notify func(map[string]any)) (map[string]any, *rpcError) { - method := strings.TrimSpace(request.Method) - switch method { - case "acp.capabilities": - providers := detectACPProviders() - singleAgent := len(providers) > 0 - multiAgent := boolArg(envOrDefault("ACP_MULTI_AGENT_ENABLED", "true"), true) - result := map[string]any{ - "singleAgent": singleAgent, - "multiAgent": multiAgent, - "providers": providers, - "capabilities": map[string]any{ - "single_agent": singleAgent, - "multi_agent": multiAgent, - "providers": providers, - }, - } - return result, nil - case "session.start", "session.message": - params := request.Params - sessionID := strings.TrimSpace(stringArg(params, "sessionId", "")) - if sessionID == "" { - return nil, &rpcError{Code: -32602, Message: "sessionId is required"} - } - threadID := strings.TrimSpace(stringArg(params, "threadId", sessionID)) - if threadID == "" { - threadID = sessionID - } - if method == "session.start" { - s.resetSession(sessionID, threadID) - } - result, rpcErr := s.enqueue(threadID, acpTask{ - req: request, - notify: notify, - done: make(chan acpTaskResult, 1), - }) - if rpcErr != nil { - return nil, rpcErr - } - return result, nil - case "session.cancel": - params := request.Params - sessionID := strings.TrimSpace(stringArg(params, "sessionId", "")) - if sessionID == "" { - return nil, &rpcError{Code: -32602, Message: "sessionId is required"} - } - cancelled := s.cancelSession(sessionID) - return map[string]any{"accepted": true, "cancelled": cancelled}, nil - case "session.close": - params := request.Params - sessionID := strings.TrimSpace(stringArg(params, "sessionId", "")) - if sessionID == "" { - return nil, &rpcError{Code: -32602, Message: "sessionId is required"} - } - closed := s.closeSession(sessionID) - return map[string]any{"accepted": true, "closed": closed}, nil - default: - return nil, &rpcError{Code: -32601, Message: fmt.Sprintf("unknown method: %s", method)} - } -} - -func (s *acpServer) enqueue(threadID string, task acpTask) (map[string]any, *rpcError) { - queue := s.ensureQueue(threadID) - queue <- task - result := <-task.done - return result.response, result.err -} - -func (s *acpServer) ensureQueue(threadID string) chan acpTask { - s.mu.Lock() - defer s.mu.Unlock() - queue, ok := s.queues[threadID] - if ok { - return queue - } - queue = make(chan acpTask, 32) - s.queues[threadID] = queue - go s.runQueue(queue) - return queue -} - -func (s *acpServer) runQueue(queue chan acpTask) { - for task := range queue { - response, err := s.executeSessionTask(task) - task.done <- acpTaskResult{response: response, err: err} - } -} - -func (s *acpServer) executeSessionTask(task acpTask) (map[string]any, *rpcError) { - params := task.req.Params - sessionID := strings.TrimSpace(stringArg(params, "sessionId", "")) - threadID := strings.TrimSpace(stringArg(params, "threadId", sessionID)) - mode := strings.TrimSpace(stringArg(params, "mode", "single-agent")) - provider := strings.TrimSpace(stringArg(params, "provider", "")) - if mode == "single-agent" && provider == "" { - provider = "codex" - } - - session := s.getOrCreateSession(sessionID, threadID) - session.mode = mode - if provider != "" { - session.provider = provider - } - - prompt := strings.TrimSpace(stringArg(params, "taskPrompt", "")) - if prompt != "" { - session.history = append(session.history, prompt) - } - turnID := fmt.Sprintf("turn-%d", time.Now().UnixNano()) - - ctx, cancel := context.WithCancel(context.Background()) - s.setSessionCancel(sessionID, cancel) - defer s.clearSessionCancel(sessionID) - - notify := task.notify - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "status", - "event": "started", - "message": "session started", - "pending": true, - "error": false, - }) - - if mode == "multi-agent" { - result := s.runMultiAgent(ctx, session, params, turnID, notify) - if result.err != nil { - return nil, result.err - } - return result.response, nil - } - - result := s.runSingleAgent(ctx, session, params, turnID, notify) - if result.err != nil { - return nil, result.err - } - return result.response, nil -} - -func (s *acpServer) runSingleAgent( - ctx context.Context, - session *acpSession, - params map[string]any, - turnID string, - notify func(map[string]any), -) acpTaskResult { - provider := session.provider - if provider == "" { - provider = strings.TrimSpace(stringArg(params, "provider", "codex")) - } - workingDirectory := strings.TrimSpace(stringArg(params, "workingDirectory", "")) - model := strings.TrimSpace(stringArg(params, "model", "")) - prompt := strings.TrimSpace(stringArg(params, "taskPrompt", "")) - prompt = augmentPromptWithAttachments(prompt, params) - - output, err := runProviderCommand(ctx, provider, model, prompt, workingDirectory) - if err != nil { - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "status", - "event": "completed", - "message": err.Error(), - "pending": false, - "error": true, - }) - return acpTaskResult{ - response: map[string]any{ - "success": false, - "error": err.Error(), - "turnId": turnID, - "mode": "single-agent", - "provider": provider, - }, - } - } - - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "delta", - "delta": output, - "pending": false, - "error": false, - }) - - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "status", - "event": "completed", - "message": "single-agent completed", - "pending": false, - "error": false, - }) - - return acpTaskResult{ - response: map[string]any{ - "success": true, - "output": output, - "turnId": turnID, - "mode": "single-agent", - "provider": provider, - }, - } -} - -func (s *acpServer) runMultiAgent( - ctx context.Context, - session *acpSession, - params map[string]any, - turnID string, - notify func(map[string]any), -) acpTaskResult { - prompt := composeHistoryPrompt(session.history) - if prompt == "" { - prompt = strings.TrimSpace(stringArg(params, "taskPrompt", "")) - } - prompt = augmentPromptWithAttachments(prompt, params) - - baseURL := normalizeBaseURL(stringArg(params, "aiGatewayBaseUrl", "")) - apiKey := strings.TrimSpace(stringArg(params, "aiGatewayApiKey", "")) - model := strings.TrimSpace(stringArg(params, "model", envOrDefault("ACP_MULTI_AGENT_MODEL", "gpt-4o"))) - if model == "" { - model = "gpt-4o" - } - - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "step", - "mode": "multi-agent", - "title": "Planner", - "message": "Preparing multi-agent run", - "pending": false, - "error": false, - "role": "architect", - "iteration": 1, - "score": 0, - }) - - if apiKey == "" { - errMsg := "aiGatewayApiKey is required for multi-agent mode" - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "status", - "mode": "multi-agent", - "message": errMsg, - "pending": false, - "error": true, - }) - return acpTaskResult{ - response: map[string]any{ - "success": false, - "error": errMsg, - "turnId": turnID, - "mode": "multi-agent", - }, - } - } - - messages := []map[string]string{ - {"role": "system", "content": "You are a multi-agent coordinator. Return concise actionable output."}, - {"role": "user", "content": prompt}, - } - output, err := callOpenAICompatibleCtx(ctx, baseURL, apiKey, model, messages) - if err != nil { - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "status", - "mode": "multi-agent", - "message": err.Error(), - "pending": false, - "error": true, - }) - return acpTaskResult{ - response: map[string]any{ - "success": false, - "error": err.Error(), - "turnId": turnID, - "mode": "multi-agent", - }, - } - } - - s.emitSessionUpdate(session, notify, turnID, map[string]any{ - "type": "step", - "mode": "multi-agent", - "title": "Reviewer", - "message": output, - "pending": false, - "error": false, - "role": "tester", - "iteration": 1, - "score": 9, - }) - - return acpTaskResult{ - response: map[string]any{ - "success": true, - "summary": output, - "finalScore": 9, - "iterations": 1, - "turnId": turnID, - "mode": "multi-agent", - }, - } -} - -func (s *acpServer) emitSessionUpdate( - session *acpSession, - notify func(map[string]any), - turnID string, - payload map[string]any, -) { - if notify == nil { - return - } - s.mu.Lock() - session.seq++ - seq := session.seq - s.mu.Unlock() - params := map[string]any{ - "sessionId": session.sessionID, - "threadId": session.threadID, - "turnId": turnID, - "seq": seq, - } - for key, value := range payload { - params[key] = value - } - notify(notificationEnvelope("session.update", params)) -} - -func (s *acpServer) getOrCreateSession(sessionID, threadID string) *acpSession { - s.mu.Lock() - defer s.mu.Unlock() - if session, ok := s.sessions[sessionID]; ok { - if threadID != "" { - session.threadID = threadID - } - session.closed = false - return session - } - session := &acpSession{sessionID: sessionID, threadID: threadID} - s.sessions[sessionID] = session - return session -} - -func (s *acpServer) resetSession(sessionID, threadID string) { - s.mu.Lock() - defer s.mu.Unlock() - s.sessions[sessionID] = &acpSession{ - sessionID: sessionID, - threadID: threadID, - history: []string{}, - } -} - -func (s *acpServer) setSessionCancel(sessionID string, cancel context.CancelFunc) { - s.mu.Lock() - defer s.mu.Unlock() - if session, ok := s.sessions[sessionID]; ok { - session.cancel = cancel - } -} - -func (s *acpServer) clearSessionCancel(sessionID string) { - s.mu.Lock() - defer s.mu.Unlock() - if session, ok := s.sessions[sessionID]; ok { - session.cancel = nil - } -} - -func (s *acpServer) cancelSession(sessionID string) bool { - s.mu.Lock() - session, ok := s.sessions[sessionID] - if !ok { - s.mu.Unlock() - return false - } - cancel := session.cancel - s.mu.Unlock() - if cancel != nil { - cancel() - return true - } - return false -} - -func (s *acpServer) closeSession(sessionID string) bool { - s.mu.Lock() - session, ok := s.sessions[sessionID] - if !ok { - s.mu.Unlock() - return false - } - cancel := session.cancel - session.closed = true - delete(s.sessions, sessionID) - s.mu.Unlock() - if cancel != nil { - cancel() - } - return true + + toolbridge.Run(os.Stdin, os.Stdout) } diff --git a/go/go_core/main_tools.go b/go/go_core/main_tools.go index 2222c9fe..c12a7d0c 100644 --- a/go/go_core/main_tools.go +++ b/go/go_core/main_tools.go @@ -1,486 +1,34 @@ package main import ( - "bytes" - "context" - "encoding/json" - "errors" - "fmt" - "io" - "net/http" - "os" - "os/exec" - "sort" - "strings" "time" + + "xworkmate/go_core/internal/shared" ) -func detectACPProviders() []string { - candidates := []struct { - provider string - envKey string - binary string - }{ - {provider: "codex", envKey: "ACP_CODEX_BIN", binary: "codex"}, - {provider: "opencode", envKey: "ACP_OPENCODE_BIN", binary: "opencode"}, - {provider: "claude", envKey: "ACP_CLAUDE_BIN", binary: "claude"}, - {provider: "gemini", envKey: "ACP_GEMINI_BIN", binary: "gemini"}, - } - providers := make([]string, 0, len(candidates)) - for _, candidate := range candidates { - binary := strings.TrimSpace(envOrDefault(candidate.envKey, candidate.binary)) - if binary == "" { - continue - } - if _, err := exec.LookPath(binary); err == nil { - providers = append(providers, candidate.provider) - } - } - sort.Strings(providers) - return providers +func parseClaudeJSON(raw string) (map[string]any, error) { + return shared.ParseClaudeJSON(raw) } -func runProviderCommand( - ctx context.Context, - provider, - model, - prompt, - workingDirectory string, -) (string, error) { - command, args := resolveProviderCommand(provider, model, prompt, workingDirectory) - if command == "" { - return "", fmt.Errorf("unsupported provider: %s", provider) - } - cmd := exec.CommandContext(ctx, command, args...) - if strings.TrimSpace(workingDirectory) != "" { - cmd.Dir = strings.TrimSpace(workingDirectory) - } - var stdout bytes.Buffer - var stderr bytes.Buffer - cmd.Stdout = &stdout - cmd.Stderr = &stderr - - if err := cmd.Run(); err != nil { - if errors.Is(ctx.Err(), context.Canceled) { - return "", errors.New("run canceled") - } - message := strings.TrimSpace(stderr.String()) - if message == "" { - message = err.Error() - } - return "", fmt.Errorf("%s run failed: %s", provider, message) - } - output := strings.TrimSpace(stdout.String()) - if output == "" { - output = strings.TrimSpace(stderr.String()) - } - if output == "" { - return "", fmt.Errorf("%s returned empty output", provider) - } - return output, nil -} - -func resolveProviderCommand(provider, model, prompt, cwd string) (string, []string) { - switch strings.TrimSpace(strings.ToLower(provider)) { - case "codex": - binary := strings.TrimSpace(envOrDefault("ACP_CODEX_BIN", "codex")) - args := []string{"exec", "--skip-git-repo-check", "--color", "never"} - if strings.TrimSpace(cwd) != "" { - args = append(args, "-C", strings.TrimSpace(cwd)) - } - if strings.TrimSpace(model) != "" { - args = append(args, "-m", strings.TrimSpace(model)) - } - args = append(args, prompt) - return binary, args - case "opencode": - binary := strings.TrimSpace(envOrDefault("ACP_OPENCODE_BIN", "opencode")) - args := []string{"run", "--format", "default"} - if strings.TrimSpace(cwd) != "" { - args = append(args, "--dir", strings.TrimSpace(cwd)) - } - if strings.TrimSpace(model) != "" { - args = append(args, "-m", strings.TrimSpace(model)) - } - args = append(args, prompt) - return binary, args - case "claude": - binary := strings.TrimSpace(envOrDefault("ACP_CLAUDE_BIN", "claude")) - if strings.TrimSpace(model) == "" { - return binary, []string{"-p", prompt} - } - return binary, []string{"--model", strings.TrimSpace(model), "-p", prompt} - case "gemini": - binary := strings.TrimSpace(envOrDefault("ACP_GEMINI_BIN", "gemini")) - if strings.TrimSpace(model) == "" { - return binary, []string{"-p", prompt} - } - return binary, []string{"--model", strings.TrimSpace(model), "-p", prompt} - default: - return "", nil - } -} - -func augmentPromptWithAttachments(prompt string, params map[string]any) string { - attachmentsRaw := listArg(params, "attachments") - if len(attachmentsRaw) == 0 { - return prompt - } - lines := make([]string, 0, len(attachmentsRaw)) - for _, raw := range attachmentsRaw { - entry, ok := raw.(map[string]any) - if !ok { - continue - } - name := strings.TrimSpace(stringArg(entry, "name", "attachment")) - path := strings.TrimSpace(stringArg(entry, "path", "")) - if path == "" { - continue - } - lines = append(lines, fmt.Sprintf("- %s: %s", name, path)) - } - if len(lines) == 0 { - return prompt - } - var builder strings.Builder - builder.WriteString("User-selected local attachments:\n") - builder.WriteString(strings.Join(lines, "\n")) - builder.WriteString("\n\n") - builder.WriteString(prompt) - return builder.String() -} - -func composeHistoryPrompt(history []string) string { - if len(history) == 0 { - return "" - } - var builder strings.Builder - for index, turn := range history { - builder.WriteString(fmt.Sprintf("## User Turn %d\n", index+1)) - builder.WriteString(turn) - builder.WriteString("\n\n") - } - return strings.TrimSpace(builder.String()) -} - -func callOpenAICompatibleCtx( - ctx context.Context, +func callOpenAICompatible( baseURL, apiKey, model string, messages []map[string]string, ) (string, error) { - payload := map[string]any{ - "model": model, - "messages": messages, - "max_tokens": 4096, - "stream": false, - } - body, _ := json.Marshal(payload) - request, err := http.NewRequestWithContext( - ctx, - http.MethodPost, - strings.TrimRight(baseURL, "/")+"/chat/completions", - bytes.NewReader(body), - ) - if err != nil { - return "", err - } - request.Header.Set("Content-Type", "application/json") - request.Header.Set("Authorization", "Bearer "+apiKey) - - client := &http.Client{Timeout: 120 * time.Second} - response, err := client.Do(request) - if err != nil { - return "", err - } - defer response.Body.Close() - responseBody, err := io.ReadAll(response.Body) - if err != nil { - return "", err - } - if response.StatusCode < 200 || response.StatusCode >= 300 { - return "", fmt.Errorf("api error %d: %s", response.StatusCode, strings.TrimSpace(string(responseBody))) - } - - var decoded map[string]any - if err := json.Unmarshal(responseBody, &decoded); err != nil { - return "", err - } - choices, _ := decoded["choices"].([]any) - if len(choices) == 0 { - return "", errors.New("missing choices in response") - } - choice, _ := choices[0].(map[string]any) - message, _ := choice["message"].(map[string]any) - content := strings.TrimSpace(fmt.Sprint(message["content"])) - if content == "" || content == "" { - return "", errors.New("empty response content") - } - return content, nil -} - -func decodeRpcRequest(payload []byte) (rpcRequest, error) { - var request rpcRequest - if err := json.Unmarshal(payload, &request); err != nil { - return rpcRequest{}, fmt.Errorf("invalid json: %w", err) - } - if strings.TrimSpace(request.Method) == "" { - return rpcRequest{}, errors.New("missing method") - } - if request.Params == nil { - request.Params = map[string]any{} - } - return request, nil -} - -func writeSSE(w http.ResponseWriter, payload map[string]any) { - encoded, _ := json.Marshal(payload) - _, _ = fmt.Fprintf(w, "data: %s\n\n", encoded) -} - -func resultEnvelope(id any, result map[string]any) map[string]any { - return map[string]any{ - "jsonrpc": "2.0", - "id": id, - "result": result, - } -} - -func errorEnvelope(id any, code int, message string) map[string]any { - return map[string]any{ - "jsonrpc": "2.0", - "id": id, - "error": map[string]any{ - "code": code, - "message": message, - }, - } -} - -func notificationEnvelope(method string, params map[string]any) map[string]any { - return map[string]any{ - "jsonrpc": "2.0", - "method": method, - "params": params, - } -} - -func errorResponse(id any, code int, message string) map[string]any { - return map[string]any{ - "jsonrpc": "2.0", - "id": id, - "error": map[string]any{ - "code": code, - "message": message, - }, - } -} - -func toolTextResult(id any, content string) map[string]any { - return map[string]any{ - "jsonrpc": "2.0", - "id": id, - "result": map[string]any{ - "content": []map[string]any{ - {"type": "text", "text": content}, - }, - }, - } -} - -func toolErrorResult(id any, err error) map[string]any { - return map[string]any{ - "jsonrpc": "2.0", - "id": id, - "result": map[string]any{ - "content": []map[string]any{ - {"type": "text", "text": fmt.Sprintf("Error: %v", err)}, - }, - "isError": true, - }, - } + return shared.CallOpenAICompatible(baseURL, apiKey, model, messages) } func handleChatTool(arguments map[string]any) (string, error) { - apiKey := strings.TrimSpace(envOrDefault("LLM_API_KEY", "")) - if apiKey == "" { - return "", errors.New("LLM_API_KEY environment variable not set") - } - baseURL := normalizeBaseURL(envOrDefault("LLM_BASE_URL", "https://api.openai.com/v1")) - model := stringArg(arguments, "model", envOrDefault("LLM_MODEL", "gpt-4o")) - prompt := strings.TrimSpace(stringArg(arguments, "prompt", "")) - if prompt == "" { - return "", errors.New("prompt is required") - } - system := strings.TrimSpace(stringArg(arguments, "system", "")) - - messages := make([]map[string]string, 0, 2) - if system != "" { - messages = append(messages, map[string]string{"role": "system", "content": system}) - } - messages = append(messages, map[string]string{"role": "user", "content": prompt}) - return callOpenAICompatible(baseURL, apiKey, model, messages) + return shared.HandleChatTool(arguments) } -func handleClaudeReviewTool(arguments map[string]any) (string, error) { - prompt := strings.TrimSpace(stringArg(arguments, "prompt", "")) - if prompt == "" { - return "", errors.New("prompt is required") - } - model := strings.TrimSpace(stringArg(arguments, "model", envOrDefault("CLAUDE_REVIEW_MODEL", ""))) - system := strings.TrimSpace(stringArg(arguments, "system", envOrDefault("CLAUDE_REVIEW_SYSTEM", ""))) - tools := strings.TrimSpace(stringArg(arguments, "tools", envOrDefault("CLAUDE_REVIEW_TOOLS", ""))) - timeout := intArg(envOrDefault("CLAUDE_REVIEW_TIMEOUT_SEC", "600"), 600) - return runClaudeReview(prompt, model, system, tools, time.Duration(timeout)*time.Second) -} - -func callOpenAICompatible(baseURL, apiKey, model string, messages []map[string]string) (string, error) { - return callOpenAICompatibleCtx(context.Background(), baseURL, apiKey, model, messages) -} - -func runClaudeReview(prompt, model, system, tools string, timeout time.Duration) (string, error) { - claudeBin := strings.TrimSpace(envOrDefault("CLAUDE_BIN", "claude")) - resolved, err := exec.LookPath(claudeBin) - if err != nil { - return "", fmt.Errorf("Claude CLI not found: %s", claudeBin) - } - - args := []string{"-p", prompt, "--output-format", "json", "--permission-mode", "plan"} - if model != "" { - args = append(args, "--model", model) - } - if system != "" { - args = append(args, "--system-prompt", system) - } - if tools != "" { - args = append(args, "--tools", tools) - } - - ctx, cancel := context.WithTimeout(context.Background(), timeout) - defer cancel() - - cmd := exec.CommandContext(ctx, resolved, args...) - cmd.Stdin = nil - var stdout bytes.Buffer - var stderr bytes.Buffer - cmd.Stdout = &stdout - cmd.Stderr = &stderr - - if err := cmd.Run(); err != nil { - if errors.Is(ctx.Err(), context.DeadlineExceeded) { - return "", fmt.Errorf("Claude review timed out after %s", timeout) - } - message := strings.TrimSpace(stderr.String()) - if message == "" { - message = err.Error() - } - return "", fmt.Errorf("Claude review failed: %s", message) - } - - payload, err := parseClaudeJSON(stdout.String()) - if err != nil { - message := strings.TrimSpace(stderr.String()) - if message != "" { - return "", fmt.Errorf("%v. stderr: %s", err, message) - } - return "", err - } - if isError, _ := payload["is_error"].(bool); isError { - return "", fmt.Errorf("%v", payload["result"]) - } - response := strings.TrimSpace(fmt.Sprint(payload["result"])) - if response == "" || response == "" { - return "", errors.New("Claude review returned empty output") - } - return response, nil -} - -func parseClaudeJSON(raw string) (map[string]any, error) { - lines := strings.Split(raw, "\n") - for i := len(lines) - 1; i >= 0; i-- { - candidate := strings.TrimSpace(lines[i]) - if candidate == "" { - continue - } - var payload map[string]any - if err := json.Unmarshal([]byte(candidate), &payload); err == nil { - return payload, nil - } - } - return nil, errors.New("Claude CLI did not return JSON output") -} - -func normalizeBaseURL(raw string) string { - trimmed := strings.TrimSpace(raw) - if trimmed == "" { - return "https://api.openai.com/v1" - } - if strings.HasSuffix(trimmed, "/v1") { - return trimmed - } - return strings.TrimRight(trimmed, "/") + "/v1" -} - -func envOrDefault(key, fallback string) string { - value := strings.TrimSpace(os.Getenv(key)) - if value == "" { - return fallback - } - return value -} - -func stringArg(arguments map[string]any, key, fallback string) string { - if arguments == nil { - return fallback - } - value, ok := arguments[key] - if !ok { - return fallback - } - text := strings.TrimSpace(fmt.Sprint(value)) - if text == "" || text == "" { - return fallback - } - return text -} - -func listArg(arguments map[string]any, key string) []any { - if arguments == nil { - return nil - } - raw, ok := arguments[key] - if !ok || raw == nil { - return nil - } - if values, ok := raw.([]any); ok { - return values - } - if values, ok := raw.([]interface{}); ok { - return values - } - return nil -} - -func intArg(raw string, fallback int) int { - var parsed int - if _, err := fmt.Sscanf(raw, "%d", &parsed); err != nil || parsed <= 0 { - return fallback - } - return parsed -} - -func boolArg(raw string, fallback bool) bool { - trimmed := strings.TrimSpace(strings.ToLower(raw)) - if trimmed == "" { - return fallback - } - switch trimmed { - case "1", "true", "yes", "on": - return true - case "0", "false", "no", "off": - return false - default: - return fallback - } +func runClaudeReview( + prompt, + model, + system, + tools string, + timeout time.Duration, +) (string, error) { + return shared.RunClaudeReview(prompt, model, system, tools, timeout) }