diff --git a/browser_gateway/agent.py b/browser_gateway/agent.py index 84a3fed..ea8a48d 100644 --- a/browser_gateway/agent.py +++ b/browser_gateway/agent.py @@ -96,7 +96,7 @@ class OutboundAgent: return self._stop_requested.wait(seconds) def _connect(self) -> websocket.WebSocket: - url = self.config.platform_url + url = self.config.platform_url.rstrip("/") + "/v1/agent" if url.startswith("https://"): url = "wss://" + url[len("https://"):] elif url.startswith("http://"): @@ -226,14 +226,12 @@ class OutboundAgent: def _execute_task(self, method: str, path: str, payload: Any) -> tuple[int, bytes, str]: from .server.http import RequestError, route_gateway_request - body: dict = {} - if payload: - try: - body = json.loads(payload) - except ValueError as exc: - return 400, b"", f"task payload invalid: {exc}" - if not isinstance(body, dict): - return 400, b"", "task payload must be a JSON object" + if payload is None: + body: dict = {} + elif isinstance(payload, dict): + body = payload + else: + return 400, b"", "task payload must be a JSON object" try: result = route_gateway_request(self.gateway, method, path, {}, body) except RequestError as exc: diff --git a/browser_gateway/test_agent.py b/browser_gateway/test_agent.py index 6058f6e..08f349f 100644 --- a/browser_gateway/test_agent.py +++ b/browser_gateway/test_agent.py @@ -6,9 +6,11 @@ import threading import time import unittest from importlib import import_module -from unittest.mock import Mock +from unittest.mock import Mock, patch import websocket +from websockets.datastructures import Headers +from websockets.http11 import Response from websockets.sync.server import serve agent_module = import_module(f"{__package__}.agent") @@ -36,12 +38,13 @@ class FakePlatform: self.keep_open = threading.Event() def url(self) -> str: - return f"ws://127.0.0.1:{self.port}/v1/agent" + return f"http://127.0.0.1:{self.port}" def _process_request(self, connection, request): + if request.path != "/v1/agent": + return Response(404, "Not Found", Headers(), b"invalid agent path\r\n") if self.reject: - from websockets.http11 import Response - return Response(403, "Forbidden", b"invalid access key\r\n") + return Response(403, "Forbidden", Headers(), b"invalid access key\r\n") return None def _handler(self, ws) -> None: @@ -113,6 +116,21 @@ def make_agent(gateway: Mock, platform: FakePlatform) -> OutboundAgent: class OutboundAgentConnectionTests(unittest.TestCase): + def test_platform_root_url_resolves_to_agent_endpoint(self) -> None: + cases = ( + ("http://127.0.0.1:3010", "ws://127.0.0.1:3010/v1/agent"), + ("http://127.0.0.1:3010/", "ws://127.0.0.1:3010/v1/agent"), + ("https://platform.example", "wss://platform.example/v1/agent"), + ("https://platform.example/", "wss://platform.example/v1/agent"), + ("https://platform.example/creator/", "wss://platform.example/creator/v1/agent"), + ) + for root, expected in cases: + with self.subTest(root=root): + agent = OutboundAgent(make_gateway(), AgentConfig(platform_url=root, access_key="test-key")) + with patch.object(agent_module.websocket, "create_connection") as connect: + agent._connect() + self.assertEqual(connect.call_args.args[0], expected) + def test_heartbeats_flow_with_version_until_stopped(self) -> None: platform = FakePlatform() agent = make_agent(make_gateway(), platform) @@ -143,6 +161,7 @@ class OutboundAgentConnectionTests(unittest.TestCase): time.sleep(0.01) self.assertGreaterEqual(agent.reconnect_attempts, 2, "auth failure must trigger retry") self.assertTrue(agent.last_error, "auth failure must be reported") + self.assertIn("403", agent.last_error, "fixture must return the intended rejection, not an internal error") finally: agent.stop() runner.join(timeout=5) @@ -150,6 +169,65 @@ class OutboundAgentConnectionTests(unittest.TestCase): class OutboundAgentTaskTests(unittest.TestCase): + def test_json_object_payload_is_forwarded_without_decoding(self) -> None: + gateway = make_gateway() + agent = OutboundAgent(gateway, AgentConfig(platform_url="http://127.0.0.1:3010", access_key="test-key")) + payload = {"alias": "account-a", "target_uid": "7354367890123456789", "text": "测试消息", "options": {"enabled": True}} + with patch.object(gateway_module, "route_gateway_request", return_value={"success": True}) as route: + status, body, error = agent._execute_task("POST", "/v1/social/action", payload) + self.assertEqual(status, 200) + self.assertEqual(json.loads(body), {"success": True}) + self.assertEqual(error, "") + route.assert_called_once_with(gateway, "POST", "/v1/social/action", {}, payload) + + def test_null_and_empty_object_payloads_are_valid(self) -> None: + for payload in (None, {}): + with self.subTest(payload=payload): + gateway = make_gateway() + agent = OutboundAgent(gateway, AgentConfig(platform_url="http://127.0.0.1:3010", access_key="test-key")) + with patch.object(gateway_module, "route_gateway_request", return_value={}) as route: + status, _, error = agent._execute_task("GET", "/v1/info", payload) + self.assertEqual(status, 200) + self.assertEqual(error, "") + route.assert_called_once_with(gateway, "GET", "/v1/info", {}, {}) + + def test_non_object_payloads_are_rejected_before_routing(self) -> None: + for payload in ('{"alias":"account-a"}', "", [], [1], 0, 1, False, True): + with self.subTest(payload=payload): + agent = OutboundAgent(make_gateway(), AgentConfig(platform_url="http://127.0.0.1:3010", access_key="test-key")) + with patch.object(gateway_module, "route_gateway_request", return_value={}) as route: + status, body, error = agent._execute_task("POST", "/v1/social/action", payload) + self.assertEqual(status, 400) + self.assertEqual(body, b"") + self.assertIn("JSON object", error) + route.assert_not_called() + + def test_post_task_frame_preserves_json_object_and_result(self) -> None: + platform = FakePlatform() + gateway = make_gateway() + agent = make_agent(gateway, platform) + payload = {"alias": "account-a", "target_uid": "7354367890123456789", "text": "测试消息"} + result_body = {"success": True, "message_id": "7354367890123456790"} + with patch.object(gateway_module, "route_gateway_request", return_value=result_body) as route: + runner = threading.Thread(target=agent.run, daemon=True) + runner.start() + try: + self.wait_until_connected(agent) + platform.latest_connection().send(json.dumps({ + "type": "task", "id": 17, "method": "POST", "path": "/v1/social/action", "payload": payload, + }, ensure_ascii=False)) + result = platform.wait_for(lambda f: f.get("type") == "result" and f.get("id") == 17) + self.assertIsNotNone(result, "POST result missing") + self.assertEqual(result.get("status"), 200) + import base64 + self.assertEqual(json.loads(base64.b64decode(result.get("body"))), result_body) + self.assertEqual(result.get("error"), "") + route.assert_called_once_with(gateway, "POST", "/v1/social/action", {}, payload) + finally: + agent.stop() + runner.join(timeout=5) + platform.close() + def test_task_dispatches_route_and_returns_result(self) -> None: platform = FakePlatform() gateway = make_gateway()