From 2b5535e772ea21e51186822db34055a964b9136d Mon Sep 17 00:00:00 2001 From: Haitao Pan Date: Sat, 2 May 2026 21:46:01 +0800 Subject: [PATCH] fix: connect local openclaw gateway before requests --- internal/acp/gateway.go | 76 +++++++++++++++++- internal/acp/orchestrator.go | 3 + internal/acp/routing_test.go | 152 +++++++++++++++++++++++++++++++++++ 3 files changed, 229 insertions(+), 2 deletions(-) diff --git a/internal/acp/gateway.go b/internal/acp/gateway.go index 3749e4f..9b8702d 100644 --- a/internal/acp/gateway.go +++ b/internal/acp/gateway.go @@ -2,14 +2,24 @@ package acp import ( "context" + "crypto/ed25519" + "crypto/rand" + "crypto/sha256" + "encoding/base64" "net/url" "strings" + "sync" "time" "xworkmate-bridge/internal/gatewayruntime" "xworkmate-bridge/internal/shared" ) +var bridgeGatewayIdentity = struct { + sync.Mutex + value gatewayruntime.DeviceIdentity +}{} + func (s *Server) handleGatewayMethod(ctx context.Context, method string, params map[string]any, notify func(map[string]any)) (map[string]any, *shared.RPCError) { switch method { case "xworkmate.gateway.connect": @@ -72,11 +82,11 @@ func handleGatewayConnect( } request = applyProductionGatewayRouting(server, request) request.ReportedRemoteAddress = resolveGatewayReportedRemoteAddress(server, request) - + if server.gateway == nil { server.gateway = gatewayruntime.NewManager() } - + result := server.gateway.Connect(request, notify) return map[string]any{ "ok": result.OK, @@ -166,4 +176,66 @@ func handleGatewayDisconnect( return map[string]any{"accepted": true} } +func ensureProductionGatewayConnected( + server *Server, + mode string, + notify func(map[string]any), +) *shared.RPCError { + normalizedMode := strings.TrimSpace(strings.ToLower(mode)) + if normalizedMode == "" { + normalizedMode = "openclaw" + } + if normalizedMode != "openclaw" { + return nil + } + if server.gateway == nil { + server.gateway = gatewayruntime.NewManager() + } + + request := applyProductionGatewayRouting( + server, + gatewayruntime.ConnectRequest{ + RuntimeID: "xworkmate-bridge-openclaw", + Mode: "openclaw", + ClientID: "xworkmate-bridge", + Locale: "en_US", + UserAgent: "xworkmate-bridge", + Endpoint: gatewayruntime.Endpoint{Host: "127.0.0.1", Port: 18789, TLS: false}, + PackageInfo: gatewayruntime.PackageInfo{AppName: "XWorkmate Bridge", PackageName: "xworkmate-bridge", Version: "bridge", BuildNumber: "0"}, + DeviceInfo: gatewayruntime.DeviceInfo{Platform: "linux", DeviceFamily: "bridge"}, + Identity: newBridgeGatewayIdentity(), + }, + ) + request.ReportedRemoteAddress = resolveGatewayReportedRemoteAddress(server, request) + result := server.gateway.Connect(request, notify) + if result.OK { + return nil + } + message := strings.TrimSpace(shared.StringArg(result.Error, "message", "gateway connect failed")) + code := strings.TrimSpace(shared.StringArg(result.Error, "code", "")) + if code != "" { + message = code + ": " + message + } + return &shared.RPCError{Code: -32002, Message: "GATEWAY_CONNECT_FAILED: " + message} +} + +func newBridgeGatewayIdentity() gatewayruntime.DeviceIdentity { + bridgeGatewayIdentity.Lock() + defer bridgeGatewayIdentity.Unlock() + if strings.TrimSpace(bridgeGatewayIdentity.value.DeviceID) != "" { + return bridgeGatewayIdentity.value + } + publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return gatewayruntime.DeviceIdentity{} + } + sum := sha256.Sum256(publicKey) + bridgeGatewayIdentity.value = gatewayruntime.DeviceIdentity{ + DeviceID: "xworkmate-bridge-" + base64.RawURLEncoding.EncodeToString(sum[:9]), + PublicKeyBase64URL: base64.RawURLEncoding.EncodeToString(publicKey), + PrivateKeyBase64URL: base64.RawURLEncoding.EncodeToString(privateKey), + } + return bridgeGatewayIdentity.value +} + // Helper functions are now in helpers.go diff --git a/internal/acp/orchestrator.go b/internal/acp/orchestrator.go index ab201c4..ba5b526 100644 --- a/internal/acp/orchestrator.go +++ b/internal/acp/orchestrator.go @@ -127,6 +127,9 @@ func (o *SessionOrchestrator) runGateway( if gatewayProvider == "" { return nil, &shared.RPCError{Code: -32602, Message: "GATEWAY_PROVIDER_REQUIRED"} } + if rpcErr := ensureProductionGatewayConnected(o.server, gatewayProvider, notify); rpcErr != nil { + return nil, rpcErr + } params = withResolvedGatewayProvider(params, gatewayProvider) result := o.server.gateway.RequestByMode( gatewayProvider, diff --git a/internal/acp/routing_test.go b/internal/acp/routing_test.go index dbefbbb..3d8846d 100644 --- a/internal/acp/routing_test.go +++ b/internal/acp/routing_test.go @@ -3,13 +3,17 @@ package acp import ( "context" "encoding/json" + "net" "net/http" "net/http/httptest" "os" "path/filepath" "strings" + "sync/atomic" "testing" + "time" + "github.com/gorilla/websocket" "xworkmate-bridge/internal/shared" ) @@ -446,6 +450,47 @@ func TestExecuteSessionTaskExplicitGatewayUsesResolvedGatewayProvider(t *testing } } +func TestExecuteSessionTaskGatewayAutoConnectsLocalOpenClaw(t *testing.T) { + gateway := newAcpFakeOpenClawGateway(t) + defer gateway.Close() + + t.Setenv("GATEWAY_RPC_URL", gateway.URL()) + t.Setenv("BRIDGE_AUTH_TOKEN", "bridge-token") + + server := NewServer() + response, rpcErr := server.executeSessionTask(task{ + req: shared.RPCRequest{ + Method: "session.start", + Params: map[string]any{ + "sessionId": "session-openclaw", + "threadId": "thread-openclaw", + "taskPrompt": "say pong", + "workingDirectory": t.TempDir(), + "routing": map[string]any{ + "routingMode": "explicit", + "explicitExecutionTarget": "gateway", + "preferredGatewayProviderId": "openclaw", + }, + }, + }, + }) + if rpcErr != nil { + t.Fatalf("expected gateway response, got rpc error: %#v", rpcErr) + } + if got := response["output"]; got != "gateway pong" { + t.Fatalf("expected gateway pong output, got %#v", response) + } + if got := response["resolvedGatewayProviderId"]; got != "openclaw" { + t.Fatalf("expected openclaw gateway provider, got %#v", response) + } + if gateway.ConnectCount() != 1 { + t.Fatalf("expected one automatic gateway connect, got %d", gateway.ConnectCount()) + } + if gateway.SessionStartCount() != 1 { + t.Fatalf("expected one session.start request, got %d", gateway.SessionStartCount()) + } +} + func TestExecuteSessionTaskDefaultsExplicitGatewayToOpenClaw(t *testing.T) { server := NewServer() @@ -471,6 +516,113 @@ func TestExecuteSessionTaskDefaultsExplicitGatewayToOpenClaw(t *testing.T) { } } +type acpFakeOpenClawGateway struct { + server *http.Server + listener net.Listener + connectCount atomic.Int32 + sessionStartCount atomic.Int32 +} + +func newAcpFakeOpenClawGateway(t *testing.T) *acpFakeOpenClawGateway { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen fake openclaw gateway: %v", err) + } + fake := &acpFakeOpenClawGateway{listener: listener} + upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} + mux := http.NewServeMux() + mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer func() { + _ = conn.Close() + }() + _ = conn.WriteJSON(map[string]any{ + "type": "event", + "event": "connect.challenge", + "payload": map[string]any{ + "nonce": "nonce-1", + }, + }) + for { + _, payload, err := conn.ReadMessage() + if err != nil { + return + } + var frame map[string]any + if err := json.Unmarshal(payload, &frame); err != nil { + continue + } + if strings.TrimSpace(shared.StringArg(frame, "type", "")) != "req" { + continue + } + id := frame["id"] + switch strings.TrimSpace(shared.StringArg(frame, "method", "")) { + case "connect": + fake.connectCount.Add(1) + _ = conn.WriteJSON(map[string]any{ + "type": "res", + "id": id, + "ok": true, + "payload": map[string]any{ + "server": map[string]any{"host": "127.0.0.1"}, + "snapshot": map[string]any{ + "sessionDefaults": map[string]any{"mainSessionKey": "main"}, + }, + "auth": map[string]any{ + "role": "operator", + "scopes": []string{"operator.read", "operator.write"}, + "deviceToken": "device-token-1", + }, + }, + }) + case "session.start": + fake.sessionStartCount.Add(1) + _ = conn.WriteJSON(map[string]any{ + "type": "res", + "id": id, + "ok": true, + "payload": map[string]any{ + "success": true, + "output": "gateway pong", + }, + }) + default: + _ = conn.WriteJSON(map[string]any{ + "type": "res", + "id": id, + "ok": true, + "payload": map[string]any{}, + }) + } + } + }) + fake.server = &http.Server{Handler: mux, ReadHeaderTimeout: 2 * time.Second} + go func() { + _ = fake.server.Serve(listener) + }() + return fake +} + +func (f *acpFakeOpenClawGateway) URL() string { + return "ws://" + f.listener.Addr().String() + "/" +} + +func (f *acpFakeOpenClawGateway) ConnectCount() int { + return int(f.connectCount.Load()) +} + +func (f *acpFakeOpenClawGateway) SessionStartCount() int { + return int(f.sessionStartCount.Load()) +} + +func (f *acpFakeOpenClawGateway) Close() { + _ = f.server.Close() +} + func TestExecuteSessionTaskAutoRoutingUsesBridgeProductionProviderOrder(t *testing.T) { workspaceDir := filepath.Join(t.TempDir(), "workspace") if err := os.MkdirAll(workspaceDir, 0o755); err != nil {