diff --git a/go/go_core/internal/acp/execution.go b/go/go_core/internal/acp/execution.go new file mode 100644 index 00000000..0f9ba939 --- /dev/null +++ b/go/go_core/internal/acp/execution.go @@ -0,0 +1,326 @@ +package acp + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/url" + "strings" + "time" + + "github.com/gorilla/websocket" + + "xworkmate/go_core/internal/router" + "xworkmate/go_core/internal/shared" +) + +func buildResolvedExecutionParams( + params map[string]any, + resolved router.Result, +) map[string]any { + next := make(map[string]any, len(params)+8) + for key, value := range params { + next[key] = value + } + switch resolved.ResolvedExecutionTarget { + case router.ExecutionTargetGateway: + next["mode"] = router.ExecutionTargetGatewayChat + next["executionTarget"] = resolved.ResolvedEndpointTarget + case router.ExecutionTargetMultiAgent: + next["mode"] = router.ExecutionTargetMultiAgent + default: + next["mode"] = router.ExecutionTargetSingleAgent + } + if strings.TrimSpace(resolved.ResolvedProviderID) != "" { + next["provider"] = strings.TrimSpace(resolved.ResolvedProviderID) + } + if strings.TrimSpace(resolved.ResolvedModel) != "" { + next["model"] = strings.TrimSpace(resolved.ResolvedModel) + } + if len(resolved.ResolvedSkills) > 0 { + next["selectedSkills"] = append([]string(nil), resolved.ResolvedSkills...) + } + next["resolvedExecutionTarget"] = resolved.ResolvedExecutionTarget + next["resolvedEndpointTarget"] = resolved.ResolvedEndpointTarget + next["resolvedProviderId"] = resolved.ResolvedProviderID + next["resolvedModel"] = resolved.ResolvedModel + next["resolvedSkills"] = append([]string(nil), resolved.ResolvedSkills...) + return next +} + +func (s *Server) runGateway( + ctx context.Context, + method string, + session *session, + params map[string]any, + turnID string, + notify func(map[string]any), +) taskResult { + _ = ctx + executionTarget := strings.TrimSpace(shared.StringArg(params, "executionTarget", "")) + if executionTarget == "" { + executionTarget = router.EndpointTargetLocal + } + result := s.gateway.RequestByMode( + executionTarget, + method, + params, + 2*time.Minute, + notify, + ) + if !result.OK { + errMessage := strings.TrimSpace(shared.StringArg(result.Error, "message", "gateway execution failed")) + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "event": "completed", + "message": errMessage, + "pending": false, + "error": true, + }) + return taskResult{ + response: map[string]any{ + "success": false, + "error": errMessage, + "turnId": turnID, + "mode": router.ExecutionTargetGatewayChat, + }, + } + } + payload := asMap(result.Payload) + if len(payload) == 0 { + payload = map[string]any{ + "success": true, + "turnId": turnID, + "mode": router.ExecutionTargetGatewayChat, + } + } + if _, ok := payload["turnId"]; !ok { + payload["turnId"] = turnID + } + if _, ok := payload["mode"]; !ok { + payload["mode"] = router.ExecutionTargetGatewayChat + } + return taskResult{response: payload} +} + +func (s *Server) runSingleAgentViaExternalProvider( + ctx context.Context, + provider syncedProvider, + method string, + params map[string]any, + notify func(map[string]any), +) (map[string]any, error) { + endpoint := strings.TrimSpace(provider.Endpoint) + if endpoint == "" { + return nil, fmt.Errorf("external provider endpoint is missing") + } + return requestExternalACP( + ctx, + endpoint, + provider.AuthorizationHeader, + method, + params, + notify, + ) +} + +func requestExternalACP( + ctx context.Context, + endpoint, + authorization, + method string, + params map[string]any, + notify func(map[string]any), +) (map[string]any, error) { + parsed, err := httpOrWebsocketEndpoint(endpoint) + if err != nil { + return nil, err + } + switch parsed.Scheme { + case "http", "https": + return requestExternalACPHTTP(ctx, parsed, authorization, method, params) + default: + return requestExternalACPWebSocket(ctx, parsed, authorization, method, params, notify) + } +} + +func requestExternalACPHTTP( + ctx context.Context, + endpoint *urlSpec, + authorization, + method string, + params map[string]any, +) (map[string]any, error) { + requestBody, _ := json.Marshal(map[string]any{ + "jsonrpc": "2.0", + "id": fmt.Sprintf("req-%d", time.Now().UnixNano()), + "method": method, + "params": params, + }) + req, err := http.NewRequestWithContext( + ctx, + http.MethodPost, + endpoint.httpRPCEndpoint(), + strings.NewReader(string(requestBody)), + ) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json; charset=utf-8") + req.Header.Set("Accept", "application/json") + if strings.TrimSpace(authorization) != "" { + req.Header.Set("Authorization", strings.TrimSpace(authorization)) + } + response, err := (&http.Client{Timeout: 2 * time.Minute}).Do(req) + if err != nil { + return nil, err + } + defer response.Body.Close() + var decoded map[string]any + if err := json.NewDecoder(response.Body).Decode(&decoded); err != nil { + return nil, err + } + if errPayload := asMap(decoded["error"]); len(errPayload) > 0 { + return nil, fmt.Errorf( + "%s", + strings.TrimSpace(shared.StringArg(errPayload, "message", "external ACP request failed")), + ) + } + return decoded, nil +} + +func requestExternalACPWebSocket( + ctx context.Context, + endpoint *urlSpec, + authorization, + method string, + params map[string]any, + notify func(map[string]any), +) (map[string]any, error) { + headers := http.Header{} + if strings.TrimSpace(authorization) != "" { + headers.Set("Authorization", strings.TrimSpace(authorization)) + } + conn, _, err := websocket.DefaultDialer.DialContext( + ctx, + endpoint.webSocketEndpoint(), + headers, + ) + if err != nil { + return nil, err + } + defer conn.Close() + + requestID := fmt.Sprintf("req-%d", time.Now().UnixNano()) + if err := conn.WriteJSON(map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "method": method, + "params": params, + }); err != nil { + return nil, err + } + + for { + if err := conn.SetReadDeadline(time.Now().Add(2 * time.Minute)); err != nil { + return nil, err + } + var payload map[string]any + if err := conn.ReadJSON(&payload); err != nil { + return nil, err + } + if strings.TrimSpace(shared.StringArg(payload, "id", "")) == requestID && + (payload["result"] != nil || payload["error"] != nil) { + if errPayload := asMap(payload["error"]); len(errPayload) > 0 { + return nil, fmt.Errorf( + "%s", + strings.TrimSpace(shared.StringArg(errPayload, "message", "external ACP request failed")), + ) + } + return payload, nil + } + if notify != nil && strings.TrimSpace(shared.StringArg(payload, "method", "")) != "" { + notify(payload) + } + } +} + +type urlSpec struct { + Scheme string + Host string + Port string + Path string +} + +func httpOrWebsocketEndpoint(raw string) (*urlSpec, error) { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return nil, fmt.Errorf("missing external ACP endpoint") + } + parsed, err := url.ParseRequestURI(trimmed) + if err != nil { + return nil, err + } + scheme := strings.ToLower(strings.TrimSpace(parsed.Scheme)) + if scheme != "http" && scheme != "https" && scheme != "ws" && scheme != "wss" { + return nil, fmt.Errorf("unsupported external ACP scheme: %s", scheme) + } + return &urlSpec{ + Scheme: scheme, + Host: parsed.Host, + Path: strings.TrimRight(parsed.Path, "/"), + }, nil +} + +func (u *urlSpec) basePath() string { + path := strings.TrimSpace(u.Path) + if path == "" || path == "/" { + return "" + } + if strings.HasSuffix(path, "/acp/rpc") { + path = strings.TrimSuffix(path, "/acp/rpc") + } else if strings.HasSuffix(path, "/acp") { + path = strings.TrimSuffix(path, "/acp") + } + path = strings.TrimRight(path, "/") + if path == "" || path == "/" { + return "" + } + if !strings.HasPrefix(path, "/") { + return "/" + path + } + return path +} + +func (u *urlSpec) httpRPCEndpoint() string { + scheme := u.Scheme + if scheme == "ws" { + scheme = "http" + } else if scheme == "wss" { + scheme = "https" + } + basePath := u.basePath() + if basePath == "" { + basePath = "/acp/rpc" + } else { + basePath += "/acp/rpc" + } + return fmt.Sprintf("%s://%s%s", scheme, u.Host, basePath) +} + +func (u *urlSpec) webSocketEndpoint() string { + scheme := u.Scheme + if scheme == "http" { + scheme = "ws" + } else if scheme == "https" { + scheme = "wss" + } + basePath := u.basePath() + if basePath == "" { + basePath = "/acp" + } else { + basePath += "/acp" + } + return fmt.Sprintf("%s://%s%s", scheme, u.Host, basePath) +} diff --git a/go/go_core/internal/acp/providers_sync.go b/go/go_core/internal/acp/providers_sync.go new file mode 100644 index 00000000..850df7e6 --- /dev/null +++ b/go/go_core/internal/acp/providers_sync.go @@ -0,0 +1,99 @@ +package acp + +import ( + "sort" + "strings" + + "xworkmate/go_core/internal/shared" +) + +type syncedProvider struct { + ProviderID string + Label string + Endpoint string + AuthorizationHeader string + Enabled bool +} + +func parseSyncedProviders(raw any) []syncedProvider { + list, ok := raw.([]any) + if !ok { + return nil + } + providers := make([]syncedProvider, 0, len(list)) + for _, item := range list { + entry := asMap(item) + providerID := strings.TrimSpace(sharedString(entry, "providerId")) + if providerID == "" { + continue + } + providers = append(providers, syncedProvider{ + ProviderID: providerID, + Label: strings.TrimSpace(sharedString(entry, "label")), + Endpoint: strings.TrimSpace(sharedString(entry, "endpoint")), + AuthorizationHeader: strings.TrimSpace(sharedString(entry, "authorizationHeader")), + Enabled: parseBool(entry["enabled"]), + }) + } + return providers +} + +func (s *Server) syncProviders(providers []syncedProvider) map[string]any { + s.mu.Lock() + defer s.mu.Unlock() + s.providerCatalog = make(map[string]syncedProvider, len(providers)) + for _, provider := range providers { + if strings.TrimSpace(provider.ProviderID) == "" { + continue + } + s.providerCatalog[provider.ProviderID] = provider + } + return map[string]any{ + "ok": true, + "providers": syncedProvidersResult(providers), + } +} + +func (s *Server) syncedProviderByID(providerID string) (syncedProvider, bool) { + s.mu.Lock() + defer s.mu.Unlock() + provider, ok := s.providerCatalog[strings.TrimSpace(providerID)] + if !ok || !provider.Enabled || strings.TrimSpace(provider.Endpoint) == "" { + return syncedProvider{}, false + } + return provider, true +} + +func (s *Server) availableProviders() []string { + providers := make(map[string]struct{}) + for _, provider := range shared.DetectACPProviders() { + providers[provider] = struct{}{} + } + s.mu.Lock() + for _, provider := range s.providerCatalog { + if !provider.Enabled || strings.TrimSpace(provider.Endpoint) == "" { + continue + } + providers[provider.ProviderID] = struct{}{} + } + s.mu.Unlock() + ordered := make([]string, 0, len(providers)) + for providerID := range providers { + ordered = append(ordered, providerID) + } + sort.Strings(ordered) + return ordered +} + +func syncedProvidersResult(providers []syncedProvider) []map[string]any { + result := make([]map[string]any, 0, len(providers)) + for _, provider := range providers { + result = append(result, map[string]any{ + "providerId": provider.ProviderID, + "label": provider.Label, + "endpoint": provider.Endpoint, + "enabled": provider.Enabled, + }) + } + return result +} diff --git a/go/go_core/internal/acp/providers_sync_test.go b/go/go_core/internal/acp/providers_sync_test.go new file mode 100644 index 00000000..2e6d648c --- /dev/null +++ b/go/go_core/internal/acp/providers_sync_test.go @@ -0,0 +1,127 @@ +package acp + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "xworkmate/go_core/internal/shared" +) + +func TestProvidersSyncUpdatesCapabilities(t *testing.T) { + server := NewServer() + + _, rpcErr := server.handleRequest(shared.RPCRequest{ + Method: "xworkmate.providers.sync", + Params: map[string]any{ + "providers": []any{ + map[string]any{ + "providerId": "claude", + "label": "Claude", + "endpoint": "http://127.0.0.1:9999", + "authorizationHeader": "Bearer test", + "enabled": true, + }, + }, + }, + }, func(map[string]any) {}) + if rpcErr != nil { + t.Fatalf("expected sync success, got %v", rpcErr) + } + + result, rpcErr := server.handleRequest(shared.RPCRequest{ + Method: "acp.capabilities", + Params: map[string]any{}, + }, func(map[string]any) {}) + if rpcErr != nil { + t.Fatalf("expected capabilities success, got %v", rpcErr) + } + providers, _ := result["providers"].([]string) + if len(providers) == 0 { + t.Fatalf("expected synced provider in capabilities, got %#v", result) + } + found := false + for _, provider := range providers { + if provider == "claude" { + found = true + break + } + } + if !found { + t.Fatalf("expected claude provider after sync, got %#v", providers) + } +} + +func TestExecuteSessionTaskUsesSyncedExternalProvider(t *testing.T) { + externalServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/acp/rpc" { + http.NotFound(w, r) + return + } + defer 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, _ := request["method"].(string) + switch method { + case "session.start": + _ = json.NewEncoder(w).Encode(map[string]any{ + "jsonrpc": "2.0", + "id": request["id"], + "result": map[string]any{ + "success": true, + "output": "external-provider-ok", + "turnId": "turn-external", + "provider": "claude", + "mode": "single-agent", + }, + }) + default: + _ = json.NewEncoder(w).Encode(map[string]any{ + "jsonrpc": "2.0", + "id": request["id"], + "result": map[string]any{"ok": true}, + }) + } + })) + defer externalServer.Close() + + server := NewServer() + server.syncProviders([]syncedProvider{ + { + ProviderID: "claude", + Label: "Claude", + Endpoint: externalServer.URL, + AuthorizationHeader: "Bearer test", + Enabled: true, + }, + }) + + response, rpcErr := server.executeSessionTask(task{ + req: shared.RPCRequest{ + Method: "session.start", + Params: map[string]any{ + "sessionId": "session-external", + "threadId": "thread-external", + "taskPrompt": "hello from external provider", + "workingDirectory": t.TempDir(), + "routing": map[string]any{ + "routingMode": "explicit", + "explicitExecutionTarget": "singleAgent", + "explicitProviderId": "claude", + }, + }, + }, + }) + if rpcErr != nil { + t.Fatalf("expected success, got rpc error: %v", rpcErr) + } + if got := response["output"]; got != "external-provider-ok" { + t.Fatalf("expected external provider output, got %#v", response) + } + if got := response["resolvedProviderId"]; got != "claude" { + t.Fatalf("expected resolved provider claude, got %#v", response) + } +} diff --git a/go/go_core/internal/acp/routing.go b/go/go_core/internal/acp/routing.go index ca7682ef..2811f779 100644 --- a/go/go_core/internal/acp/routing.go +++ b/go/go_core/internal/acp/routing.go @@ -11,15 +11,23 @@ import ( ) func handleRoutingResolve(params map[string]any) map[string]any { - result, _ := resolveRoutingMetadata(params) + result, _ := resolveRoutingMetadataWithProviders(params, nil) return mergeRoutingResponse(map[string]any{"ok": true}, result) } func resolveRoutingMetadata(params map[string]any) (router.Result, bool) { + return resolveRoutingMetadataWithProviders(params, nil) +} + +func resolveRoutingMetadataWithProviders( + params map[string]any, + availableProviders []string, +) (router.Result, bool) { routingParams := asMap(params["routing"]) if len(routingParams) == 0 { return router.Result{}, false } + installApproval := asMap(routingParams["installApproval"]) resolver := router.NewResolver() result := resolver.Resolve(router.Request{ @@ -32,9 +40,14 @@ func resolveRoutingMetadata(params map[string]any) (router.Result, bool) { ExplicitModel: strings.TrimSpace(sharedString(routingParams, "explicitModel")), ExplicitSkills: parseRoutingStringSlice(routingParams["explicitSkills"]), AllowSkillInstall: parseBool(routingParams["allowSkillInstall"]), - AvailableSkills: parseRoutingSkillCandidates(routingParams["availableSkills"]), - AIGatewayBaseURL: strings.TrimSpace(sharedString(params, "aiGatewayBaseUrl")), - AIGatewayAPIKey: strings.TrimSpace(sharedString(params, "aiGatewayApiKey")), + InstallApproval: skills.InstallApproval{ + RequestID: strings.TrimSpace(sharedString(installApproval, "requestId")), + ApprovedSkillKeys: parseRoutingStringSlice(installApproval["approvedSkillKeys"]), + }, + AvailableSkills: parseRoutingSkillCandidates(routingParams["availableSkills"]), + AvailableProviders: append([]string(nil), availableProviders...), + AIGatewayBaseURL: strings.TrimSpace(sharedString(params, "aiGatewayBaseUrl")), + AIGatewayAPIKey: strings.TrimSpace(sharedString(params, "aiGatewayApiKey")), }) return result, true } @@ -50,6 +63,16 @@ func mergeRoutingResponse(response map[string]any, result router.Result) map[str response["resolvedSkills"] = append([]string(nil), result.ResolvedSkills...) response["skillResolutionSource"] = result.SkillResolutionSource response["needsSkillInstall"] = result.NeedsSkillInstall + response["unavailable"] = result.Unavailable + if strings.TrimSpace(result.UnavailableCode) != "" { + response["unavailableCode"] = result.UnavailableCode + } + if strings.TrimSpace(result.UnavailableMessage) != "" { + response["unavailableMessage"] = result.UnavailableMessage + } + if strings.TrimSpace(result.SkillInstallRequestID) != "" { + response["skillInstallRequestId"] = result.SkillInstallRequestID + } if len(result.SkillCandidates) > 0 { response["skillCandidates"] = routingSkillCandidatesMap(result.SkillCandidates) } @@ -92,37 +115,6 @@ func recordRoutingSuccess( }) } -func applyResolvedRouting(params map[string]any, result router.Result) map[string]any { - if len(params) == 0 { - return params - } - next := make(map[string]any, len(params)+6) - for key, value := range params { - next[key] = value - } - switch result.ResolvedExecutionTarget { - case router.ExecutionTargetSingleAgent: - next["mode"] = router.ExecutionTargetSingleAgent - case router.ExecutionTargetMultiAgent: - next["mode"] = router.ExecutionTargetMultiAgent - case router.ExecutionTargetGateway: - next["mode"] = router.ExecutionTargetGatewayChat - if strings.TrimSpace(result.ResolvedEndpointTarget) != "" { - next["executionTarget"] = strings.TrimSpace(result.ResolvedEndpointTarget) - } - } - if strings.TrimSpace(result.ResolvedProviderID) != "" { - next["provider"] = strings.TrimSpace(result.ResolvedProviderID) - } - if strings.TrimSpace(result.ResolvedModel) != "" { - next["model"] = strings.TrimSpace(result.ResolvedModel) - } - if len(result.ResolvedSkills) > 0 { - next["selectedSkills"] = append([]string(nil), result.ResolvedSkills...) - } - return next -} - func parseRoutingSkillCandidates(raw any) []skills.Candidate { list, ok := raw.([]any) if !ok { diff --git a/go/go_core/internal/acp/routing_test.go b/go/go_core/internal/acp/routing_test.go index 6385731a..7e452dce 100644 --- a/go/go_core/internal/acp/routing_test.go +++ b/go/go_core/internal/acp/routing_test.go @@ -157,7 +157,6 @@ func TestExecuteSessionTaskAutoRoutingRecordsProjectMemory(t *testing.T) { Params: map[string]any{ "sessionId": "session-auto", "threadId": "thread-auto", - "mode": "single-agent", "provider": "claude", "taskPrompt": "create a powerpoint deck for launch", "workingDirectory": workspaceDir, @@ -183,25 +182,26 @@ func TestExecuteSessionTaskAutoRoutingRecordsProjectMemory(t *testing.T) { t.Fatalf("expected success response, got %#v", response) } + projectLocalMemory := filepath.Join(workspaceDir, ".xworkmate", "memory.md") + content, err := os.ReadFile(projectLocalMemory) + if err != nil { + t.Fatalf("expected memory file %s: %v", projectLocalMemory, err) + } + text := string(content) + if !strings.Contains(text, "preferred-route: single-agent") { + t.Fatalf("expected preferred route in %s, got %q", projectLocalMemory, text) + } + if !strings.Contains(text, "preferred-skills: PPTX") { + t.Fatalf("expected preferred skills in %s, got %q", projectLocalMemory, text) + } projectHomeMemory := filepath.Join( homeDir, "self-improving", "projects", filepath.Base(workspaceDir)+".md", ) - projectLocalMemory := filepath.Join(workspaceDir, ".xworkmate", "memory.md") - for _, target := range []string{projectHomeMemory, projectLocalMemory} { - content, err := os.ReadFile(target) - if err != nil { - t.Fatalf("expected memory file %s: %v", target, err) - } - text := string(content) - if !strings.Contains(text, "preferred-route: single-agent") { - t.Fatalf("expected preferred route in %s, got %q", target, text) - } - if !strings.Contains(text, "preferred-skills: PPTX") { - t.Fatalf("expected preferred skills in %s, got %q", target, text) - } + if _, err := os.Stat(projectHomeMemory); !os.IsNotExist(err) { + t.Fatalf("expected auto memory write to stay project-local only, got stat err=%v", err) } } @@ -230,7 +230,6 @@ func TestExecuteSessionTaskExplicitRoutingDoesNotRecordProjectMemory(t *testing. Params: map[string]any{ "sessionId": "session-explicit", "threadId": "thread-explicit", - "mode": "single-agent", "provider": "claude", "taskPrompt": "create a powerpoint deck for launch", "workingDirectory": workspaceDir, @@ -271,6 +270,27 @@ func TestExecuteSessionTaskExplicitRoutingDoesNotRecordProjectMemory(t *testing. } } +func TestExecuteSessionTaskRequiresRouting(t *testing.T) { + server := NewServer() + _, rpcErr := server.executeSessionTask(task{ + req: shared.RPCRequest{ + ID: "request-1", + Method: "session.start", + Params: map[string]any{ + "sessionId": "session-missing-routing", + "threadId": "thread-missing-routing", + "taskPrompt": "hello", + }, + }, + }) + if rpcErr == nil { + t.Fatalf("expected routing-required error") + } + if rpcErr.Message != "ROUTING_REQUIRED" { + t.Fatalf("expected ROUTING_REQUIRED, got %#v", rpcErr) + } +} + func TestExecuteSessionTaskAutoRoutingPromotesComplexRequestToMultiAgent(t *testing.T) { workspaceDir := filepath.Join(t.TempDir(), "workspace") if err := os.MkdirAll(workspaceDir, 0o755); err != nil { @@ -291,7 +311,6 @@ func TestExecuteSessionTaskAutoRoutingPromotesComplexRequestToMultiAgent(t *test Params: map[string]any{ "sessionId": "session-complex", "threadId": "thread-complex", - "mode": "single-agent", "provider": "claude", "taskPrompt": "collect latest news and summarize it into a report for review", "workingDirectory": workspaceDir, @@ -359,11 +378,40 @@ func TestHandleRoutingResolveAllowsSkillInstallRetry(t *testing.T) { if got := result["skillResolutionSource"]; got != "find_skills" { t.Fatalf("expected find_skills source, got %#v", got) } - if got := result["needsSkillInstall"]; got != false { + if got := result["needsSkillInstall"]; got != true { + t.Fatalf("expected first pass to request install approval, got %#v", got) + } + requestID, _ := result["skillInstallRequestId"].(string) + if strings.TrimSpace(requestID) == "" { + t.Fatalf("expected install request id, got %#v", result) + } + + retried := handleRoutingResolve(map[string]any{ + "taskPrompt": "translate and dub this video with subtitles", + "workingDirectory": "/tmp/workspace", + "routing": map[string]any{ + "routingMode": "auto", + "allowSkillInstall": true, + "installApproval": map[string]any{ + "requestId": requestID, + "approvedSkillKeys": []any{"video-translator"}, + }, + "availableSkills": []any{ + map[string]any{ + "id": "docx", + "label": "docx", + "description": "docs", + "installed": true, + }, + }, + }, + }) + + if got := retried["needsSkillInstall"]; got != false { t.Fatalf("expected install retry to clear needsSkillInstall, got %#v", got) } - resolvedSkills, _ := result["resolvedSkills"].([]string) + resolvedSkills, _ := retried["resolvedSkills"].([]string) if len(resolvedSkills) != 1 || resolvedSkills[0] != "video-translator" { - t.Fatalf("expected installed skill to resolve, got %#v", result["resolvedSkills"]) + t.Fatalf("expected installed skill to resolve, got %#v", retried["resolvedSkills"]) } } diff --git a/go/go_core/internal/acp/server.go b/go/go_core/internal/acp/server.go index 8c0c5e43..83b78d1f 100644 --- a/go/go_core/internal/acp/server.go +++ b/go/go_core/internal/acp/server.go @@ -48,6 +48,7 @@ type Server struct { sessions map[string]*session queues map[string]chan task gateway *gatewayruntime.Manager + providerCatalog map[string]syncedProvider } var wsUpgrader = websocket.Upgrader{ @@ -97,6 +98,7 @@ func NewServer() *Server { sessions: make(map[string]*session), queues: make(map[string]chan task), gateway: gatewayruntime.NewManager(), + providerCatalog: make(map[string]syncedProvider), } } @@ -247,7 +249,7 @@ func (s *Server) handleRequest( method := strings.TrimSpace(request.Method) switch method { case "acp.capabilities": - providers := shared.DetectACPProviders() + providers := s.availableProviders() singleAgent := len(providers) > 0 multiAgent := shared.BoolArg( shared.EnvOrDefault("ACP_MULTI_AGENT_ENABLED", "true"), @@ -316,7 +318,13 @@ func (s *Server) handleRequest( case "xworkmate.dispatch.resolve": return handleDispatchResolve(request.Params), nil case "xworkmate.routing.resolve": - return handleRoutingResolve(request.Params), nil + result, _ := resolveRoutingMetadataWithProviders( + request.Params, + s.availableProviders(), + ) + return mergeRoutingResponse(map[string]any{"ok": true}, result), nil + case "xworkmate.providers.sync": + return s.syncProviders(parseSyncedProviders(request.Params["providers"])), nil case "xworkmate.mounts.reconcile": return handleMountReconcile(request.Params), nil case "xworkmate.gateway.connect": @@ -533,18 +541,32 @@ func (s *Server) runQueue(queue chan task) { func (s *Server) executeSessionTask(task task) (map[string]any, *shared.RPCError) { params := task.req.Params - resolvedRouting, hasResolvedRouting := resolveRoutingMetadata(params) - if hasResolvedRouting { - params = applyResolvedRouting(params, resolvedRouting) + resolvedRouting, hasResolvedRouting := resolveRoutingMetadataWithProviders( + params, + s.availableProviders(), + ) + if !hasResolvedRouting { + return nil, &shared.RPCError{ + Code: -32602, + Message: "ROUTING_REQUIRED", + } } 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" + if resolvedRouting.Unavailable { + response := mergeRoutingResponse(map[string]any{ + "success": false, + "error": resolvedRouting.UnavailableMessage, + "unavailable": true, + "unavailableCode": resolvedRouting.UnavailableCode, + "unavailableMessage": resolvedRouting.UnavailableMessage, + }, resolvedRouting) + return response, nil } + executionParams := buildResolvedExecutionParams(params, resolvedRouting) + mode := strings.TrimSpace(shared.StringArg(executionParams, "mode", "single-agent")) + provider := strings.TrimSpace(shared.StringArg(executionParams, "provider", "")) session := s.getOrCreateSession(sessionID, threadID) session.mode = mode @@ -552,7 +574,7 @@ func (s *Server) executeSessionTask(task task) (map[string]any, *shared.RPCError session.provider = provider } - prompt := strings.TrimSpace(shared.StringArg(params, "taskPrompt", "")) + prompt := strings.TrimSpace(shared.StringArg(executionParams, "taskPrompt", "")) if prompt != "" { session.history = append(session.history, prompt) } @@ -572,49 +594,54 @@ func (s *Server) executeSessionTask(task task) (map[string]any, *shared.RPCError }) if mode == router.ExecutionTargetGatewayChat || mode == router.ExecutionTargetGateway { - result := taskResult{ - response: map[string]any{ - "success": false, - "error": "gateway execution must be dispatched to a connected gateway ACP endpoint", - "turnId": turnID, - "mode": router.ExecutionTargetGatewayChat, - }, - } - if hasResolvedRouting { - result.response = mergeRoutingResponse(result.response, resolvedRouting) + result := s.runGateway( + ctx, + task.req.Method, + session, + executionParams, + turnID, + notify, + ) + if result.err != nil { + return nil, result.err } + result.response = mergeRoutingResponse(result.response, resolvedRouting) return result.response, nil } if mode == "multi-agent" { - result := s.runMultiAgent(ctx, session, params, turnID, notify) + result := s.runMultiAgent(ctx, session, executionParams, turnID, notify) if result.err != nil { return nil, result.err } - if hasResolvedRouting { - result.response = mergeRoutingResponse(result.response, resolvedRouting) - if err := recordRoutingSuccess(params, resolvedRouting, result.response); err != nil { - return nil, &shared.RPCError{Code: -32001, Message: err.Error()} - } - } - return result.response, nil - } - - result := s.runSingleAgent(ctx, session, params, turnID, notify) - if result.err != nil { - return nil, result.err - } - if hasResolvedRouting { result.response = mergeRoutingResponse(result.response, resolvedRouting) if err := recordRoutingSuccess(params, resolvedRouting, result.response); err != nil { return nil, &shared.RPCError{Code: -32001, Message: err.Error()} } + return result.response, nil + } + + result := s.runSingleAgent( + ctx, + task.req.Method, + session, + executionParams, + turnID, + notify, + ) + if result.err != nil { + return nil, result.err + } + result.response = mergeRoutingResponse(result.response, resolvedRouting) + if err := recordRoutingSuccess(params, resolvedRouting, result.response); err != nil { + return nil, &shared.RPCError{Code: -32001, Message: err.Error()} } return result.response, nil } func (s *Server) runSingleAgent( ctx context.Context, + method string, session *session, params map[string]any, turnID string, @@ -631,6 +658,55 @@ func (s *Server) runSingleAgent( prompt := strings.TrimSpace(shared.StringArg(params, "taskPrompt", "")) prompt = shared.AugmentPromptWithAttachments(prompt, params) + if syncedProvider, ok := s.syncedProviderByID(provider); ok { + response, err := s.runSingleAgentViaExternalProvider( + ctx, + syncedProvider, + method, + params, + notify, + ) + if err == nil { + result := asMap(response["result"]) + if len(result) == 0 { + result = response + } + if _, exists := result["provider"]; !exists { + result["provider"] = provider + } + if _, exists := result["mode"]; !exists { + result["mode"] = "single-agent" + } + if _, exists := result["turnId"]; !exists { + result["turnId"] = turnID + } + return taskResult{response: result} + } + s.emitSessionUpdate(session, notify, turnID, map[string]any{ + "type": "status", + "event": "completed", + "message": err.Error(), + "pending": false, + "error": true, + }) + 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, + }, + } + } + output, err := shared.RunProviderCommand( ctx, provider, diff --git a/go/go_core/internal/gatewayruntime/runtime.go b/go/go_core/internal/gatewayruntime/runtime.go index 0fcf39b7..cfa7f5c4 100644 --- a/go/go_core/internal/gatewayruntime/runtime.go +++ b/go/go_core/internal/gatewayruntime/runtime.go @@ -163,6 +163,27 @@ func (m *Manager) Request( return current.request(method, params, timeout) } +func (m *Manager) RequestByMode( + mode string, + method string, + params map[string]any, + timeout time.Duration, + notify func(map[string]any), +) RequestResult { + current := m.lookupConnectedByMode(mode) + if current == nil { + return RequestResult{ + OK: false, + Error: (&GatewayError{ + Message: "gateway not connected", + Code: "OFFLINE", + }).Map(), + } + } + current.setNotify(notify) + return current.request(method, params, timeout) +} + func (m *Manager) Disconnect(runtimeID string, notify func(map[string]any)) { current := m.lookup(runtimeID) if current == nil { @@ -178,6 +199,25 @@ func (m *Manager) lookup(runtimeID string) *session { return m.sessions[strings.TrimSpace(runtimeID)] } +func (m *Manager) lookupConnectedByMode(mode string) *session { + normalizedMode := strings.TrimSpace(mode) + m.mu.Lock() + defer m.mu.Unlock() + for _, current := range m.sessions { + if current == nil { + continue + } + current.mu.Lock() + connected := current.snapshot.Status == "connected" + currentMode := current.snapshot.Mode + current.mu.Unlock() + if connected && strings.TrimSpace(currentMode) == normalizedMode { + return current + } + } + return nil +} + type session struct { manager *Manager runtimeID string diff --git a/go/go_core/internal/memory/provider.go b/go/go_core/internal/memory/provider.go index f5aeed17..a5dee4bb 100644 --- a/go/go_core/internal/memory/provider.go +++ b/go/go_core/internal/memory/provider.go @@ -17,6 +17,7 @@ type Preferences struct { PreferredRoute string PreferredModel string PreferredSkills []string + Provider string } type LoadResult struct { @@ -91,30 +92,39 @@ func (s Service) RecordSuccess(workingDirectory string, entry SuccessEntry) erro if projectName == "" { return nil } - targets := []string{ - filepath.Join(s.HomeDir, "self-improving", "projects", projectName+".md"), - filepath.Join(workingDirectory, ".xworkmate", "memory.md"), + target := s.projectWriteTarget(workingDirectory, projectName) + if target == "" { + return nil } block := formatSuccessEntry(entry) - for _, target := range targets { - if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { - return err - } - file, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644) - if err != nil { - return err - } - if _, err := file.WriteString(block); err != nil { - _ = file.Close() - return err - } - if err := file.Close(); err != nil { - return err - } + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return err + } + file, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644) + if err != nil { + return err + } + if _, err := file.WriteString(block); err != nil { + _ = file.Close() + return err + } + if err := file.Close(); err != nil { + return err } return nil } +func (s Service) projectWriteTarget( + workingDirectory string, + projectName string, +) string { + repoLocalDir := filepath.Join(workingDirectory, ".xworkmate") + if err := os.MkdirAll(repoLocalDir, 0o755); err == nil { + return filepath.Join(repoLocalDir, "memory.md") + } + return filepath.Join(s.HomeDir, "self-improving", "projects", projectName+".md") +} + func formatSuccessEntry(entry SuccessEntry) string { lines := []string{ "", @@ -162,6 +172,8 @@ func parsePreferences(text string) Preferences { prefs.PreferredSkills = append(prefs.PreferredSkills, value) } } + case strings.HasPrefix(strings.ToLower(trimmed), "provider:"): + prefs.Provider = strings.TrimSpace(strings.TrimPrefix(trimmed, "provider:")) } } return prefs @@ -177,6 +189,9 @@ func mergePreferences(dst *Preferences, src Preferences) { if len(src.PreferredSkills) > 0 { dst.PreferredSkills = append([]string(nil), src.PreferredSkills...) } + if strings.TrimSpace(src.Provider) != "" { + dst.Provider = strings.TrimSpace(src.Provider) + } } func sanitizeMemoryText(text string) string { @@ -192,7 +207,8 @@ func sanitizeMemoryText(text string) string { strings.Contains(normalized, "password") || strings.Contains(normalized, "secret") || strings.Contains(normalized, "api_key") || - strings.Contains(normalized, "apikey") { + strings.Contains(normalized, "apikey") || + strings.Contains(normalized, "api key") { continue } filtered = append(filtered, line) diff --git a/go/go_core/internal/memory/provider_test.go b/go/go_core/internal/memory/provider_test.go index d9500b91..ede54b55 100644 --- a/go/go_core/internal/memory/provider_test.go +++ b/go/go_core/internal/memory/provider_test.go @@ -65,22 +65,21 @@ func TestRecordSuccessWritesProjectLevelMemoryFiles(t *testing.T) { t.Fatalf("record success: %v", err) } - targets := []string{ - filepath.Join(homeDir, "self-improving", "projects", "repo.md"), - filepath.Join(workingDir, ".xworkmate", "memory.md"), + repoLocalTarget := filepath.Join(workingDir, ".xworkmate", "memory.md") + content, err := os.ReadFile(repoLocalTarget) + if err != nil { + t.Fatalf("read target %s: %v", repoLocalTarget, err) } - for _, target := range targets { - content, err := os.ReadFile(target) - if err != nil { - t.Fatalf("read target %s: %v", target, err) - } - text := string(content) - if !strings.Contains(text, "preferred-route: single-agent") { - t.Fatalf("missing preferred route in %s: %q", target, text) - } - if strings.Contains(strings.ToLower(text), "token") { - t.Fatalf("unexpected sensitive content in %s: %q", target, text) - } + text := string(content) + if !strings.Contains(text, "preferred-route: single-agent") { + t.Fatalf("missing preferred route in %s: %q", repoLocalTarget, text) + } + if strings.Contains(strings.ToLower(text), "token") { + t.Fatalf("unexpected sensitive content in %s: %q", repoLocalTarget, text) + } + homeProjectTarget := filepath.Join(homeDir, "self-improving", "projects", "repo.md") + if _, err := os.Stat(homeProjectTarget); !os.IsNotExist(err) { + t.Fatalf("expected single project-level write target, got stat err=%v", err) } } diff --git a/go/go_core/internal/router/router.go b/go/go_core/internal/router/router.go index 49d253b9..9311003d 100644 --- a/go/go_core/internal/router/router.go +++ b/go/go_core/internal/router/router.go @@ -2,6 +2,7 @@ package router import ( "os" + "sort" "strings" "xworkmate/go_core/internal/memory" @@ -32,7 +33,9 @@ type Request struct { ExplicitModel string ExplicitSkills []string AllowSkillInstall bool + InstallApproval skills.InstallApproval AvailableSkills []skills.Candidate + AvailableProviders []string AIGatewayBaseURL string AIGatewayAPIKey string } @@ -46,7 +49,11 @@ type Result struct { SkillResolutionSource string SkillCandidates []skills.Candidate NeedsSkillInstall bool + SkillInstallRequestID string MemorySources []memory.Source + Unavailable bool + UnavailableCode string + UnavailableMessage string } type Resolver struct { @@ -68,14 +75,20 @@ func NewResolver() Resolver { func (r Resolver) Resolve(req Request) Result { mem := r.MemoryService.Load(req.WorkingDirectory) + availableProviders := normalizeProviders(req.AvailableProviders) result := Result{ - ResolvedProviderID: strings.TrimSpace(req.ExplicitProviderID), ResolvedModel: strings.TrimSpace(req.ExplicitModel), MemorySources: mem.Sources, } result.ResolvedExecutionTarget, result.ResolvedEndpointTarget = r.resolveExecution(req, mem.Preferences) + result.ResolvedProviderID, result.Unavailable, result.UnavailableCode, result.UnavailableMessage = resolveProvider( + req, + mem.Preferences, + availableProviders, + result.ResolvedExecutionTarget, + ) if result.ResolvedModel == "" { result.ResolvedModel = strings.TrimSpace(mem.Preferences.PreferredModel) } @@ -85,12 +98,14 @@ func (r Resolver) Resolve(req Request) Result { ExplicitSkills: req.ExplicitSkills, AvailableSkills: req.AvailableSkills, AllowSkillInstall: req.AllowSkillInstall, + InstallApproval: req.InstallApproval, } skillResult := skills.Resolve(skillRequest, r.SkillFinder, r.SkillInstaller) result.ResolvedSkills = skillResult.ResolvedSkills result.SkillResolutionSource = skillResult.Source result.SkillCandidates = skillResult.Candidates result.NeedsSkillInstall = skillResult.NeedsInstall + result.SkillInstallRequestID = skillResult.InstallRequestID if len(result.ResolvedSkills) == 0 && len(mem.Preferences.PreferredSkills) > 0 { result.ResolvedSkills = append([]string(nil), mem.Preferences.PreferredSkills...) @@ -102,10 +117,18 @@ func (r Resolver) Resolve(req Request) Result { result.SkillResolutionSource = "none" } if result.ResolvedExecutionTarget == "" { - result.ResolvedExecutionTarget = ExecutionTargetSingleAgent + if len(availableProviders) > 0 { + result.ResolvedExecutionTarget = ExecutionTargetSingleAgent + } else { + result.ResolvedExecutionTarget = ExecutionTargetGateway + } } if result.ResolvedEndpointTarget == "" { - result.ResolvedEndpointTarget = EndpointTargetSingleAgent + if result.ResolvedExecutionTarget == ExecutionTargetGateway { + result.ResolvedEndpointTarget = normalizeGatewayTarget(req.PreferredGatewayTarget) + } else { + result.ResolvedEndpointTarget = EndpointTargetSingleAgent + } } return result } @@ -149,8 +172,15 @@ func (r Resolver) resolveExecution(req Request, prefs memory.Preferences) (strin return ExecutionTargetGateway, normalizeGatewayTarget(req.PreferredGatewayTarget) case ExecutionTargetMultiAgent: return ExecutionTargetMultiAgent, EndpointTargetSingleAgent + case ExecutionTargetSingleAgent: + if len(normalizeProviders(req.AvailableProviders)) > 0 { + return ExecutionTargetSingleAgent, EndpointTargetSingleAgent + } } - return ExecutionTargetSingleAgent, EndpointTargetSingleAgent + if len(normalizeProviders(req.AvailableProviders)) > 0 { + return ExecutionTargetSingleAgent, EndpointTargetSingleAgent + } + return ExecutionTargetGateway, normalizeGatewayTarget(req.PreferredGatewayTarget) } func (r Resolver) classify(req Request) string { @@ -181,13 +211,82 @@ func mapExplicitTarget(value string) (string, string) { func normalizeGatewayTarget(value string) string { switch strings.TrimSpace(value) { - case EndpointTargetLocal: + case EndpointTargetLocal, "": return EndpointTargetLocal default: return EndpointTargetRemote } } +func resolveProvider( + req Request, + prefs memory.Preferences, + availableProviders []string, + executionTarget string, +) (string, bool, string, string) { + explicitProviderID := normalize(strings.TrimSpace(req.ExplicitProviderID)) + if explicitProviderID != "" { + if len(availableProviders) == 0 { + return explicitProviderID, false, "", "" + } + if containsProvider(availableProviders, explicitProviderID) { + return explicitProviderID, false, "", "" + } + return "", true, "PROVIDER_UNAVAILABLE", "explicit provider is unavailable" + } + + if executionTarget != ExecutionTargetSingleAgent { + preferredProvider := normalize(strings.TrimSpace(prefs.Provider)) + if containsProvider(availableProviders, preferredProvider) { + return preferredProvider, false, "", "" + } + return "", false, "", "" + } + + preferredProvider := normalize(strings.TrimSpace(prefs.Provider)) + if containsProvider(availableProviders, preferredProvider) { + return preferredProvider, false, "", "" + } + if len(availableProviders) > 0 { + return availableProviders[0], false, "", "" + } + return "", true, "PROVIDER_UNAVAILABLE", "no single-agent provider is available" +} + +func normalizeProviders(values []string) []string { + if len(values) == 0 { + return nil + } + unique := make(map[string]struct{}, len(values)) + normalized := make([]string, 0, len(values)) + for _, value := range values { + providerID := normalize(value) + if providerID == "" { + continue + } + if _, ok := unique[providerID]; ok { + continue + } + unique[providerID] = struct{}{} + normalized = append(normalized, providerID) + } + sort.Strings(normalized) + return normalized +} + +func containsProvider(values []string, want string) bool { + want = normalize(want) + if want == "" { + return false + } + for _, value := range values { + if normalize(value) == want { + return true + } + } + return false +} + func looksLocal(prompt string) bool { return containsAny(prompt, []string{ "ppt", "pptx", "powerpoint", "word", "docx", "excel", "xlsx", "pdf", diff --git a/go/go_core/internal/skills/resolver.go b/go/go_core/internal/skills/resolver.go index 0c3255ba..34ebdece 100644 --- a/go/go_core/internal/skills/resolver.go +++ b/go/go_core/internal/skills/resolver.go @@ -2,6 +2,7 @@ package skills import ( "fmt" + "sort" "strings" ) @@ -25,6 +26,12 @@ type ResolveRequest struct { ExplicitSkills []string AvailableSkills []Candidate AllowSkillInstall bool + InstallApproval InstallApproval +} + +type InstallApproval struct { + RequestID string + ApprovedSkillKeys []string } type ResolveResult struct { @@ -32,6 +39,7 @@ type ResolveResult struct { Candidates []Candidate Source string NeedsInstall bool + InstallRequestID string } type StaticFinder struct{} @@ -97,8 +105,23 @@ func Resolve(req ResolveRequest, finder Finder, installer Installer) ResolveResu } } - if req.AllowSkillInstall && installer != nil && len(uninstalled) > 0 { - installedCandidates, err := installer.Install(uninstalled) + installRequestID := buildInstallRequestID(uninstalled) + if shouldInstallApprovedCandidates(req, installRequestID) && + installer != nil && + len(uninstalled) > 0 { + approvedCandidates := filterApprovedCandidates( + uninstalled, + req.InstallApproval.ApprovedSkillKeys, + ) + if len(approvedCandidates) == 0 { + return ResolveResult{ + Candidates: fallback, + Source: "find_skills", + NeedsInstall: true, + InstallRequestID: installRequestID, + } + } + installedCandidates, err := installer.Install(approvedCandidates) if err == nil && len(installedCandidates) > 0 { mergedAvailable := dedupeCandidates( append(append([]Candidate(nil), available...), installedCandidates...), @@ -116,12 +139,71 @@ func Resolve(req ResolveRequest, finder Finder, installer Installer) ResolveResu } return ResolveResult{ - Candidates: fallback, - Source: "find_skills", - NeedsInstall: len(uninstalled) > 0, + Candidates: fallback, + Source: "find_skills", + NeedsInstall: len(uninstalled) > 0, + InstallRequestID: installRequestID, } } +func shouldInstallApprovedCandidates( + req ResolveRequest, + expectedRequestID string, +) bool { + if !req.AllowSkillInstall || expectedRequestID == "" { + return false + } + if strings.TrimSpace(req.InstallApproval.RequestID) != expectedRequestID { + return false + } + return len(dedupeStrings(req.InstallApproval.ApprovedSkillKeys)) > 0 +} + +func filterApprovedCandidates( + candidates []Candidate, + approvedSkillKeys []string, +) []Candidate { + if len(candidates) == 0 { + return nil + } + approved := make(map[string]struct{}, len(approvedSkillKeys)) + for _, key := range approvedSkillKeys { + normalized := normalize(key) + if normalized == "" { + continue + } + approved[normalized] = struct{}{} + } + filtered := make([]Candidate, 0, len(candidates)) + for _, candidate := range candidates { + if _, ok := approved[normalize(candidate.ID)]; ok { + filtered = append(filtered, candidate) + } + } + return filtered +} + +func buildInstallRequestID(candidates []Candidate) string { + if len(candidates) == 0 { + return "" + } + keys := make([]string, 0, len(candidates)) + for _, candidate := range candidates { + key := normalize(candidate.ID) + if key == "" { + key = normalize(candidate.Label) + } + if key != "" { + keys = append(keys, key) + } + } + if len(keys) == 0 { + return "" + } + sort.Strings(keys) + return "skill-install:" + strings.Join(keys, ",") +} + type builtinSkill struct { id string label string diff --git a/go/go_core/internal/skills/resolver_test.go b/go/go_core/internal/skills/resolver_test.go index 42814794..7b658488 100644 --- a/go/go_core/internal/skills/resolver_test.go +++ b/go/go_core/internal/skills/resolver_test.go @@ -72,11 +72,38 @@ func TestResolveFallsBackToFindSkillsCandidates(t *testing.T) { } func TestResolveInstallsMissingSkillsWhenAuthorized(t *testing.T) { + initial := Resolve( + ResolveRequest{ + Prompt: "translate and dub this video with subtitles", + AvailableSkills: []Candidate{{ID: "docx", Label: "docx", Installed: true}}, + AllowSkillInstall: true, + }, + fakeFinder{ + {ID: "video-translator", Label: "video-translator", Installed: false}, + }, + fakeInstaller{ + installed: []Candidate{ + {ID: "video-translator", Label: "video-translator", Installed: true}, + }, + }, + ) + + if !initial.NeedsInstall { + t.Fatalf("expected install approval flow to pause first, got %#v", initial) + } + if initial.InstallRequestID == "" { + t.Fatalf("expected install request id, got %#v", initial) + } + result := Resolve( ResolveRequest{ Prompt: "translate and dub this video with subtitles", AvailableSkills: []Candidate{{ID: "docx", Label: "docx", Installed: true}}, AllowSkillInstall: true, + InstallApproval: InstallApproval{ + RequestID: initial.InstallRequestID, + ApprovedSkillKeys: []string{"video-translator"}, + }, }, fakeFinder{ {ID: "video-translator", Label: "video-translator", Installed: false}, diff --git a/lib/app/app_controller_desktop_core.dart b/lib/app/app_controller_desktop_core.dart index fbe2f612..bad7edf3 100644 --- a/lib/app/app_controller_desktop_core.dart +++ b/lib/app/app_controller_desktop_core.dart @@ -315,6 +315,11 @@ class AppController extends ChangeNotifier { {}; final Map singleAgentRuntimeModelBySessionInternal = {}; + final Map> + latestRoutingResolutionBySessionInternal = + >{}; + final Map syncedGoAgentProvidersInternal = + {}; final DesktopThreadArtifactService threadArtifactServiceInternal = DesktopThreadArtifactService(); List singleAgentSharedImportedSkillsInternal = diff --git a/lib/app/app_controller_desktop_go_agent_core_routing.dart b/lib/app/app_controller_desktop_go_agent_core_routing.dart new file mode 100644 index 00000000..e7db8be5 --- /dev/null +++ b/lib/app/app_controller_desktop_go_agent_core_routing.dart @@ -0,0 +1,101 @@ +// ignore_for_file: unused_import, unnecessary_import + +import 'dart:async'; +import 'dart:convert'; +import 'dart:io'; +import 'package:flutter/material.dart'; +import 'app_metadata.dart'; +import 'app_capabilities.dart'; +import 'app_store_policy.dart'; +import 'ui_feature_manifest.dart'; +import '../i18n/app_language.dart'; +import '../models/app_models.dart'; +import '../runtime/device_identity_store.dart'; +import '../runtime/aris_bundle.dart'; +import '../runtime/go_core.dart'; +import '../runtime/runtime_bootstrap.dart'; +import '../runtime/desktop_platform_service.dart'; +import '../runtime/gateway_runtime.dart'; +import '../runtime/runtime_controllers.dart'; +import '../runtime/runtime_models.dart'; +import '../runtime/secure_config_store.dart'; +import '../runtime/embedded_agent_launch_policy.dart'; +import '../runtime/runtime_coordinator.dart'; +import '../runtime/direct_single_agent_app_server_client.dart'; +import '../runtime/gateway_acp_client.dart'; +import '../runtime/codex_runtime.dart'; +import '../runtime/codex_config_bridge.dart'; +import '../runtime/code_agent_node_orchestrator.dart'; +import '../runtime/assistant_artifacts.dart'; +import '../runtime/desktop_thread_artifact_service.dart'; +import '../runtime/go_agent_core_client.dart'; +import '../runtime/mode_switcher.dart'; +import '../runtime/agent_registry.dart'; +import '../runtime/multi_agent_orchestrator.dart'; +import '../runtime/platform_environment.dart'; +import '../runtime/single_agent_runner.dart'; +import '../runtime/skill_directory_access.dart'; +import 'app_controller_desktop_core.dart'; +import 'app_controller_desktop_thread_sessions.dart'; + +extension AppControllerDesktopGoAgentCoreRouting on AppController { + Future> + buildGoAgentCoreSyncedProvidersInternal() async { + final providers = []; + for (final profile in settings.externalAcpEndpoints) { + final providerId = profile.providerKey.trim(); + final endpoint = profile.endpoint.trim(); + if (providerId.isEmpty || endpoint.isEmpty) { + continue; + } + final authorizationHeader = profile.authRef.trim().isEmpty + ? '' + : await settingsControllerInternal.resolveSecretValueInternal( + refName: profile.authRef.trim(), + ); + providers.add( + GoAgentCoreSyncedProvider( + providerId: providerId, + label: profile.label, + endpoint: endpoint, + authorizationHeader: authorizationHeader, + enabled: profile.enabled, + ), + ); + } + return providers; + } + + Future syncGoAgentCoreProvidersInternal() async { + final providers = await buildGoAgentCoreSyncedProvidersInternal(); + syncedGoAgentProvidersInternal + ..clear() + ..addEntries( + providers.map((item) => MapEntry(item.providerId.trim(), item)), + ); + await goAgentCoreClientInternal.syncProviders(providers); + } + + void updateLatestRoutingResolutionInternal( + String sessionKey, + GoAgentCoreRunResult result, + ) { + final normalizedSessionKey = normalizedAssistantSessionKeyInternal( + sessionKey, + ); + latestRoutingResolutionBySessionInternal[normalizedSessionKey] = + { + 'resolvedExecutionTarget': result.resolvedExecutionTarget, + 'resolvedEndpointTarget': result.resolvedEndpointTarget, + 'resolvedProviderId': result.resolvedProviderId, + 'resolvedModel': result.resolvedModel.trim(), + 'resolvedSkills': result.resolvedSkills, + 'skillResolutionSource': result.skillResolutionSource, + 'skillCandidates': result.skillCandidates, + 'needsSkillInstall': result.needsSkillInstall, + 'skillInstallRequestId': result.skillInstallRequestId, + 'memorySources': result.memorySources, + 'updatedAtMs': DateTime.now().millisecondsSinceEpoch, + }; + } +} diff --git a/lib/app/app_controller_desktop_runtime_coordination_impl.dart b/lib/app/app_controller_desktop_runtime_coordination_impl.dart index 2b125dae..ccbbfb55 100644 --- a/lib/app/app_controller_desktop_runtime_coordination_impl.dart +++ b/lib/app/app_controller_desktop_runtime_coordination_impl.dart @@ -45,6 +45,8 @@ import 'app_controller_desktop_workspace_execution.dart'; import 'app_controller_desktop_settings_runtime.dart'; import 'app_controller_desktop_thread_storage.dart'; import 'app_controller_desktop_skill_permissions.dart'; +import 'app_controller_desktop_go_agent_core_routing.dart'; +import 'app_controller_desktop_runtime_helpers.dart'; Future refreshAcpCapabilitiesRuntimeInternal( AppController controller, { @@ -82,6 +84,7 @@ Future refreshSingleAgentCapabilitiesRuntimeInternal( AppController controller, { bool forceRefresh = false, }) async { + await controller.syncGoAgentCoreProvidersInternal(); final capabilities = await controller.goAgentCoreClientInternal .loadCapabilities( target: AssistantExecutionTarget.singleAgent, diff --git a/lib/app/app_controller_desktop_settings.dart b/lib/app/app_controller_desktop_settings.dart index 57d4a8e9..6984e420 100644 --- a/lib/app/app_controller_desktop_settings.dart +++ b/lib/app/app_controller_desktop_settings.dart @@ -270,6 +270,7 @@ extension AppControllerDesktopSettings on AppController { aiGatewayStreamingClientsInternal.clear(); aiGatewayPendingSessionKeysInternal.clear(); aiGatewayAbortedSessionKeysInternal.clear(); + latestRoutingResolutionBySessionInternal.clear(); singleAgentExternalCliPendingSessionKeysInternal.clear(); assistantThreadTurnQueuesInternal.clear(); multiAgentRunPendingInternal = false; diff --git a/lib/app/app_controller_desktop_single_agent.dart b/lib/app/app_controller_desktop_single_agent.dart index eae643fd..d863f934 100644 --- a/lib/app/app_controller_desktop_single_agent.dart +++ b/lib/app/app_controller_desktop_single_agent.dart @@ -45,6 +45,7 @@ import 'app_controller_desktop_workspace_execution.dart'; import 'app_controller_desktop_settings_runtime.dart'; import 'app_controller_desktop_thread_storage.dart'; import 'app_controller_desktop_skill_permissions.dart'; +import 'app_controller_desktop_go_agent_core_routing.dart'; import 'app_controller_desktop_runtime_helpers.dart'; extension AppControllerDesktopSingleAgent on AppController { @@ -203,6 +204,7 @@ extension AppControllerDesktopSingleAgent on AppController { }, ); final resolvedRuntimeModel = result.resolvedModel.trim(); + updateLatestRoutingResolutionInternal(sessionKey, result); if (resolvedRuntimeModel.isNotEmpty) { singleAgentRuntimeModelBySessionInternal[sessionKey] = resolvedRuntimeModel; diff --git a/lib/app/app_controller_desktop_thread_sessions.dart b/lib/app/app_controller_desktop_thread_sessions.dart index bff57027..a16292ae 100644 --- a/lib/app/app_controller_desktop_thread_sessions.dart +++ b/lib/app/app_controller_desktop_thread_sessions.dart @@ -50,6 +50,14 @@ import 'app_controller_desktop_thread_sessions_collaboration_impl.dart'; // ignore_for_file: invalid_use_of_visible_for_testing_member, invalid_use_of_protected_member extension AppControllerDesktopThreadSessions on AppController { + Map latestRoutingResolutionForSession(String sessionKey) { + final normalizedSessionKey = normalizedAssistantSessionKeyInternal( + sessionKey, + ); + return latestRoutingResolutionBySessionInternal[normalizedSessionKey] ?? + const {}; + } + int assistantSkillCountForSession(String sessionKey) { final normalizedSessionKey = normalizedAssistantSessionKeyInternal( sessionKey, @@ -97,8 +105,14 @@ extension AppControllerDesktopThreadSessions on AppController { sessionKey, ); final target = assistantExecutionTargetForSession(normalizedSessionKey); + final latestRouting = latestRoutingResolutionForSession(normalizedSessionKey); + final latestResolvedModel = + latestRouting['resolvedModel']?.toString().trim() ?? ''; if (target == AssistantExecutionTarget.singleAgent || target == AssistantExecutionTarget.auto) { + if (latestResolvedModel.isNotEmpty) { + return latestResolvedModel; + } if (singleAgentUsesAiChatFallbackForSession(normalizedSessionKey)) { final recordModel = assistantThreadRecordsInternal[normalizedSessionKey] @@ -376,6 +390,67 @@ extension AppControllerDesktopThreadSessions on AppController { final target = assistantExecutionTargetForSession(normalizedSessionKey); if (target == AssistantExecutionTarget.singleAgent || target == AssistantExecutionTarget.auto) { + final latestRouting = latestRoutingResolutionForSession(normalizedSessionKey); + final latestResolvedExecutionTarget = + latestRouting['resolvedExecutionTarget']?.toString().trim() ?? ''; + final latestResolvedEndpointTarget = + latestRouting['resolvedEndpointTarget']?.toString().trim() ?? ''; + final latestResolvedProviderId = + latestRouting['resolvedProviderId']?.toString().trim() ?? ''; + final latestResolvedModel = + latestRouting['resolvedModel']?.toString().trim() ?? ''; + final primaryLabel = target == AssistantExecutionTarget.auto + ? 'Auto' + : target.label; + final actualDetailPrefix = target == AssistantExecutionTarget.auto + ? appText('当前: ', 'Current: ') + : ''; + if (target == AssistantExecutionTarget.auto && + latestResolvedExecutionTarget.isEmpty) { + return AssistantThreadConnectionState( + executionTarget: target, + status: RuntimeConnectionStatus.offline, + primaryLabel: primaryLabel, + detailLabel: appText('待服务端路由', 'Waiting for server routing'), + ready: false, + pairingRequired: false, + gatewayTokenMissing: false, + lastError: null, + ); + } + if (target == AssistantExecutionTarget.auto && + latestResolvedExecutionTarget.isNotEmpty) { + final detail = switch (latestResolvedExecutionTarget) { + 'gateway' => joinConnectionPartsInternal([ + latestResolvedEndpointTarget.isEmpty + ? appText('OpenClaw Gateway', 'OpenClaw Gateway') + : latestResolvedEndpointTarget, + latestResolvedModel, + ]), + 'multi-agent' => joinConnectionPartsInternal([ + appText('Multi-Agent', 'Multi-Agent'), + latestResolvedModel, + ]), + _ => joinConnectionPartsInternal([ + latestResolvedProviderId.isEmpty + ? appText('Single Agent', 'Single Agent') + : latestResolvedProviderId, + latestResolvedModel, + ]), + }; + return AssistantThreadConnectionState( + executionTarget: target, + status: RuntimeConnectionStatus.connected, + primaryLabel: primaryLabel, + detailLabel: detail.isEmpty + ? appText('待服务端路由', 'Waiting for server routing') + : '$actualDetailPrefix$detail', + ready: true, + pairingRequired: false, + gatewayTokenMissing: false, + lastError: null, + ); + } final provider = singleAgentProviderForSession(normalizedSessionKey); final resolvedProvider = singleAgentResolvedProviderForSession( normalizedSessionKey, @@ -410,12 +485,6 @@ extension AppControllerDesktopThreadSessions on AppController { '当前线程的外部 Agent ACP 连接尚未就绪。', 'The external Agent ACP connection for this thread is not ready yet.', ); - final primaryLabel = target == AssistantExecutionTarget.auto - ? 'Auto' - : target.label; - final actualDetailPrefix = target == AssistantExecutionTarget.auto - ? appText('当前: ', 'Current: ') - : ''; return AssistantThreadConnectionState( executionTarget: target, status: providerReady || fallbackReady diff --git a/lib/runtime/go_agent_core_client.dart b/lib/runtime/go_agent_core_client.dart index 042fdac6..55d567bd 100644 --- a/lib/runtime/go_agent_core_client.dart +++ b/lib/runtime/go_agent_core_client.dart @@ -20,6 +20,32 @@ class GoAgentCoreCapabilities { final Map raw; } +class GoAgentCoreSyncedProvider { + const GoAgentCoreSyncedProvider({ + required this.providerId, + required this.label, + required this.endpoint, + required this.authorizationHeader, + required this.enabled, + }); + + final String providerId; + final String label; + final String endpoint; + final String authorizationHeader; + final bool enabled; + + Map toJson() { + return { + 'providerId': providerId.trim(), + 'label': label.trim(), + 'endpoint': endpoint.trim(), + 'authorizationHeader': authorizationHeader.trim(), + 'enabled': enabled, + }; + } +} + enum GoAgentCoreRoutingMode { auto, explicit } class GoAgentCoreAvailableSkill { @@ -55,6 +81,7 @@ class GoAgentCoreRoutingConfig { required this.explicitSkills, required this.allowSkillInstall, required this.availableSkills, + this.installApproval, }); const GoAgentCoreRoutingConfig.auto({ @@ -65,7 +92,8 @@ class GoAgentCoreRoutingConfig { explicitProviderId = '', explicitModel = '', explicitSkills = const [], - allowSkillInstall = false; + allowSkillInstall = false, + installApproval = null; final GoAgentCoreRoutingMode mode; final String preferredGatewayTarget; @@ -75,6 +103,7 @@ class GoAgentCoreRoutingConfig { final List explicitSkills; final bool allowSkillInstall; final List availableSkills; + final GoAgentCoreSkillInstallApproval? installApproval; bool get isAuto => mode == GoAgentCoreRoutingMode.auto; @@ -97,6 +126,27 @@ class GoAgentCoreRoutingConfig { 'availableSkills': availableSkills .map((item) => item.toJson()) .toList(growable: false), + if (installApproval != null) 'installApproval': installApproval!.toJson(), + }; + } +} + +class GoAgentCoreSkillInstallApproval { + const GoAgentCoreSkillInstallApproval({ + required this.requestId, + required this.approvedSkillKeys, + }); + + final String requestId; + final List approvedSkillKeys; + + Map toJson() { + return { + 'requestId': requestId.trim(), + 'approvedSkillKeys': approvedSkillKeys + .map((item) => item.trim()) + .where((item) => item.isNotEmpty) + .toList(growable: false), }; } } @@ -168,7 +218,11 @@ class GoAgentCoreSessionRequest { bool get hasInlineAttachments => inlineAttachments.isNotEmpty; + GoAgentCoreRoutingConfig get effectiveRouting => + routing ?? _synthesizedRouting(); + Map toAcpParams() { + final resolvedRouting = effectiveRouting; final params = { 'sessionId': sessionId, 'threadId': threadId, @@ -210,7 +264,7 @@ class GoAgentCoreSessionRequest { 'aiGatewayBaseUrl': aiGatewayBaseUrl.trim(), if (aiGatewayApiKey.trim().isNotEmpty) 'aiGatewayApiKey': aiGatewayApiKey.trim(), - if (routing != null) 'routing': routing!.toJson(), + 'routing': resolvedRouting.toJson(), if (_usesGatewaySessionMode(mode)) ...{ 'executionTarget': target.promptValue, if (agentId.trim().isNotEmpty) 'agentId': agentId.trim(), @@ -219,6 +273,47 @@ class GoAgentCoreSessionRequest { }; return params; } + + GoAgentCoreRoutingConfig _synthesizedRouting() { + final preferredGatewayTarget = switch (target) { + AssistantExecutionTarget.remote => 'remote', + _ => 'local', + }; + final explicitExecutionTarget = switch (target) { + AssistantExecutionTarget.local => 'local', + AssistantExecutionTarget.remote => 'remote', + AssistantExecutionTarget.singleAgent => 'singleAgent', + AssistantExecutionTarget.auto => '', + }; + final explicitProviderId = provider == SingleAgentProvider.auto + ? '' + : provider.providerId; + final explicitModelValue = model.trim(); + final explicitSkillsValue = selectedSkills + .map((item) => item.trim()) + .where((item) => item.isNotEmpty) + .toList(growable: false); + final hasExplicitSelection = + explicitExecutionTarget.isNotEmpty || + explicitProviderId.isNotEmpty || + explicitModelValue.isNotEmpty || + explicitSkillsValue.isNotEmpty; + if (!hasExplicitSelection) { + return GoAgentCoreRoutingConfig.auto( + preferredGatewayTarget: preferredGatewayTarget, + ); + } + return GoAgentCoreRoutingConfig( + mode: GoAgentCoreRoutingMode.explicit, + preferredGatewayTarget: preferredGatewayTarget, + explicitExecutionTarget: explicitExecutionTarget, + explicitProviderId: explicitProviderId, + explicitModel: explicitModelValue, + explicitSkills: explicitSkillsValue, + allowSkillInstall: false, + availableSkills: const [], + ); + } } const String _gatewaySessionMode = 'gateway-chat'; @@ -302,6 +397,15 @@ class GoAgentCoreRunResult { bool get needsSkillInstall => _boolValue(raw['needsSkillInstall']) ?? false; + String get skillInstallRequestId => + raw['skillInstallRequestId']?.toString().trim() ?? ''; + + List> get skillCandidates => + _castMapList(raw['skillCandidates']); + + List> get memorySources => + _castMapList(raw['memorySources']); + WorkspaceRefKind? get resolvedWorkspaceRefKind { final rawValue = raw['resolvedWorkspaceRefKind']?.toString().trim() ?? ''; if (rawValue.isEmpty) { @@ -312,6 +416,8 @@ class GoAgentCoreRunResult { } abstract class GoAgentCoreClient { + Future syncProviders(List providers); + Future loadCapabilities({ required AssistantExecutionTarget target, bool forceRefresh = false, @@ -466,3 +572,13 @@ bool? _boolValue(Object? raw) { } return null; } + +List> _castMapList(Object? raw) { + if (raw is! List) { + return const >[]; + } + return raw + .map((item) => _castMap(item)) + .where((item) => item.isNotEmpty) + .toList(growable: false); +} diff --git a/lib/runtime/go_agent_core_desktop_transport.dart b/lib/runtime/go_agent_core_desktop_transport.dart index f8204d1d..d88fa152 100644 --- a/lib/runtime/go_agent_core_desktop_transport.dart +++ b/lib/runtime/go_agent_core_desktop_transport.dart @@ -22,7 +22,6 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { GoCoreLocator? goCoreLocator, GoAgentCoreProcessStarter? processStarter, }) : _acpClient = acpClient, - _endpointResolver = endpointResolver, _goCoreLocator = goCoreLocator ?? GoCoreLocator(), _processStarter = processStarter ?? @@ -32,11 +31,10 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { arguments, environment: environment, workingDirectory: workingDirectory, - ); + ); }); final GatewayAcpClient _acpClient; - final Uri? Function(AssistantExecutionTarget target) _endpointResolver; final GoCoreLocator _goCoreLocator; final GoAgentCoreProcessStarter _processStarter; @@ -44,12 +42,27 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { Uri? _localEndpoint; Future? _localEndpointFuture; + @override + Future syncProviders(List providers) async { + final endpoint = await _ensureLocalEndpoint(); + if (endpoint == null) { + return; + } + await _acpClient.request( + method: 'xworkmate.providers.sync', + params: { + 'providers': providers.map((item) => item.toJson()).toList(growable: false), + }, + endpointOverride: endpoint, + ); + } + @override Future loadCapabilities({ required AssistantExecutionTarget target, bool forceRefresh = false, }) async { - final endpoint = await _resolveEndpoint(target); + final endpoint = await _ensureLocalEndpoint(); if (endpoint == null) { return const GoAgentCoreCapabilities.empty(); } @@ -70,10 +83,7 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { GoAgentCoreSessionRequest request, { required void Function(GoAgentCoreSessionUpdate update) onUpdate, }) async { - final routingResult = await _resolveRouting(request); - final endpoint = await _resolveEndpoint( - _targetForRouting(request, routingResult), - ); + final endpoint = await _ensureLocalEndpoint(); if (endpoint == null) { throw const GatewayAcpException( 'Missing Go Agent-core endpoint', @@ -84,7 +94,7 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { String? completedMessage; final response = await _acpClient.request( method: request.resumeSession ? 'session.message' : 'session.start', - params: _resolvedParams(request, routingResult), + params: request.toAcpParams(), endpointOverride: endpoint, onNotification: (notification) { final update = goAgentCoreUpdateFromNotification(notification); @@ -100,11 +110,8 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { onUpdate(update); }, ); - final mergedResponse = routingResult == null - ? response - : mergeGoAgentCoreResponseResult(response, routingResult); return goAgentCoreRunResultFromResponse( - mergedResponse, + response, streamedText: streamedText, completedMessage: completedMessage, ); @@ -116,7 +123,7 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { required String sessionId, required String threadId, }) async { - final endpoint = await _resolveEndpoint(target); + final endpoint = await _ensureLocalEndpoint(); if (endpoint == null) { return; } @@ -133,7 +140,7 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { required String sessionId, required String threadId, }) async { - final endpoint = await _resolveEndpoint(target); + final endpoint = await _ensureLocalEndpoint(); if (endpoint == null) { return; } @@ -159,13 +166,6 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { } } - Future _resolveEndpoint(AssistantExecutionTarget target) async { - if (target == AssistantExecutionTarget.singleAgent) { - return _ensureLocalEndpoint(); - } - return _endpointResolver(target); - } - Future _ensureLocalEndpoint() async { if (_localEndpoint != null) { return _localEndpoint; @@ -237,129 +237,4 @@ class GoAgentCoreDesktopTransport implements GoAgentCoreClient { await dispose(); return null; } - - Future?> _resolveRouting( - GoAgentCoreSessionRequest request, - ) async { - final routing = request.routing; - if (routing == null) { - return null; - } - final endpoint = await _ensureLocalEndpoint(); - if (endpoint == null) { - return null; - } - try { - final response = await _acpClient.request( - method: 'xworkmate.routing.resolve', - params: request.toAcpParams(), - endpointOverride: endpoint, - ); - return _castRoutingResult(response['result']); - } on Object { - return null; - } - } - - Map _resolvedParams( - GoAgentCoreSessionRequest request, - Map? routingResult, - ) { - final params = Map.from(request.toAcpParams()); - if (routingResult == null || routingResult.isEmpty) { - return params; - } - final resolvedExecutionTarget = - routingResult['resolvedExecutionTarget']?.toString().trim() ?? ''; - final resolvedEndpointTarget = - routingResult['resolvedEndpointTarget']?.toString().trim() ?? ''; - final resolvedProviderId = - routingResult['resolvedProviderId']?.toString().trim() ?? ''; - final resolvedModel = - routingResult['resolvedModel']?.toString().trim() ?? ''; - final resolvedSkills = _castStringList(routingResult['resolvedSkills']); - final routedTarget = _targetForRouting(request, routingResult); - - if (routedTarget != AssistantExecutionTarget.singleAgent) { - if (resolvedExecutionTarget.isNotEmpty) { - params['mode'] = 'gateway-chat'; - } - if (resolvedEndpointTarget.isNotEmpty) { - params['executionTarget'] = resolvedEndpointTarget; - params['resolvedEndpointTarget'] = resolvedEndpointTarget; - } - if (resolvedProviderId.isNotEmpty) { - params['provider'] = resolvedProviderId; - params['resolvedProviderId'] = resolvedProviderId; - } - if (resolvedModel.isNotEmpty) { - params['model'] = resolvedModel; - params['resolvedModel'] = resolvedModel; - } - if (resolvedSkills.isNotEmpty) { - params['selectedSkills'] = resolvedSkills; - params['resolvedSkills'] = resolvedSkills; - } - } - if (resolvedExecutionTarget.isNotEmpty) { - params['resolvedExecutionTarget'] = resolvedExecutionTarget; - } - for (final key in [ - 'skillResolutionSource', - 'memorySources', - 'skillCandidates', - 'needsSkillInstall', - ]) { - if (routingResult.containsKey(key)) { - params[key] = routingResult[key]; - } - } - return params; - } - - AssistantExecutionTarget _targetForRouting( - GoAgentCoreSessionRequest request, - Map? routingResult, - ) { - if (routingResult == null || routingResult.isEmpty) { - return request.target; - } - final resolvedExecutionTarget = - routingResult['resolvedExecutionTarget']?.toString().trim() ?? ''; - if (_isGatewayExecutionTarget(resolvedExecutionTarget)) { - final endpointTarget = - routingResult['resolvedEndpointTarget']?.toString().trim() ?? ''; - return switch (endpointTarget) { - 'local' => AssistantExecutionTarget.local, - 'remote' => AssistantExecutionTarget.remote, - _ => request.target, - }; - } - return AssistantExecutionTarget.singleAgent; - } - - Map _castRoutingResult(Object? raw) { - if (raw is Map) { - return raw; - } - if (raw is Map) { - return raw.cast(); - } - return const {}; - } - - List _castStringList(Object? raw) { - if (raw is! List) { - return const []; - } - return raw - .map((item) => item?.toString().trim() ?? '') - .where((item) => item.isNotEmpty) - .toList(growable: false); - } - - bool _isGatewayExecutionTarget(String value) { - final normalized = value.trim(); - return normalized == 'gateway' || normalized == 'gateway-chat'; - } } diff --git a/lib/web/go_agent_core_web_transport.dart b/lib/web/go_agent_core_web_transport.dart index 7e64b1ac..aad6c2df 100644 --- a/lib/web/go_agent_core_web_transport.dart +++ b/lib/web/go_agent_core_web_transport.dart @@ -12,12 +12,29 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { final WebAcpClient _acpClient; final Uri? Function(AssistantExecutionTarget target) _endpointResolver; + Uri? get _goCoreEndpoint => _endpointResolver(AssistantExecutionTarget.singleAgent); + + @override + Future syncProviders(List providers) async { + final endpoint = _goCoreEndpoint; + if (endpoint == null) { + return; + } + await _acpClient.request( + endpoint: endpoint, + method: 'xworkmate.providers.sync', + params: { + 'providers': providers.map((item) => item.toJson()).toList(growable: false), + }, + ); + } + @override Future loadCapabilities({ required AssistantExecutionTarget target, bool forceRefresh = false, }) async { - final endpoint = _endpointResolver(target); + final endpoint = _goCoreEndpoint; if (endpoint == null) { return const GoAgentCoreCapabilities.empty(); } @@ -35,10 +52,7 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { GoAgentCoreSessionRequest request, { required void Function(GoAgentCoreSessionUpdate update) onUpdate, }) async { - final routingResult = await _resolveRouting(request); - final endpoint = _endpointResolver( - _targetForRouting(request, routingResult), - ); + final endpoint = _goCoreEndpoint; if (endpoint == null) { throw const WebAcpException( 'Missing Go Agent-core endpoint', @@ -50,7 +64,7 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { final response = await _acpClient.request( endpoint: endpoint, method: request.resumeSession ? 'session.message' : 'session.start', - params: _resolvedParams(request, routingResult), + params: request.toAcpParams(), onNotification: (notification) { final update = goAgentCoreUpdateFromNotification(notification); if (update == null) { @@ -65,11 +79,8 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { onUpdate(update); }, ); - final mergedResponse = routingResult == null - ? response - : mergeGoAgentCoreResponseResult(response, routingResult); return goAgentCoreRunResultFromResponse( - mergedResponse, + response, streamedText: streamedText, completedMessage: completedMessage, ); @@ -81,7 +92,7 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { required String sessionId, required String threadId, }) async { - final endpoint = _endpointResolver(target); + final endpoint = _goCoreEndpoint; if (endpoint == null) { return; } @@ -98,7 +109,7 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { required String sessionId, required String threadId, }) async { - final endpoint = _endpointResolver(target); + final endpoint = _goCoreEndpoint; if (endpoint == null) { return; } @@ -111,129 +122,4 @@ class GoAgentCoreWebTransport implements GoAgentCoreClient { @override Future dispose() async {} - - Future?> _resolveRouting( - GoAgentCoreSessionRequest request, - ) async { - final routing = request.routing; - if (routing == null) { - return null; - } - final endpoint = _endpointResolver(AssistantExecutionTarget.singleAgent); - if (endpoint == null) { - return null; - } - try { - final response = await _acpClient.request( - endpoint: endpoint, - method: 'xworkmate.routing.resolve', - params: request.toAcpParams(), - ); - return _castRoutingResult(response['result']); - } on Object { - return null; - } - } - - Map _resolvedParams( - GoAgentCoreSessionRequest request, - Map? routingResult, - ) { - final params = Map.from(request.toAcpParams()); - if (routingResult == null || routingResult.isEmpty) { - return params; - } - final resolvedExecutionTarget = - routingResult['resolvedExecutionTarget']?.toString().trim() ?? ''; - final resolvedEndpointTarget = - routingResult['resolvedEndpointTarget']?.toString().trim() ?? ''; - final resolvedProviderId = - routingResult['resolvedProviderId']?.toString().trim() ?? ''; - final resolvedModel = - routingResult['resolvedModel']?.toString().trim() ?? ''; - final resolvedSkills = _castStringList(routingResult['resolvedSkills']); - final routedTarget = _targetForRouting(request, routingResult); - - if (routedTarget != AssistantExecutionTarget.singleAgent) { - if (resolvedExecutionTarget.isNotEmpty) { - params['mode'] = 'gateway-chat'; - } - if (resolvedEndpointTarget.isNotEmpty) { - params['executionTarget'] = resolvedEndpointTarget; - params['resolvedEndpointTarget'] = resolvedEndpointTarget; - } - if (resolvedProviderId.isNotEmpty) { - params['provider'] = resolvedProviderId; - params['resolvedProviderId'] = resolvedProviderId; - } - if (resolvedModel.isNotEmpty) { - params['model'] = resolvedModel; - params['resolvedModel'] = resolvedModel; - } - if (resolvedSkills.isNotEmpty) { - params['selectedSkills'] = resolvedSkills; - params['resolvedSkills'] = resolvedSkills; - } - } - if (resolvedExecutionTarget.isNotEmpty) { - params['resolvedExecutionTarget'] = resolvedExecutionTarget; - } - for (final key in [ - 'skillResolutionSource', - 'memorySources', - 'skillCandidates', - 'needsSkillInstall', - ]) { - if (routingResult.containsKey(key)) { - params[key] = routingResult[key]; - } - } - return params; - } - - AssistantExecutionTarget _targetForRouting( - GoAgentCoreSessionRequest request, - Map? routingResult, - ) { - if (routingResult == null || routingResult.isEmpty) { - return request.target; - } - final resolvedExecutionTarget = - routingResult['resolvedExecutionTarget']?.toString().trim() ?? ''; - if (_isGatewayExecutionTarget(resolvedExecutionTarget)) { - final endpointTarget = - routingResult['resolvedEndpointTarget']?.toString().trim() ?? ''; - return switch (endpointTarget) { - 'local' => AssistantExecutionTarget.local, - 'remote' => AssistantExecutionTarget.remote, - _ => request.target, - }; - } - return AssistantExecutionTarget.singleAgent; - } - - Map _castRoutingResult(Object? raw) { - if (raw is Map) { - return raw; - } - if (raw is Map) { - return raw.cast(); - } - return const {}; - } - - bool _isGatewayExecutionTarget(String value) { - final normalized = value.trim(); - return normalized == 'gateway' || normalized == 'gateway-chat'; - } - - List _castStringList(Object? raw) { - if (raw is! List) { - return const []; - } - return raw - .map((item) => item?.toString().trim() ?? '') - .where((item) => item.isNotEmpty) - .toList(growable: false); - } } diff --git a/test/runtime/app_controller_ai_gateway_chat_suite_fakes.dart b/test/runtime/app_controller_ai_gateway_chat_suite_fakes.dart index f64e079e..af0ebd6f 100644 --- a/test/runtime/app_controller_ai_gateway_chat_suite_fakes.dart +++ b/test/runtime/app_controller_ai_gateway_chat_suite_fakes.dart @@ -135,6 +135,9 @@ class FakeGoAgentCoreClientInternal implements GoAgentCoreClient { final List requests = []; + @override + Future syncProviders(List providers) async {} + @override Future loadCapabilities({ required AssistantExecutionTarget target, diff --git a/test/runtime/app_controller_assistant_flow_suite.dart b/test/runtime/app_controller_assistant_flow_suite.dart index 40ca52ea..1055bbb1 100644 --- a/test/runtime/app_controller_assistant_flow_suite.dart +++ b/test/runtime/app_controller_assistant_flow_suite.dart @@ -631,6 +631,9 @@ class _FakeGoAgentCoreClient implements GoAgentCoreClient { GoAgentCoreSessionRequest? lastRequest; final void Function(GoAgentCoreSessionRequest request)? onExecute; + @override + Future syncProviders(List providers) async {} + @override Future loadCapabilities({ required AssistantExecutionTarget target, diff --git a/test/runtime/go_agent_core_client_suite.dart b/test/runtime/go_agent_core_client_suite.dart index 287bc96d..15ef61cc 100644 --- a/test/runtime/go_agent_core_client_suite.dart +++ b/test/runtime/go_agent_core_client_suite.dart @@ -95,6 +95,39 @@ void main() { }); }); + test('session request synthesizes routing when caller omits it', () { + const request = GoAgentCoreSessionRequest( + sessionId: 'session-implicit-routing', + threadId: 'thread-implicit-routing', + target: AssistantExecutionTarget.singleAgent, + prompt: 'hello world', + workingDirectory: '/tmp/workspace', + model: 'codex-sonnet', + thinking: '', + selectedSkills: ['PPTX'], + inlineAttachments: [], + localAttachments: [], + aiGatewayBaseUrl: '', + aiGatewayApiKey: '', + agentId: '', + metadata: {}, + provider: SingleAgentProvider.opencode, + ); + + final params = request.toAcpParams(); + + expect(params['routing'], { + 'routingMode': 'explicit', + 'preferredGatewayTarget': 'local', + 'explicitExecutionTarget': 'singleAgent', + 'explicitProviderId': 'opencode', + 'explicitModel': 'codex-sonnet', + 'explicitSkills': const ['PPTX'], + 'allowSkillInstall': false, + 'availableSkills': const >[], + }); + }); + test('routing execution target uses gateway while session mode stays compatible', () { const request = GoAgentCoreSessionRequest( sessionId: 'session-2', @@ -120,6 +153,14 @@ void main() { expect(params['mode'], 'gateway-chat'); expect(params['executionTarget'], 'local'); expect(params['agentId'], 'agent-1'); + expect(params['routing'], { + 'routingMode': 'explicit', + 'preferredGatewayTarget': 'local', + 'explicitExecutionTarget': 'local', + 'explicitSkills': const [], + 'allowSkillInstall': false, + 'availableSkills': const >[], + }); }); test(