oclaw/tests/test_ws_gateway.py
oliver 4a23b715a2 重构仓库目录为统一的 runtime 分层并清理历史 openclaw 残留。
本次迁移将网关/通道/工具/技能/脚本与协议资源集中到新结构,统一路径常量与脚本转发机制,减少顶层噪音并保证运行与测试行为一致。

Made-with: Cursor
2026-04-25 01:24:23 +08:00

320 lines
13 KiB
Python

import unittest
from unittest import mock
from fastapi.testclient import TestClient
from oclaw.interfaces.http.fastapi_app import create_app
from oclaw.runtime.gateway import OclawGatewayResult
def _connect_params() -> dict:
return {
"minProtocol": 3,
"maxProtocol": 3,
"client": {"id": "test", "version": "0.0.0", "platform": "pytest", "mode": "operator"},
"role": "operator",
"scopes": ["operator.read"],
"caps": [],
"commands": [],
"permissions": {},
}
class WsGatewayTests(unittest.TestCase):
def setUp(self) -> None:
self.client = TestClient(create_app())
def test_ws_requires_connect_first(self) -> None:
with self.client.websocket_connect("/ws") as ws:
# server sends connect.challenge first
evt = ws.receive_json()
assert evt["type"] == "event"
assert evt["event"] == "connect.challenge"
ws.send_json({"type": "req", "id": "1", "method": "agent", "params": {"message": "hi", "idempotencyKey": "k"}})
res = ws.receive_json()
assert res["type"] == "res"
assert res["id"] == "1"
assert res["ok"] is False
assert (res.get("error") or {}).get("code") == "INVALID_REQUEST"
def test_ws_connect_returns_hello_ok(self) -> None:
with self.client.websocket_connect("/ws") as ws:
ws.receive_json() # connect.challenge
ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()})
res = ws.receive_json()
assert res["type"] == "res"
assert res["id"] == "c1"
assert res["ok"] is True
hello = res["payload"]
assert hello["type"] == "hello-ok"
assert int(hello["protocol"]) >= 1
assert "snapshot" in hello and "policy" in hello
assert "sessions.send" in list((hello.get("features") or {}).get("methods") or [])
def test_ws_rejects_invalid_frame(self) -> None:
with self.client.websocket_connect("/ws") as ws:
ws.receive_json() # connect.challenge
ws.send_text("not-json")
res = ws.receive_json()
assert res["type"] == "res"
assert res["ok"] is False
def test_ws_agent_run_emits_events_and_response(self) -> None:
def _fake_handle_turn(self, **kwargs): # noqa: ANN001
on_progress = kwargs.get("on_progress")
on_token = kwargs.get("on_token")
on_tool_ui = kwargs.get("on_tool_ui")
rid = kwargs.get("run_id") or "run_test"
if callable(on_progress):
on_progress("oclaw: think (1)…")
if callable(on_tool_ui):
on_tool_ui("skill", {"ok": True})
if callable(on_token):
on_token("hello")
return OclawGatewayResult(
run_id=str(rid),
reply_text="done",
trace_id="trace_test",
elapsed_ms=5,
mode="sync_direct",
task_id=None,
selected_specialist="generalist",
interaction_mode="comprehensive",
)
with mock.patch("oclaw.runtime.gateway.OclawGateway.handle_turn", new=_fake_handle_turn):
with self.client.websocket_connect("/ws") as ws:
ws.receive_json() # connect.challenge
ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()})
ws.receive_json()
ws.send_json(
{
"type": "req",
"id": "a1",
"method": "agent",
"params": {"message": "hi", "idempotencyKey": "k", "sessionId": "sess1"},
}
)
# Events may race ahead of the response; drain until we see res(id=a1).
res = None
for _ in range(10):
msg = ws.receive_json()
if msg.get("type") == "res" and msg.get("id") == "a1":
res = msg
break
assert res is not None
assert res["ok"] is True
payload = res["payload"]
assert payload["runId"]
# We should see at least one agent.event (token/progress/tool) and the terminal end.
seen_agent_event = False
seen_end = False
for _ in range(6):
msg = ws.receive_json()
if msg.get("type") != "event":
continue
if msg.get("event") != "agent.event":
continue
seen_agent_event = True
data = (msg.get("payload") or {}).get("data") or {}
if data.get("phase") == "end":
seen_end = True
break
assert seen_agent_event is True
assert seen_end is True
def test_ws_sessions_send_routes_to_agent_flow(self) -> None:
def _fake_handle_turn(self, **kwargs): # noqa: ANN001
rid = kwargs.get("run_id") or "run_test"
return OclawGatewayResult(
run_id=str(rid),
reply_text="done",
trace_id="trace_sessions_send",
elapsed_ms=7,
mode="sync_direct",
task_id=None,
selected_specialist="generalist",
interaction_mode="comprehensive",
)
with mock.patch("oclaw.runtime.gateway.OclawGateway.handle_turn", new=_fake_handle_turn):
with self.client.websocket_connect("/ws") as ws:
ws.receive_json()
ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()})
ws.receive_json()
ws.send_json(
{
"type": "req",
"id": "s1",
"method": "sessions.send",
"params": {"key": "sessX", "message": "hello from sessions", "idempotencyKey": "idem1"},
}
)
res = None
for _ in range(10):
msg = ws.receive_json()
if msg.get("type") == "res" and msg.get("id") == "s1":
res = msg
break
assert res is not None
assert res["ok"] is True
assert str((res.get("payload") or {}).get("runId") or "").strip() != ""
def test_ws_agent_run_alias_works(self) -> None:
def _fake_handle_turn(self, **kwargs): # noqa: ANN001
rid = kwargs.get("run_id") or "run_alias"
return OclawGatewayResult(
run_id=str(rid),
reply_text="ok",
trace_id="trace_alias",
elapsed_ms=2,
mode="sync_direct",
task_id=None,
selected_specialist="generalist",
interaction_mode="comprehensive",
)
with mock.patch("oclaw.runtime.gateway.OclawGateway.handle_turn", new=_fake_handle_turn):
with self.client.websocket_connect("/ws") as ws:
ws.receive_json()
ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()})
hello = ws.receive_json()
methods = list(((hello.get("payload") or {}).get("features") or {}).get("methods") or [])
assert "agent.run" in methods
ws.send_json(
{
"type": "req",
"id": "r1",
"method": "agent.run",
"params": {"message": "alias", "idempotencyKey": "idem_alias", "sessionId": "sessAlias"},
}
)
res = None
for _ in range(10):
msg = ws.receive_json()
if msg.get("type") == "res" and msg.get("id") == "r1":
res = msg
break
assert res is not None
assert res["ok"] is True
def test_ws_chat_send_ack_then_final_event(self) -> None:
def _fake_handle_turn(self, **kwargs): # noqa: ANN001
rid = kwargs.get("run_id") or "run_chat_send"
on_token = kwargs.get("on_token")
if callable(on_token):
on_token("he")
on_token("llo")
return OclawGatewayResult(
run_id=str(rid),
reply_text="hello",
trace_id="trace_chat_send",
elapsed_ms=3,
mode="sync_direct",
task_id=None,
selected_specialist="generalist",
interaction_mode="comprehensive",
)
with mock.patch("oclaw.runtime.gateway.OclawGateway.handle_turn", new=_fake_handle_turn):
with self.client.websocket_connect("/ws") as ws:
ws.receive_json()
ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()})
ws.receive_json()
ws.send_json(
{
"type": "req",
"id": "cs1",
"method": "chat.send",
"params": {"sessionKey": "sess-chat", "message": "hi", "idempotencyKey": "idem-chat-1"},
}
)
ack = ws.receive_json()
assert ack.get("type") == "res"
assert ack.get("id") == "cs1"
assert ack.get("ok") is True
assert str((ack.get("payload") or {}).get("status") or "") == "started"
run_id = str((ack.get("payload") or {}).get("runId") or "")
assert run_id.strip() != ""
seen_final = False
for _ in range(12):
msg = ws.receive_json()
if msg.get("type") != "event":
continue
if msg.get("event") != "chat":
continue
payload = msg.get("payload") or {}
if str(payload.get("runId") or "") != run_id:
continue
if str(payload.get("state") or "") == "final":
seen_final = True
break
assert seen_final is True
def test_ws_chat_send_emits_delta_and_session_tool(self) -> None:
def _fake_handle_turn(self, **kwargs): # noqa: ANN001
rid = kwargs.get("run_id") or "run_chat_send_stream"
on_token = kwargs.get("on_token")
on_tool_ui = kwargs.get("on_tool_ui")
if callable(on_tool_ui):
on_tool_ui("skill", {"ok": True})
if callable(on_token):
on_token("Hel")
on_token("lo")
return OclawGatewayResult(
run_id=str(rid),
reply_text="Hello",
trace_id="trace_chat_send_stream",
elapsed_ms=3,
mode="sync_direct",
task_id=None,
selected_specialist="generalist",
interaction_mode="comprehensive",
)
with mock.patch("oclaw.runtime.gateway.OclawGateway.handle_turn", new=_fake_handle_turn):
with self.client.websocket_connect("/ws") as ws:
ws.receive_json()
ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()})
ws.receive_json()
ws.send_json(
{
"type": "req",
"id": "cs2",
"method": "chat.send",
"params": {"sessionKey": "sess-chat2", "message": "hi", "idempotencyKey": "idem-chat-2"},
}
)
ack = ws.receive_json()
assert ack.get("type") == "res"
assert ack.get("id") == "cs2"
assert ack.get("ok") is True
seen_delta = False
seen_tool = False
seen_final = False
for _ in range(30):
msg = ws.receive_json()
if msg.get("type") != "event":
continue
if msg.get("event") == "session.tool":
seen_tool = True
if msg.get("event") != "chat":
continue
payload = msg.get("payload") or {}
st = str(payload.get("state") or "")
if st == "delta":
seen_delta = True
if st == "final":
seen_final = True
break
assert seen_delta is True
assert seen_tool is True
assert seen_final is True