Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 24 additions & 8 deletions arkclaw/message.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,14 +175,14 @@ def _build_websocket_url(*, endpoint: str, chat_token: str, claw_instance_id: st
return f"{base}{separator}{query}"


def _build_connect_message(*, request_id: str | None = None) -> str:
def _build_connect_message(*, request_id: str | None = None, protocol_version: int = 4) -> str:
payload = {
"type": "req",
"id": request_id or str(uuid4()),
"method": "connect",
"params": {
"minProtocol": 3,
"maxProtocol": 3,
"minProtocol": protocol_version,
"maxProtocol": protocol_version,
"client": {
"id": "openclaw-control-ui",
"version": "dev",
Expand Down Expand Up @@ -234,6 +234,7 @@ def __init__(
receive_timeout: float = 30.0,
connect_retries: int = 2,
session_key: str = "agent:main:main",
protocol_version: int = 4,
) -> None:
if connect_retries < 0:
raise ValidationError("connect_retries must be >= 0")
Expand All @@ -247,6 +248,9 @@ def __init__(
self.receive_timeout = receive_timeout
self.connect_retries = connect_retries
self.session_key = session_key
if protocol_version not in (3, 4):
raise ValidationError(f"protocol_version must be 3 or 4, got {protocol_version}")
self.protocol_version = protocol_version

self._websocket_module: Any | None = None
self._timeout_exc: type[BaseException] = TimeoutError
Expand Down Expand Up @@ -381,15 +385,27 @@ def _ensure_connection(self, *, refresh_access: bool) -> None:
except Exception as exc:
raise ValidationError(f"WebSocket chat connection failed: {type(exc).__name__}") from exc
ws.settimeout(self.receive_timeout)
try:
ws.send(_build_connect_message())
self._wait_for_connect_ack(ws)
except Exception:
# Try configured protocol first, fall back to the other on mismatch
fallback_protocol = 3 if self.protocol_version == 4 else 4
protocols_to_try = [self.protocol_version, fallback_protocol]
last_error: Exception | None = None
for proto in protocols_to_try:
try:
ws.send(_build_connect_message(protocol_version=proto))
self._wait_for_connect_ack(ws)
last_error = None
break
except ValidationError as exc:
if 'PROTOCOL_MISMATCH' not in str(exc):
raise
last_error = exc
continue
if last_error:
try:
ws.close()
except Exception:
pass
raise
raise last_error
self._ws = ws

def _refresh_chat_access(self) -> None:
Expand Down
168 changes: 167 additions & 1 deletion tests/test_message.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
import unittest
from unittest.mock import MagicMock, patch

from arkclaw import ArkClawClient
from arkclaw import ArkClawClient, ValidationError
from arkclaw.message import (
ArkClawMessageSession,
_build_chat_send_message,
Expand Down Expand Up @@ -66,6 +66,21 @@ def test_build_connect_message_uses_webchat_client(self) -> None:
self.assertEqual(payload["method"], "connect")
self.assertEqual(payload["params"]["client"]["id"], "openclaw-control-ui")

def test_build_connect_message_defaults_to_v4(self) -> None:
payload = json.loads(_build_connect_message(request_id="connect-v3"))
self.assertEqual(payload["params"]["minProtocol"], 4)
self.assertEqual(payload["params"]["maxProtocol"], 4)

def test_build_connect_message_v4(self) -> None:
payload = json.loads(_build_connect_message(request_id="connect-v4", protocol_version=4))
self.assertEqual(payload["params"]["minProtocol"], 4)
self.assertEqual(payload["params"]["maxProtocol"], 4)

def test_build_connect_message_v3_explicit(self) -> None:
payload = json.loads(_build_connect_message(request_id="connect-v3-exp", protocol_version=3))
self.assertEqual(payload["params"]["minProtocol"], 3)
self.assertEqual(payload["params"]["maxProtocol"], 3)

def test_render_stream_message_pretty_formats_assistant(self) -> None:
data = json.dumps(
{
Expand Down Expand Up @@ -231,6 +246,157 @@ def close(self) -> None:
self.assertEqual(result["receive_timeout"], False)
self.assertEqual(mock_ws_mod.create_connection.call_count, 2)

def test_session_defaults_to_v4_protocol(self) -> None:
session = self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False)
self.assertEqual(session.protocol_version, 4)

def test_session_accepts_v3_protocol(self) -> None:
session = self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False, protocol_version=3)
self.assertEqual(session.protocol_version, 3)

def test_session_accepts_v4_protocol(self) -> None:
session = self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False, protocol_version=4)
self.assertEqual(session.protocol_version, 4)

def test_session_rejects_invalid_protocol(self) -> None:
with self.assertRaises(ValidationError):
self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False, protocol_version=2)

@patch("arkclaw.message._load_websocket_module")
def test_send_message_uses_v3_connect(self, mock_load_ws) -> None:
self.client.workflows.prepare_chat_access = MagicMock(
return_value={
"token": {
"ChatToken": "secret",
"Endpoint": "example.com/ws",
"InstanceId": "ci-xxx",
}
}
)

sent_messages: list[str] = []

class FakeConnection:
def __init__(self) -> None:
self.responses = [
json.dumps({"type": "event", "event": "connect.challenge"}),
json.dumps({"type": "res", "ok": True}),
json.dumps({"type": "event", "event": "chat", "payload": {"text": "reply"}}),
]

def send(self, data: str) -> None:
sent_messages.append(data)

def recv(self) -> str:
return self.responses.pop(0)

def settimeout(self, timeout: float) -> None:
return None

def close(self) -> None:
return None

mock_ws_mod = MagicMock()
mock_ws_mod.create_connection.return_value = FakeConnection()
mock_load_ws.return_value = mock_ws_mod

session = self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False, protocol_version=3)
session.send_message("hello")

connect_msg = json.loads(sent_messages[0])
self.assertEqual(connect_msg["params"]["minProtocol"], 3)
self.assertEqual(connect_msg["params"]["maxProtocol"], 3)

@patch("arkclaw.message._load_websocket_module")
def test_send_message_uses_v4_connect(self, mock_load_ws) -> None:
self.client.workflows.prepare_chat_access = MagicMock(
return_value={
"token": {
"ChatToken": "secret",
"Endpoint": "example.com/ws",
"InstanceId": "ci-xxx",
}
}
)

sent_messages: list[str] = []

class FakeConnection:
def __init__(self) -> None:
self.responses = [
json.dumps({"type": "event", "event": "connect.challenge"}),
json.dumps({"type": "res", "ok": True}),
json.dumps({"type": "event", "event": "chat", "payload": {"text": "reply"}}),
]

def send(self, data: str) -> None:
sent_messages.append(data)

def recv(self) -> str:
return self.responses.pop(0)

def settimeout(self, timeout: float) -> None:
return None

def close(self) -> None:
return None

mock_ws_mod = MagicMock()
mock_ws_mod.create_connection.return_value = FakeConnection()
mock_load_ws.return_value = mock_ws_mod

session = self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False, protocol_version=4)
session.send_message("hello")

connect_msg = json.loads(sent_messages[0])
self.assertEqual(connect_msg["params"]["minProtocol"], 4)
self.assertEqual(connect_msg["params"]["maxProtocol"], 4)

@patch("arkclaw.message._load_websocket_module")
def test_stream_message_uses_v4_connect(self, mock_load_ws) -> None:
self.client.workflows.prepare_chat_access = MagicMock(
return_value={
"token": {
"ChatToken": "secret",
"Endpoint": "example.com/ws",
"InstanceId": "ci-xxx",
}
}
)

sent_messages: list[str] = []

class FakeConnection:
def __init__(self) -> None:
self.responses = [
json.dumps({"type": "event", "event": "connect.challenge"}),
json.dumps({"type": "res", "ok": True}),
json.dumps({"type": "event", "event": "chat", "payload": {"state": "final", "text": "done"}}),
]

def send(self, data: str) -> None:
sent_messages.append(data)

def recv(self) -> str:
return self.responses.pop(0)

def settimeout(self, timeout: float) -> None:
return None

def close(self) -> None:
return None

mock_ws_mod = MagicMock()
mock_ws_mod.create_connection.return_value = FakeConnection()
mock_load_ws.return_value = mock_ws_mod

session = self.client.create_message_session(space_id="csi-xxx", instance_id="ci-xxx", wait=False, protocol_version=4)
session.stream_message("hello", on_event=lambda _: None)

connect_msg = json.loads(sent_messages[0])
self.assertEqual(connect_msg["params"]["minProtocol"], 4)
self.assertEqual(connect_msg["params"]["maxProtocol"], 4)

@patch("arkclaw.message._load_websocket_module")
def test_stream_message_emits_relevant_frames(self, mock_load_ws) -> None:
self.client.workflows.prepare_chat_access = MagicMock(
Expand Down
Loading