fix: connect local openclaw gateway before requests
This commit is contained in:
parent
54142767e4
commit
2b5535e772
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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 {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user