Fix gateway provider routing resolution

This commit is contained in:
Haitao Pan 2026-04-26 10:17:38 +08:00
parent f9c9b3cb68
commit c9afbe7423
2 changed files with 53 additions and 19 deletions

View File

@ -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...),
}
}

View File

@ -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)
}
}