diff --git a/arkclaw/message.py b/arkclaw/message.py index 57403bb..bea0f8d 100644 --- a/arkclaw/message.py +++ b/arkclaw/message.py @@ -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", @@ -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") @@ -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 @@ -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: diff --git a/tests/test_message.py b/tests/test_message.py index 229b98f..dd33144 100644 --- a/tests/test_message.py +++ b/tests/test_message.py @@ -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, @@ -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( { @@ -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(