Fix gateway provider routing resolution
This commit is contained in:
parent
f9c9b3cb68
commit
c9afbe7423
@ -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...),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user