diff --git a/internal/acp/orchestrator.go b/internal/acp/orchestrator.go index d8463bd..b46f923 100644 --- a/internal/acp/orchestrator.go +++ b/internal/acp/orchestrator.go @@ -56,7 +56,7 @@ func (o *SessionOrchestrator) Process(ctx context.Context, method string, params }) if res.TargetID == "gateway" { - result, rpcErr := o.runGateway(ctx, method, params, turnID, notify) + result, rpcErr := o.runGateway(ctx, method, params, res, turnID, notify) if rpcErr != nil { return nil, rpcErr } @@ -96,6 +96,7 @@ func (o *SessionOrchestrator) runGateway( ctx context.Context, method string, params map[string]any, + routing RoutingResult, turnID string, notify func(map[string]any), ) (map[string]any, *shared.RPCError) { @@ -103,10 +104,11 @@ func (o *SessionOrchestrator) runGateway( return nil, &shared.RPCError{Code: -32001, Message: "GATEWAY_NOT_INITIALIZED"} } - gatewayProvider := strings.TrimSpace(shared.StringArg(params, "gatewayProvider", "")) + gatewayProvider := resolvedGatewayProviderID(params, routing) if gatewayProvider == "" { return nil, &shared.RPCError{Code: -32602, Message: "GATEWAY_PROVIDER_REQUIRED"} } + params = withResolvedGatewayProvider(params, gatewayProvider) result := o.server.gateway.RequestByMode( gatewayProvider, method, @@ -135,18 +137,51 @@ func (o *SessionOrchestrator) runGateway( return payload, nil } +func resolvedGatewayProviderID(params map[string]any, routing RoutingResult) string { + for _, value := range []string{ + routing.GatewayProviderID, + shared.StringArg(params, "gatewayProvider", ""), + shared.StringArg(params, "gatewayProviderId", ""), + } { + if provider := strings.TrimSpace(value); provider != "" { + return provider + } + } + routingParams := shared.AsMap(params["routing"]) + for _, key := range []string{ + "gatewayProvider", + "gatewayProviderId", + "preferredGatewayProviderId", + } { + if provider := strings.TrimSpace(shared.StringArg(routingParams, key, "")); provider != "" { + return provider + } + } + return "" +} + +func withResolvedGatewayProvider(params map[string]any, gatewayProvider string) map[string]any { + next := make(map[string]any, len(params)+2) + for key, value := range params { + next[key] = value + } + next["gatewayProvider"] = gatewayProvider + next["gatewayProviderId"] = gatewayProvider + return next +} + func (o *SessionOrchestrator) formatUnavailable(res RoutingResult) map[string]any { return map[string]any{ - "success": false, - "status": "unavailable", - "unavailable": true, - "unavailableCode": res.UnavailableCode, - "unavailableMessage": res.UnavailableMsg, + "success": false, + "status": "unavailable", + "unavailable": true, + "unavailableCode": res.UnavailableCode, + "unavailableMessage": res.UnavailableMsg, "resolvedExecutionTarget": res.TargetID, - "resolvedProviderId": res.ProviderID, + "resolvedProviderId": res.ProviderID, "resolvedGatewayProviderId": res.GatewayProviderID, - "resolvedModel": res.Model, - "resolvedSkills": append([]string(nil), res.Skills...), + "resolvedModel": res.Model, + "resolvedSkills": append([]string(nil), res.Skills...), } } diff --git a/internal/acp/routing_test.go b/internal/acp/routing_test.go index f60d692..fa882a7 100644 --- a/internal/acp/routing_test.go +++ b/internal/acp/routing_test.go @@ -385,7 +385,7 @@ func TestExecuteSessionTaskExplicitProviderRequiresAdvertisedBridgeProvider(t *t } } -func TestExecuteSessionTaskExplicitGatewayIgnoresExplicitProvider(t *testing.T) { +func TestExecuteSessionTaskExplicitGatewayUsesResolvedGatewayProvider(t *testing.T) { server := NewServer() response, rpcErr := server.executeSessionTask(task{ @@ -400,20 +400,19 @@ func TestExecuteSessionTaskExplicitGatewayIgnoresExplicitProvider(t *testing.T) "explicitExecutionTarget": "gateway", "explicitProviderId": "claude", "preferredGatewayProviderId": "openclaw", - "gatewayProvider": "openclaw", }, }, }, }) if rpcErr == nil { - t.Fatalf("expected gateway provider required rpc error, got response: %v", response) + t.Fatalf("expected gateway connectivity rpc error, got response: %v", response) } - if rpcErr.Message != "GATEWAY_PROVIDER_REQUIRED" { - t.Fatalf("expected GATEWAY_PROVIDER_REQUIRED, got %q", rpcErr.Message) + if rpcErr.Message == "GATEWAY_PROVIDER_REQUIRED" { + t.Fatalf("expected resolved gateway provider to be reused, got %q", rpcErr.Message) } } -func TestExecuteSessionTaskRequiresExplicitGatewayProvider(t *testing.T) { +func TestExecuteSessionTaskDefaultsExplicitGatewayToOpenClaw(t *testing.T) { server := NewServer() _, rpcErr := server.executeSessionTask(task{ @@ -431,10 +430,10 @@ func TestExecuteSessionTaskRequiresExplicitGatewayProvider(t *testing.T) { }, }) if rpcErr == nil { - t.Fatal("expected gateway provider required error") + t.Fatal("expected gateway connectivity error") } - if rpcErr.Message != "GATEWAY_PROVIDER_REQUIRED" { - t.Fatalf("expected GATEWAY_PROVIDER_REQUIRED, got %#v", rpcErr) + if rpcErr.Message == "GATEWAY_PROVIDER_REQUIRED" { + t.Fatalf("expected openclaw default from routing result, got %#v", rpcErr) } }