diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 6724dc3a..3a99bbca 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -19,14 +19,25 @@ permissions: jobs: test: + name: Python (${{ matrix.name }}) runs-on: ${{ matrix.os }} timeout-minutes: 20 strategy: fail-fast: false matrix: - os: ${{ fromJSON('["ubuntu-latest","windows-latest"]') }} - # CI concentrates on newer runtimes (3.11/3.12 still supported per pyproject requires-python). - python-version: ${{ fromJSON('["3.13","3.14"]') }} + include: + - name: minimum, 3.11 + os: ubuntu-latest + python-version: "3.11" + coverage: false + - name: latest, 3.14 + coverage + os: ubuntu-latest + python-version: "3.14" + coverage: true + - name: Windows, 3.14 + os: windows-latest + python-version: "3.14" + coverage: false steps: - uses: actions/checkout@v4 @@ -47,10 +58,21 @@ jobs: run: uv sync --all-extras --dev - name: Lint with ruff - run: uv run ruff check nanobot --select F + if: matrix.coverage + run: uv run ruff check nanobot tests - - name: Run tests - run: uv run python -m pytest tests/ --cov=nanobot --cov-report=term-missing:skip-covered + - name: Run tests with coverage + if: matrix.coverage + run: >- + uv run python -m pytest tests/ + --cov=nanobot --cov-report=term-missing:skip-covered + --durations=25 --durations-min=1.0 + + - name: Run compatibility tests + if: ${{ !matrix.coverage }} + run: >- + uv run python -m pytest tests/ + --durations=25 --durations-min=1.0 webui: runs-on: ubuntu-latest diff --git a/nanobot/api/server.py b/nanobot/api/server.py index ae639d75..2f21131f 100644 --- a/nanobot/api/server.py +++ b/nanobot/api/server.py @@ -41,6 +41,26 @@ __all__ = ( API_SESSION_KEY = "api:default" API_CHAT_ID = "default" +_AGENT_LOOP_KEY = web.AppKey[Any]("agent_loop") +_MODEL_NAME_KEY = web.AppKey[str]("model_name") +_REQUEST_TIMEOUT_KEY = web.AppKey[float]("request_timeout") +_SESSION_LOCKS_KEY = web.AppKey[dict]("session_locks") +_MISSING = object() + + +def _app_value( + app: Any, + key: web.AppKey[Any], + legacy_key: str, + default: Any = _MISSING, +) -> Any: + """Read typed aiohttp state while accepting lightweight dict test doubles.""" + try: + return app[key] + except KeyError: + if default is _MISSING: + return app[legacy_key] + return app.get(legacy_key, default) # --------------------------------------------------------------------------- @@ -209,9 +229,14 @@ async def handle_chat_completions(request: web.Request) -> web.Response: if not isinstance(content_type, str): content_type = "" - agent_loop = request.app["agent_loop"] - timeout_s: float = request.app.get("request_timeout", 120.0) - model_name: str = request.app.get("model_name", "nanobot") + agent_loop = _app_value(request.app, _AGENT_LOOP_KEY, "agent_loop") + timeout_s: float = _app_value( + request.app, + _REQUEST_TIMEOUT_KEY, + "request_timeout", + 120.0, + ) + model_name: str = _app_value(request.app, _MODEL_NAME_KEY, "model_name", "nanobot") stream = False try: @@ -238,7 +263,11 @@ async def handle_chat_completions(request: web.Request) -> web.Response: return _error_json(400, f"Only configured model '{model_name}' is available") session_key = f"api:{session_id}" if session_id else API_SESSION_KEY - session_locks: dict[str, asyncio.Lock] = request.app["session_locks"] + session_locks: dict[str, asyncio.Lock] = _app_value( + request.app, + _SESSION_LOCKS_KEY, + "session_locks", + ) session_lock = session_locks.setdefault(session_key, asyncio.Lock()) logger.info( @@ -366,7 +395,7 @@ async def handle_chat_completions(request: web.Request) -> web.Response: async def handle_models(request: web.Request) -> web.Response: """GET /v1/models""" - model_name = request.app.get("model_name", "nanobot") + model_name = _app_value(request.app, _MODEL_NAME_KEY, "model_name", "nanobot") return web.json_response( { "object": "list", @@ -407,10 +436,10 @@ def create_app( api_key: Optional API key for Bearer-token authentication on API routes. """ app = web.Application(client_max_size=20 * 1024 * 1024) # 20MB for base64 images - app["agent_loop"] = agent_loop - app["model_name"] = model_name - app["request_timeout"] = request_timeout - app["session_locks"] = {} # per-user locks, keyed by session_key + app[_AGENT_LOOP_KEY] = agent_loop + app[_MODEL_NAME_KEY] = model_name + app[_REQUEST_TIMEOUT_KEY] = request_timeout + app[_SESSION_LOCKS_KEY] = {} # per-user locks, keyed by session_key @web.middleware async def auth_middleware(request: web.Request, handler) -> web.StreamResponse: diff --git a/tests/agent/test_dream_tools.py b/tests/agent/test_dream_tools.py index 530a90fe..22e5c02e 100644 --- a/tests/agent/test_dream_tools.py +++ b/tests/agent/test_dream_tools.py @@ -1,7 +1,7 @@ -from nanobot.config.schema import Config -from nanobot.agent.tools.loader import ToolLoader from nanobot.agent.tools.context import ToolContext +from nanobot.agent.tools.loader import ToolLoader from nanobot.agent.tools.registry import ToolRegistry +from nanobot.config.schema import Config def test_tool_loader_scope_memory_only_returns_memory_tools(): diff --git a/tests/agent/test_git_store.py b/tests/agent/test_git_store.py index b8bfc47e..09b16b60 100644 --- a/tests/agent/test_git_store.py +++ b/tests/agent/test_git_store.py @@ -1,10 +1,9 @@ """Tests for GitStore — git-backed version control for memory files.""" + import pytest -from pathlib import Path - -from nanobot.utils.gitstore import GitStore, CommitInfo +from nanobot.utils.gitstore import CommitInfo, GitStore TRACKED = ["SOUL.md", "USER.md", "memory/MEMORY.md"] @@ -64,7 +63,11 @@ class TestBuildGitignore: content = gs._build_gitignore() assert "!a.md\n" in content assert "!b.md\n" in content - dir_lines = [l for l in content.split("\n") if l.startswith("!") and l.endswith("/")] + dir_lines = [ + line + for line in content.split("\n") + if line.startswith("!") and line.endswith("/") + ] assert dir_lines == [] diff --git a/tests/agent/test_loop_runner_integration.py b/tests/agent/test_loop_runner_integration.py index 18092a77..9ce7291b 100644 --- a/tests/agent/test_loop_runner_integration.py +++ b/tests/agent/test_loop_runner_integration.py @@ -565,6 +565,7 @@ async def test_subagent_max_iterations_announces_existing_fallback(tmp_path, mon workspace=tmp_path, bus=bus, max_tool_result_chars=_MAX_TOOL_RESULT_CHARS, + max_iterations=2, ) mgr._announce_result = AsyncMock() diff --git a/tests/agent/test_mcp_reconnect_crash.py b/tests/agent/test_mcp_reconnect_crash.py index bea76533..9e2eb59d 100644 --- a/tests/agent/test_mcp_reconnect_crash.py +++ b/tests/agent/test_mcp_reconnect_crash.py @@ -27,7 +27,8 @@ from nanobot.bus.queue import MessageBus from nanobot.config.schema import MCPServerConfig from nanobot.security import network as security_network -_IDLE_TIMEOUT_SECONDS = 5 +_IDLE_TIMEOUT_SECONDS = 0.25 +_IDLE_EXPIRY_GRACE_SECONDS = 0.25 _TOOL_TIMEOUT_SECONDS = 10 @@ -179,7 +180,7 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url): assert "Hello, first" in output # Wait for the server-side idle timeout to terminate the session. - await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1) + await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + _IDLE_EXPIRY_GRACE_SECONDS) output = await asyncio.create_task(tool.execute(name="second")) assert "Hello, second" in output @@ -207,7 +208,7 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash( assert isinstance(tool, MCPToolWrapper) await asyncio.create_task(tool.execute(name="first")) - await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1) + await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + _IDLE_EXPIRY_GRACE_SECONDS) reconnect_started = asyncio.Event() finish_reconnect = asyncio.Event() diff --git a/tests/agent/test_skill_creator_scripts.py b/tests/agent/test_skill_creator_scripts.py index 4207c6fd..cd84edd7 100644 --- a/tests/agent/test_skill_creator_scripts.py +++ b/tests/agent/test_skill_creator_scripts.py @@ -1,10 +1,8 @@ import importlib -import shutil import sys import zipfile from pathlib import Path - SCRIPT_DIR = Path("nanobot/skills/skill-creator/scripts").resolve() if str(SCRIPT_DIR) not in sys.path: sys.path.insert(0, str(SCRIPT_DIR)) diff --git a/tests/agent/test_stop_preserves_context.py b/tests/agent/test_stop_preserves_context.py index c7e766be..e38c88a7 100644 --- a/tests/agent/test_stop_preserves_context.py +++ b/tests/agent/test_stop_preserves_context.py @@ -12,14 +12,12 @@ from __future__ import annotations import asyncio from pathlib import Path from types import SimpleNamespace -from typing import Any -from unittest.mock import MagicMock, patch, AsyncMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest from nanobot.agent.loop import AgentLoop from nanobot.bus.queue import MessageBus -from nanobot.providers.base import LLMProvider def _make_provider(): @@ -38,8 +36,8 @@ def _make_loop(tmp_path: Path) -> AgentLoop: provider = _make_provider() with patch("nanobot.agent.loop.ContextBuilder"), \ patch("nanobot.agent.loop.SessionManager"), \ - patch("nanobot.agent.loop.SubagentManager") as MockSubMgr: - MockSubMgr.return_value.cancel_by_session = AsyncMock(return_value=0) + patch("nanobot.agent.loop.SubagentManager") as mock_subagent_manager: + mock_subagent_manager.return_value.cancel_by_session = AsyncMock(return_value=0) return AgentLoop(bus=bus, provider=provider, workspace=tmp_path) @@ -110,8 +108,8 @@ async def test_dispatch_cancellation_restores_checkpoint(): with patch("nanobot.agent.loop.ContextBuilder"), \ patch("nanobot.agent.loop.SessionManager"), \ - patch("nanobot.agent.loop.SubagentManager") as MockSubMgr: - MockSubMgr.return_value.cancel_by_session = AsyncMock(return_value=0) + patch("nanobot.agent.loop.SubagentManager") as mock_subagent_manager: + mock_subagent_manager.return_value.cancel_by_session = AsyncMock(return_value=0) loop = AgentLoop(bus=bus, provider=provider, workspace=workspace) checkpoint_key = loop._RUNTIME_CHECKPOINT_KEY diff --git a/tests/channels/test_channel_manager_delta_coalescing.py b/tests/channels/test_channel_manager_delta_coalescing.py index df8c3f0e..6e03df5e 100644 --- a/tests/channels/test_channel_manager_delta_coalescing.py +++ b/tests/channels/test_channel_manager_delta_coalescing.py @@ -352,6 +352,11 @@ class TestProgressFiltering: content="legacy progress-shaped message", metadata={"_progress": True}, )) + await bus.publish_outbound(OutboundMessage( + channel="mock", + chat_id="chat1", + content="processing sentinel", + )) task = asyncio.create_task(manager._dispatch_outbound()) try: @@ -366,7 +371,9 @@ class TestProgressFiltering: except asyncio.CancelledError: pass - assert manager.channels["mock"]._send_mock.await_count == 0 + send_mock = manager.channels["mock"]._send_mock + assert send_mock.await_count == 1 + assert send_mock.await_args.args[0].content == "processing sentinel" @pytest.mark.asyncio async def test_channel_override_can_enable_tool_hints(self, manager, bus): diff --git a/tests/channels/test_feishu_domain.py b/tests/channels/test_feishu_domain.py index caa1c414..87d78510 100644 --- a/tests/channels/test_feishu_domain.py +++ b/tests/channels/test_feishu_domain.py @@ -1,8 +1,6 @@ """Tests for Feishu/Lark domain configuration.""" from unittest.mock import MagicMock -import pytest - from nanobot.bus.queue import MessageBus from nanobot.channels.feishu import FeishuChannel, FeishuConfig diff --git a/tests/channels/test_feishu_mention.py b/tests/channels/test_feishu_mention.py index fb81f229..660cdb5e 100644 --- a/tests/channels/test_feishu_mention.py +++ b/tests/channels/test_feishu_mention.py @@ -2,8 +2,6 @@ from types import SimpleNamespace -import pytest - from nanobot.channels.feishu import FeishuChannel diff --git a/tests/channels/test_telegram_channel.py b/tests/channels/test_telegram_channel.py index c052f6e7..4d3cfbdf 100644 --- a/tests/channels/test_telegram_channel.py +++ b/tests/channels/test_telegram_channel.py @@ -629,7 +629,9 @@ async def test_send_delta_stream_end_treats_not_modified_as_success() -> None: @pytest.mark.asyncio -async def test_send_delta_stream_end_does_not_fallback_on_network_timeout() -> None: +async def test_send_delta_stream_end_does_not_fallback_on_network_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: """TimedOut during HTML edit should propagate, never fall back to plain text.""" from telegram.error import TimedOut @@ -638,6 +640,7 @@ async def test_send_delta_stream_end_does_not_fallback_on_network_timeout() -> N MessageBus(), ) channel._app = _FakeApp(lambda: None) + monkeypatch.setattr("nanobot.channels.telegram._SEND_RETRY_BASE_DELAY", 0) # _call_with_retry retries TimedOut up to 3 times, so the mock will be called # multiple times – but all calls must be with parse_mode="HTML" (no plain fallback). channel._app.bot.edit_message_text = AsyncMock(side_effect=TimedOut("network timeout")) @@ -650,6 +653,7 @@ async def test_send_delta_stream_end_does_not_fallback_on_network_timeout() -> N # no plain-text fallback call should have been made. for call in channel._app.bot.edit_message_text.call_args_list: assert call.kwargs.get("parse_mode") == "HTML" + assert channel._app.bot.edit_message_text.await_count == 3 # Buffer should still be present (not cleaned up on error) assert "123" in channel._stream_bufs diff --git a/tests/channels/test_websocket_channel.py b/tests/channels/test_websocket_channel.py index a743ce2c..553577ae 100644 --- a/tests/channels/test_websocket_channel.py +++ b/tests/channels/test_websocket_channel.py @@ -1,18 +1,17 @@ """Unit and lightweight integration tests for the WebSocket channel.""" import asyncio -import functools import json import time from pathlib import Path from typing import Any from unittest.mock import AsyncMock, MagicMock -import httpx import pytest import websockets from websockets.exceptions import ConnectionClosed from websockets.frames import Close +from ws_test_client import http_get as _http_get from nanobot.bus.events import OUTBOUND_META_AGENT_UI, OutboundMessage from nanobot.bus.outbound_events import ( @@ -169,13 +168,6 @@ def isolate_webui_workspace_state(tmp_path, monkeypatch) -> None: ) -async def _http_get(url: str, headers: dict[str, str] | None = None) -> httpx.Response: - """Run GET in a thread to avoid blocking the asyncio loop shared with websockets.""" - return await asyncio.to_thread( - functools.partial(httpx.get, url, headers=headers or {}, timeout=5.0, trust_env=False) - ) - - @pytest.mark.asyncio async def test_send_session_updated_broadcasts_to_other_webui_connections(bus) -> None: class Conn: diff --git a/tests/channels/test_websocket_http_routes.py b/tests/channels/test_websocket_http_routes.py index a2f03a7e..b67fb8fc 100644 --- a/tests/channels/test_websocket_http_routes.py +++ b/tests/channels/test_websocket_http_routes.py @@ -1,7 +1,6 @@ """End-to-end tests for the embedded webui's HTTP routes on the WebSocket channel.""" import asyncio -import functools import json import random import socket @@ -12,8 +11,9 @@ from typing import Any from unittest.mock import AsyncMock, MagicMock from urllib.parse import quote, urlencode -import httpx import pytest +from ws_test_client import InProcessHttpChannel +from ws_test_client import http_get as _http_get from nanobot.bus.events import OutboundMessage from nanobot.channels.base import BaseChannel @@ -134,7 +134,7 @@ def _ch( local_trigger_pending_ids=local_trigger_pending_ids, channel_feature_action=channel_feature_action, ) - return WebSocketChannel(cfg, bus, gateway=gateway) + return InProcessHttpChannel(cfg, bus, gateway=gateway) @pytest.fixture() @@ -144,14 +144,6 @@ def bus() -> MagicMock: return b -async def _http_get( - url: str, headers: dict[str, str] | None = None -) -> httpx.Response: - return await asyncio.to_thread( - functools.partial(httpx.get, url, headers=headers or {}, timeout=5.0, trust_env=False) - ) - - def _seed_session(workspace: Path, key: str = "websocket:test") -> SessionManager: sm = SessionManager(workspace) s = Session(key=key) @@ -211,7 +203,6 @@ async def test_bootstrap_returns_token_for_localhost( maxMessageBytes=1_048_576, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get("http://127.0.0.1:29901/webui/bootstrap") assert resp.status_code == 200 @@ -248,7 +239,6 @@ async def test_sessions_routes_require_bearer_token( sm = _seed_session(tmp_path, key="websocket:abc") channel = _ch(bus, session_manager=sm, port=29902) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: # Unauthenticated → 401. deny = await _http_get("http://127.0.0.1:29902/api/sessions") @@ -324,7 +314,6 @@ async def test_session_automations_route_filters_by_webui_session( port=29914, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: deny = await _http_get( "http://127.0.0.1:29914/api/sessions/websocket:abc/automations" @@ -381,7 +370,6 @@ async def test_session_automations_route_ignores_unified_owner( port=29917, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -427,7 +415,6 @@ async def test_session_automations_route_lists_local_triggers( port=port, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -487,7 +474,6 @@ async def test_webui_skills_route_requires_token_and_hides_paths( port=29920, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: deny = await _http_get("http://127.0.0.1:29920/api/webui/skills") assert deny.status_code == 401 @@ -585,7 +571,6 @@ async def test_cli_apps_routes_require_token_and_return_payload( ) channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29912) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: deny = await _http_get("http://127.0.0.1:29912/api/settings/cli-apps") assert deny.status_code == 401 @@ -621,7 +606,6 @@ async def test_nanobot_feature_routes_require_token_and_enable( _stub_matrix_feature(monkeypatch, config_path, channels=["matrix", "websocket"]) channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29916) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: deny = await _http_get("http://127.0.0.1:29916/api/settings/nanobot-features") assert deny.status_code == 401 @@ -1517,7 +1501,6 @@ async def test_cli_apps_catalog_does_not_block_other_webui_http_routes( monkeypatch.setattr("nanobot.webui.settings_routes.cli_apps_payload", slow_payload) channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29935) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -1558,7 +1541,6 @@ async def test_cli_apps_route_supports_installed_only_payload( monkeypatch.setattr("nanobot.webui.settings_routes.cli_apps_payload", payload) channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29936) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -1651,7 +1633,6 @@ async def test_mcp_presets_routes_require_token_and_return_payload( ) channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29913) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: deny = await _http_get("http://127.0.0.1:29913/api/settings/mcp-presets") assert deny.status_code == 401 @@ -1743,7 +1724,6 @@ async def test_sessions_list_only_returns_websocket_sessions_by_default( ) channel = _ch(bus, session_manager=sm, port=29906) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -1769,7 +1749,6 @@ async def test_webui_sidebar_state_routes_are_config_dir_scoped( sm = _seed_session(tmp_path, key="websocket:sidebar") channel = _ch(bus, session_manager=sm, port=29911) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -1820,7 +1799,6 @@ async def test_session_delete_removes_file( append_transcript_object("websocket:doomed", {"event": "user", "chat_id": "doomed", "text": "x"}) channel = _ch(bus, session_manager=sm, port=29903) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -1900,7 +1878,6 @@ async def test_webui_automations_route_lists_all_jobs_and_allows_user_actions( port=port, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: deny = await _http_get(f"{base_url}/api/webui/automations") assert deny.status_code == 401, deny.text @@ -2115,7 +2092,6 @@ async def test_webui_automations_route_manages_local_triggers( port=port, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2193,7 +2169,6 @@ async def test_session_delete_blocks_when_bound_automation_exists( ) channel = _ch(bus, session_manager=sm, cron_service=cron, port=29915) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2238,7 +2213,6 @@ async def test_session_delete_blocks_and_cascades_local_triggers( port=port, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2287,7 +2261,6 @@ async def test_session_delete_can_cascade_bound_automations( ) channel = _ch(bus, session_manager=sm, cron_service=cron, port=29916) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2330,7 +2303,6 @@ async def test_session_delete_blocks_origin_automation_when_unified_enabled( port=29918, ) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2362,7 +2334,6 @@ async def test_session_routes_accept_percent_encoded_websocket_keys( sm = _seed_session(tmp_path, key="websocket:encoded-key") channel = _ch(bus, session_manager=sm, port=29910) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2406,7 +2377,6 @@ async def test_session_messages_hide_persisted_runtime_context( sm.save(session) channel = _ch(bus, session_manager=sm, port=29919) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) response = await _http_get( @@ -2459,7 +2429,6 @@ async def test_webui_thread_resigns_assistant_media_urls( channel = _ch(bus, port=29914) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2497,7 +2466,6 @@ async def test_session_routes_reject_non_websocket_keys( ) channel = _ch(bus, session_manager=sm, port=29909) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2530,7 +2498,6 @@ async def test_session_routes_reject_invalid_key( sm = _seed_session(tmp_path) channel = _ch(bus, session_manager=sm, port=29904) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -2558,7 +2525,6 @@ async def test_static_serves_index_when_dist_present( sm = _seed_session(tmp_path / "ws_state") channel = _ch(bus, session_manager=sm, static_dist_path=dist, port=29905) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: # Bare ``GET /`` is a browser opening the app: it must return the SPA # index.html, not the WS-upgrade handler's 401/426. @@ -2588,7 +2554,6 @@ async def test_static_rejects_path_traversal( secret.write_text("classified") channel = _ch(bus, static_dist_path=dist, port=29906) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get("http://127.0.0.1:29906/../secret.txt") # Normalized by httpx into /secret.txt → falls back to index.html, not 'classified'. @@ -2602,7 +2567,6 @@ async def test_static_rejects_path_traversal( async def test_unknown_route_returns_404(bus: MagicMock) -> None: channel = _ch(bus, port=29907) server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get("http://127.0.0.1:29907/api/unknown") assert resp.status_code == 404 diff --git a/tests/channels/test_websocket_integration.py b/tests/channels/test_websocket_integration.py index 4059c435..f4c06069 100644 --- a/tests/channels/test_websocket_integration.py +++ b/tests/channels/test_websocket_integration.py @@ -60,7 +60,6 @@ def bus() -> MagicMock: async def test_ready_event_fields(bus: MagicMock) -> None: ch = _ch(bus, 29901) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29901/", client_id="c1") as c: r = await c.recv_ready() @@ -76,7 +75,6 @@ async def test_ready_event_fields(bus: MagicMock) -> None: async def test_anonymous_client_gets_generated_id(bus: MagicMock) -> None: ch = _ch(bus, 29902) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29902/", client_id="") as c: r = await c.recv_ready() @@ -90,7 +88,6 @@ async def test_anonymous_client_gets_generated_id(bus: MagicMock) -> None: async def test_each_connection_unique_chat_id(bus: MagicMock) -> None: ch = _ch(bus, 29903) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29903/", client_id="a") as c1: async with WsTestClient("ws://127.0.0.1:29903/", client_id="b") as c2: @@ -107,7 +104,6 @@ async def test_each_connection_unique_chat_id(bus: MagicMock) -> None: async def test_plain_text(bus: MagicMock) -> None: ch = _ch(bus, 29904) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29904/", client_id="p") as c: await c.recv_ready() @@ -125,7 +121,6 @@ async def test_plain_text(bus: MagicMock) -> None: async def test_json_content_field(bus: MagicMock) -> None: ch = _ch(bus, 29905) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29905/", client_id="j") as c: await c.recv_ready() @@ -141,7 +136,6 @@ async def test_json_content_field(bus: MagicMock) -> None: async def test_json_text_and_message_fields(bus: MagicMock) -> None: ch = _ch(bus, 29906) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29906/", client_id="x") as c: await c.recv_ready() @@ -160,7 +154,6 @@ async def test_json_text_and_message_fields(bus: MagicMock) -> None: async def test_empty_payload_ignored(bus: MagicMock) -> None: ch = _ch(bus, 29907) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29907/", client_id="e") as c: await c.recv_ready() @@ -177,7 +170,6 @@ async def test_empty_payload_ignored(bus: MagicMock) -> None: async def test_messages_preserve_order(bus: MagicMock) -> None: ch = _ch(bus, 29908) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29908/", client_id="o") as c: await c.recv_ready() @@ -198,7 +190,6 @@ async def test_messages_preserve_order(bus: MagicMock) -> None: async def test_server_send_message(bus: MagicMock) -> None: ch = _ch(bus, 29909) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29909/", client_id="r") as c: ready = await c.recv_ready() @@ -217,7 +208,6 @@ async def test_server_send_tags_tool_hint_with_kind(bus: MagicMock) -> None: """Tool-hint progress events surface as ``kind: "tool_hint"``.""" ch = _ch(bus, 29919) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29919/", client_id="h") as c: ready = await c.recv_ready() @@ -255,7 +245,6 @@ async def test_server_send_tags_tool_hint_with_kind(bus: MagicMock) -> None: async def test_server_send_with_media_and_reply(bus: MagicMock) -> None: ch = _ch(bus, 29910) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29910/", client_id="m") as c: ready = await c.recv_ready() @@ -279,7 +268,6 @@ async def test_server_send_with_media_and_reply(bus: MagicMock) -> None: async def test_streaming_deltas_and_end(bus: MagicMock) -> None: ch = _ch(bus, 29911, streaming=True) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29911/", client_id="s") as c: cid = (await c.recv_ready()).chat_id @@ -301,7 +289,6 @@ async def test_streaming_deltas_and_end(bus: MagicMock) -> None: async def test_interleaved_streams(bus: MagicMock) -> None: ch = _ch(bus, 29912, streaming=True) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29912/", client_id="i") as c: cid = (await c.recv_ready()).chat_id @@ -329,7 +316,6 @@ async def test_interleaved_streams(bus: MagicMock) -> None: async def test_independent_sessions(bus: MagicMock) -> None: ch = _ch(bus, 29913) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29913/", client_id="u1") as c1: async with WsTestClient("ws://127.0.0.1:29913/", client_id="u2") as c2: @@ -351,7 +337,6 @@ async def test_independent_sessions(bus: MagicMock) -> None: async def test_disconnected_client_cleanup(bus: MagicMock) -> None: ch = _ch(bus, 29914) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29914/", client_id="tmp") as c: chat_id = (await c.recv_ready()).chat_id @@ -373,7 +358,6 @@ async def test_disconnected_client_cleanup(bus: MagicMock) -> None: async def test_static_token_accepted(bus: MagicMock) -> None: ch = _ch(bus, 29915, token="secret") t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29915/", client_id="a", token="secret") as c: assert (await c.recv_ready()).client_id == "a" @@ -386,7 +370,6 @@ async def test_static_token_accepted(bus: MagicMock) -> None: async def test_static_token_rejected(bus: MagicMock) -> None: ch = _ch(bus, 29916, token="correct") t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: with pytest.raises(websockets.exceptions.InvalidStatus) as exc: async with WsTestClient("ws://127.0.0.1:29916/", client_id="b", token="wrong"): @@ -403,7 +386,6 @@ async def test_token_issue_full_flow(bus: MagicMock) -> None: tokenIssuePath="/auth/token", tokenIssueSecret="s", websocketRequiresToken=True) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: # no secret -> 401 _, status = await issue_token(port=29917, issue_path="/auth/token") @@ -439,7 +421,6 @@ async def test_token_issue_full_flow(bus: MagicMock) -> None: async def test_custom_path(bus: MagicMock) -> None: ch = _ch(bus, 29918, path="/my-chat") t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29918/my-chat", client_id="p") as c: assert (await c.recv_ready()).event == "ready" @@ -452,7 +433,6 @@ async def test_custom_path(bus: MagicMock) -> None: async def test_wrong_path_404(bus: MagicMock) -> None: ch = _ch(bus, 29919, path="/ws") t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: with pytest.raises(websockets.exceptions.InvalidStatus) as exc: async with WsTestClient("ws://127.0.0.1:29919/wrong", client_id="x"): @@ -467,7 +447,6 @@ async def test_wrong_path_404(bus: MagicMock) -> None: async def test_trailing_slash_normalized(bus: MagicMock) -> None: ch = _ch(bus, 29920, path="/ws") t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29920/ws/", client_id="s") as c: assert (await c.recv_ready()).event == "ready" @@ -483,7 +462,6 @@ async def test_trailing_slash_normalized(bus: MagicMock) -> None: async def test_large_message(bus: MagicMock) -> None: ch = _ch(bus, 29921) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29921/", client_id="big") as c: await c.recv_ready() @@ -500,7 +478,6 @@ async def test_large_message(bus: MagicMock) -> None: async def test_unicode_roundtrip(bus: MagicMock) -> None: ch = _ch(bus, 29922) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29922/", client_id="u") as c: ready = await c.recv_ready() @@ -521,7 +498,6 @@ async def test_unicode_roundtrip(bus: MagicMock) -> None: async def test_rapid_fire(bus: MagicMock) -> None: ch = _ch(bus, 29923) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29923/", client_id="r") as c: ready = await c.recv_ready() @@ -544,7 +520,6 @@ async def test_rapid_fire(bus: MagicMock) -> None: async def test_invalid_json_as_plain_text(bus: MagicMock) -> None: ch = _ch(bus, 29924) t = asyncio.create_task(ch.start()) - await asyncio.sleep(0.3) try: async with WsTestClient("ws://127.0.0.1:29924/", client_id="j") as c: await c.recv_ready() diff --git a/tests/channels/test_websocket_media_route.py b/tests/channels/test_websocket_media_route.py index 037e420e..937874ba 100644 --- a/tests/channels/test_websocket_media_route.py +++ b/tests/channels/test_websocket_media_route.py @@ -11,15 +11,15 @@ These tests cover the two halves end-to-end plus the adversarial edges from __future__ import annotations import asyncio -import functools import hashlib import hmac from pathlib import Path from typing import Any from unittest.mock import AsyncMock, MagicMock, patch -import httpx import pytest +from ws_test_client import InProcessHttpChannel +from ws_test_client import http_get as _http_get from nanobot.channels.websocket import WebSocketChannel, WebSocketConfig from nanobot.session.manager import Session, SessionManager @@ -67,7 +67,7 @@ def _ch( runtime_surface="browser", runtime_capabilities_overrides=None, ) - return WebSocketChannel(cfg, bus, gateway=gateway) + return InProcessHttpChannel(cfg, bus, gateway=gateway) @pytest.fixture() @@ -86,14 +86,6 @@ def _fake_media_dir(root: Path): return inner -async def _http_get( - url: str, headers: dict[str, str] | None = None -) -> httpx.Response: - return await asyncio.to_thread( - functools.partial(httpx.get, url, headers=headers or {}, timeout=5.0, trust_env=False) - ) - - # --------------------------------------------------------------------------- # gateway.media.sign_media_path: the URL minter # --------------------------------------------------------------------------- @@ -223,7 +215,6 @@ async def test_media_route_serves_signed_file( url_path = channel.gateway.media.sign_media_path(target) assert url_path is not None server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get(f"http://127.0.0.1:29920{url_path}") finally: @@ -256,7 +247,6 @@ async def test_media_route_serves_video_byte_ranges( url_path = channel.gateway.media.sign_media_path(target) assert url_path is not None server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get( f"http://127.0.0.1:29927{url_path}", @@ -288,7 +278,6 @@ async def test_media_route_serves_suffix_video_byte_ranges( url_path = channel.gateway.media.sign_media_path(target) assert url_path is not None server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get( f"http://127.0.0.1:29928{url_path}", @@ -317,7 +306,6 @@ async def test_media_route_rejects_unsatisfiable_byte_range( url_path = channel.gateway.media.sign_media_path(target) assert url_path is not None server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get( f"http://127.0.0.1:29929{url_path}", @@ -357,7 +345,6 @@ async def test_media_route_rejects_bad_signature( forged = f"/api/media/{b64url_encode(forged_mac)}/{payload}" server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get(f"http://127.0.0.1:29921{forged}") finally: @@ -391,7 +378,6 @@ async def test_media_route_rejects_path_traversal_payload( with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media): server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get(f"http://127.0.0.1:29922{url}") finally: @@ -418,7 +404,6 @@ async def test_media_route_404s_missing_file( assert url_path is not None target.unlink() # the file vanishes between signing and fetching server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get(f"http://127.0.0.1:29923{url_path}") finally: @@ -448,7 +433,6 @@ async def test_media_route_degrades_non_image_to_octet_stream( ).digest()[:16] url = f"/api/media/{b64url_encode(mac)}/{payload}" server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get(f"http://127.0.0.1:29924{url}") finally: @@ -476,7 +460,6 @@ async def test_media_route_serves_svg_with_strict_csp( url_path = channel.gateway.media.sign_media_path(target) assert url_path is not None server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: resp = await _http_get(f"http://127.0.0.1:29928{url_path}") finally: @@ -515,7 +498,6 @@ async def test_session_messages_exposes_signed_media_urls( channel = _ch(bus, session_manager=sm, port=29925) with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media): server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) auth = {"Authorization": f"Bearer {token}"} @@ -559,7 +541,6 @@ async def test_session_messages_skips_vanished_media( channel = _ch(bus, session_manager=sm, port=29926) with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media): server_task = asyncio.create_task(channel.start()) - await asyncio.sleep(0.3) try: token = channel.gateway.tokens.issue_api_token(300) resp = await _http_get( diff --git a/tests/channels/test_weixin_channel.py b/tests/channels/test_weixin_channel.py index 1e0ec7b1..f654c3ae 100644 --- a/tests/channels/test_weixin_channel.py +++ b/tests/channels/test_weixin_channel.py @@ -37,6 +37,23 @@ def _make_channel() -> tuple[WeixinChannel, MessageBus]: return channel, bus +@pytest.fixture +def no_qr_poll_delay(monkeypatch: pytest.MonkeyPatch) -> None: + """Keep QR state-machine tests event-driven without one-second polling sleeps.""" + real_sleep = asyncio.sleep + + async def yield_to_loop(_delay: float) -> None: + await real_sleep(0) + + class AsyncioProxy: + sleep = staticmethod(yield_to_loop) + + def __getattr__(self, name: str): + return getattr(asyncio, name) + + monkeypatch.setattr(weixin_mod, "asyncio", AsyncioProxy()) + + def test_make_headers_includes_route_tag_when_configured() -> None: bus = MessageBus() channel = WeixinChannel( @@ -446,7 +463,9 @@ async def test_poll_once_pauses_session_on_expired_errcode() -> None: @pytest.mark.asyncio -async def test_qr_login_refreshes_expired_qr_and_then_succeeds() -> None: +async def test_qr_login_refreshes_expired_qr_and_then_succeeds( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._save_state = lambda: None @@ -478,7 +497,9 @@ async def test_qr_login_refreshes_expired_qr_and_then_succeeds() -> None: @pytest.mark.asyncio -async def test_qr_login_returns_false_after_too_many_expired_qr_codes() -> None: +async def test_qr_login_returns_false_after_too_many_expired_qr_codes( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._print_qr_code = lambda url: None @@ -505,7 +526,9 @@ async def test_qr_login_returns_false_after_too_many_expired_qr_codes() -> None: @pytest.mark.asyncio -async def test_qr_login_switches_polling_base_url_on_redirect_status() -> None: +async def test_qr_login_switches_polling_base_url_on_redirect_status( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._save_state = lambda: None @@ -537,7 +560,9 @@ async def test_qr_login_switches_polling_base_url_on_redirect_status() -> None: @pytest.mark.asyncio -async def test_qr_login_redirect_without_host_keeps_current_polling_base_url() -> None: +async def test_qr_login_redirect_without_host_keeps_current_polling_base_url( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._save_state = lambda: None @@ -569,7 +594,9 @@ async def test_qr_login_redirect_without_host_keeps_current_polling_base_url() - @pytest.mark.asyncio -async def test_qr_login_resets_redirect_base_url_after_qr_refresh() -> None: +async def test_qr_login_resets_redirect_base_url_after_qr_refresh( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._save_state = lambda: None @@ -859,7 +886,9 @@ async def test_get_typing_ticket_failure_uses_backoff_and_cached_ticket(monkeypa @pytest.mark.asyncio -async def test_qr_login_treats_temporary_connect_error_as_wait_and_recovers() -> None: +async def test_qr_login_treats_temporary_connect_error_as_wait_and_recovers( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._save_state = lambda: None @@ -887,7 +916,9 @@ async def test_qr_login_treats_temporary_connect_error_as_wait_and_recovers() -> @pytest.mark.asyncio -async def test_qr_login_treats_5xx_gateway_response_error_as_wait_and_recovers() -> None: +async def test_qr_login_treats_5xx_gateway_response_error_as_wait_and_recovers( + no_qr_poll_delay, +) -> None: channel, _bus = _make_channel() channel._running = True channel._save_state = lambda: None diff --git a/tests/channels/ws_test_client.py b/tests/channels/ws_test_client.py index ec3ba146..e05452b4 100644 --- a/tests/channels/ws_test_client.py +++ b/tests/channels/ws_test_client.py @@ -21,6 +21,58 @@ from typing import Any import httpx import websockets from websockets.asyncio.client import ClientConnection +from websockets.datastructures import Headers +from websockets.http11 import Request as WsRequest + +from nanobot.channels.websocket import WebSocketChannel +from nanobot.webui.http_utils import http_response + +_IN_PROCESS_HTTP_CHANNELS: dict[int, InProcessHttpChannel] = {} + + +class _HttpConnection: + remote_address = ("127.0.0.1", 12345) + + @staticmethod + def respond(status: int, body: str) -> object: + return http_response(body.encode("utf-8"), status=status) + + +class InProcessHttpChannel(WebSocketChannel): + """Exercise gateway HTTP dispatch without booting a socket per route test.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._test_stop_event = asyncio.Event() + _IN_PROCESS_HTTP_CHANNELS[self.config.port] = self + + async def start(self) -> None: + self._running = True + await self._test_stop_event.wait() + self._running = False + + async def stop(self) -> None: + self._test_stop_event.set() + if _IN_PROCESS_HTTP_CHANNELS.get(self.config.port) is self: + _IN_PROCESS_HTTP_CHANNELS.pop(self.config.port, None) + + +async def _in_process_http_get( + channel: InProcessHttpChannel, + request: httpx.Request, +) -> httpx.Response: + ws_request = WsRequest( + request.url.raw_path.decode("ascii"), + Headers(list(request.headers.multi_items())), + ) + response = await channel._dispatch_http(_HttpConnection(), ws_request) + assert response is not None + return httpx.Response( + response.status_code, + headers=list(response.headers.raw_items()), + content=response.body, + request=request, + ) @dataclass @@ -89,11 +141,19 @@ class WsTestClient: self._extra_headers = extra_headers self._ws: ClientConnection | None = None - async def connect(self) -> None: - self._ws = await websockets.connect( - self._uri, - additional_headers=self._extra_headers, - ) + async def connect(self, timeout: float = 2.0) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while True: + try: + self._ws = await websockets.connect( + self._uri, + additional_headers=self._extra_headers, + ) + return + except OSError: + if asyncio.get_running_loop().time() >= deadline: + raise + await asyncio.sleep(0.01) async def close(self) -> None: if self._ws: @@ -186,6 +246,31 @@ class WsTestClient: # -- Token issuance helpers ----------------------------------------------- +async def http_get( + url: str, + headers: dict[str, str] | None = None, +) -> httpx.Response: + """GET a local test server without loading an unused TLS trust store.""" + request = httpx.Request("GET", url, headers=headers or {}) + channel = _IN_PROCESS_HTTP_CHANNELS.get(request.url.port) + if channel is not None: + return await _in_process_http_get(channel, request) + + deadline = asyncio.get_running_loop().time() + 2.0 + while True: + try: + async with httpx.AsyncClient( + timeout=5.0, + trust_env=False, + verify=False, + ) as client: + return await client.get(url, headers=headers or {}) + except httpx.ConnectError: + if asyncio.get_running_loop().time() >= deadline: + raise + await asyncio.sleep(0.01) + + async def issue_token( host: str = "127.0.0.1", port: int = 8765, @@ -201,10 +286,7 @@ async def issue_token( if secret: headers["Authorization"] = f"Bearer {secret}" - loop = asyncio.get_running_loop() - resp = await loop.run_in_executor( - None, lambda: httpx.get(url, headers=headers, timeout=5.0) - ) + resp = await http_get(url, headers) try: data = resp.json() except Exception: diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 00000000..fafe649b --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,51 @@ +"""Cross-suite test infrastructure.""" + +from __future__ import annotations + +import os +import ssl +import sys +from collections.abc import Iterator + +import certifi +import pytest + + +@pytest.fixture(scope="session", autouse=True) +def _use_windows_system_ca_for_default_http_clients() -> Iterator[None]: + """Avoid reparsing certifi's CA bundle for every offline HTTP client. + + Loading certifi takes roughly 0.7 seconds per client on Windows. The test + suite constructs hundreds of clients while mocking their I/O. System roots + preserve certificate verification for accidental local requests; explicit + ``cafile``, ``capath``, and ``cadata`` arguments still use the real loader. + """ + if sys.platform != "win32": + yield + return + + original = ssl.create_default_context + certifi_path = os.path.normcase(os.path.abspath(certifi.where())) + + def create_default_context( + purpose: ssl.Purpose = ssl.Purpose.SERVER_AUTH, + *, + cafile: str | None = None, + capath: str | None = None, + cadata: str | bytes | None = None, + ) -> ssl.SSLContext: + requested_path = os.path.normcase(os.path.abspath(cafile)) if cafile else None + if requested_path == certifi_path and capath is None and cadata is None: + return original(purpose) + return original( + purpose, + cafile=cafile, + capath=capath, + cadata=cadata, + ) + + ssl.create_default_context = create_default_context + try: + yield + finally: + ssl.create_default_context = original diff --git a/tests/providers/test_enforce_role_alternation.py b/tests/providers/test_enforce_role_alternation.py index 1195c258..335062ba 100644 --- a/tests/providers/test_enforce_role_alternation.py +++ b/tests/providers/test_enforce_role_alternation.py @@ -1,6 +1,6 @@ """Tests for LLMProvider._enforce_role_alternation.""" -from nanobot.providers.base import LLMProvider, _SYNTHETIC_USER_CONTENT +from nanobot.providers.base import _SYNTHETIC_USER_CONTENT, LLMProvider class TestEnforceRoleAlternation: diff --git a/tests/providers/test_litellm_kwargs.py b/tests/providers/test_litellm_kwargs.py index 61577c40..c8ba0b6b 100644 --- a/tests/providers/test_litellm_kwargs.py +++ b/tests/providers/test_litellm_kwargs.py @@ -520,8 +520,8 @@ async def test_openrouter_keeps_model_name_intact() -> None: mock_create = AsyncMock(return_value=_fake_chat_response()) spec = find_by_name("openrouter") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_create provider = OpenAICompatProvider( @@ -545,8 +545,8 @@ async def test_aihubmix_strips_model_prefix() -> None: mock_create = AsyncMock(return_value=_fake_chat_response()) spec = find_by_name("aihubmix") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_create provider = OpenAICompatProvider( @@ -570,8 +570,8 @@ async def test_standard_provider_passes_model_through() -> None: mock_create = AsyncMock(return_value=_fake_chat_response()) spec = find_by_name("deepseek") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_create provider = OpenAICompatProvider( @@ -594,8 +594,8 @@ async def test_openai_compat_preserves_extra_content_on_tool_calls() -> None: mock_create = AsyncMock(return_value=_fake_tool_call_response()) spec = find_by_name("gemini") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_create provider = OpenAICompatProvider( @@ -656,8 +656,8 @@ async def test_direct_openai_gpt5_uses_responses_api() -> None: mock_responses = AsyncMock(return_value=_fake_responses_response("from responses")) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -687,8 +687,8 @@ async def test_direct_openai_reasoning_prefers_responses_api() -> None: mock_responses = AsyncMock(return_value=_fake_responses_response("reasoned")) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -716,8 +716,8 @@ async def test_direct_openai_gpt4o_stays_on_chat_completions() -> None: mock_responses = AsyncMock(return_value=_fake_responses_response()) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -741,8 +741,8 @@ async def test_openrouter_gpt5_stays_on_chat_completions() -> None: mock_responses = AsyncMock(return_value=_fake_responses_response()) spec = find_by_name("openrouter") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -767,8 +767,8 @@ async def test_direct_openai_streaming_gpt5_uses_responses_api() -> None: mock_responses = AsyncMock(return_value=_fake_responses_stream("hi")) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -794,8 +794,8 @@ async def test_direct_openai_responses_404_falls_back_to_chat_completions() -> N mock_responses = AsyncMock(side_effect=_FakeResponsesError(404, "Responses endpoint not supported")) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -820,8 +820,8 @@ async def test_direct_openai_open_circuit_skips_responses_api() -> None: mock_responses = AsyncMock(return_value=_fake_responses_response("from responses")) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -851,8 +851,8 @@ async def test_direct_openai_stream_responses_unsupported_param_falls_back() -> ) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -877,8 +877,8 @@ async def test_direct_openai_responses_rate_limit_does_not_fallback() -> None: mock_responses = AsyncMock(side_effect=_FakeResponsesError(429, "rate limit")) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_chat client_instance.responses.create = mock_responses @@ -1232,8 +1232,8 @@ async def test_openai_compat_stream_watchdog_returns_error_on_stall(monkeypatch) mock_create = AsyncMock(return_value=_StalledStream()) spec = find_by_name("openai") - with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: - client_instance = MockClient.return_value + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class: + client_instance = mock_client_class.return_value client_instance.chat.completions.create = mock_create provider = OpenAICompatProvider( diff --git a/tests/providers/test_responses_circuit_breaker.py b/tests/providers/test_responses_circuit_breaker.py index ae6eb93a..2c9782e1 100644 --- a/tests/providers/test_responses_circuit_breaker.py +++ b/tests/providers/test_responses_circuit_breaker.py @@ -5,9 +5,9 @@ import time import pytest from nanobot.providers.openai_compat_provider import ( - OpenAICompatProvider, _RESPONSES_FAILURE_THRESHOLD, _RESPONSES_PROBE_INTERVAL_S, + OpenAICompatProvider, ) diff --git a/tests/test_package_version.py b/tests/test_package_version.py index 4780757d..467f9d81 100644 --- a/tests/test_package_version.py +++ b/tests/test_package_version.py @@ -3,9 +3,8 @@ from __future__ import annotations import subprocess import sys import textwrap -from pathlib import Path - import tomllib +from pathlib import Path def test_source_checkout_import_uses_pyproject_version_without_metadata() -> None: diff --git a/tests/tools/test_edit_enhancements.py b/tests/tools/test_edit_enhancements.py index 7202fc37..07572df9 100644 --- a/tests/tools/test_edit_enhancements.py +++ b/tests/tools/test_edit_enhancements.py @@ -3,9 +3,8 @@ notebook JSON editing, and create-file semantics.""" import pytest -from nanobot.agent.tools.filesystem import EditFileTool, ReadFileTool, WriteFileTool from nanobot.agent.tools import file_state - +from nanobot.agent.tools.filesystem import EditFileTool, ReadFileTool # --------------------------------------------------------------------------- # Helpers diff --git a/tests/tools/test_exec_security.py b/tests/tools/test_exec_security.py index 72263600..73a58ce2 100644 --- a/tests/tools/test_exec_security.py +++ b/tests/tools/test_exec_security.py @@ -9,7 +9,11 @@ from unittest.mock import patch import pytest from nanobot.agent.tools.shell import ExecTool -from nanobot.security.workspace_access import bind_workspace_scope, build_workspace_scope, reset_workspace_scope +from nanobot.security.workspace_access import ( + bind_workspace_scope, + build_workspace_scope, + reset_workspace_scope, +) def _fake_resolve_private(hostname, port, family=0, type_=0): diff --git a/tests/tools/test_exec_session_tools.py b/tests/tools/test_exec_session_tools.py index f1fc3eed..b62dbbe5 100644 --- a/tests/tools/test_exec_session_tools.py +++ b/tests/tools/test_exec_session_tools.py @@ -22,6 +22,30 @@ def _python_command(code: str) -> str: return f"{shlex.quote(sys.executable)} -u -c {shlex.quote(code)}" +def _waiting_shell_command(initial: str, *, delayed: str | None = None) -> str: + """Print deterministic output, then wait in the shell itself for stdin. + + Long-lived Python children keep inherited pipes open after their parent + shell is terminated on Windows. These tests exercise exec-session control, + not process-tree semantics, so keep the waiter in the managed shell. + """ + if sys.platform == "win32": + def quote(value: str) -> str: + return "'" + value.replace("'", "''") + "'" + + parts = [f"Write-Output {quote(initial)}"] + if delayed is not None: + parts.extend(("Start-Sleep -Milliseconds 100", f"Write-Output {quote(delayed)}")) + parts.append("$null = [Console]::In.ReadLine()") + return "; ".join(parts) + + parts = [f"printf '%s\\n' {shlex.quote(initial)}"] + if delayed is not None: + parts.extend(("sleep 0.1", f"printf '%s\\n' {shlex.quote(delayed)}")) + parts.append("IFS= read -r _") + return "; ".join(parts) + + def _session_id(output: str) -> str: match = re.search(r"session_id:\s*([0-9a-f]+)", output) assert match, output @@ -204,16 +228,14 @@ def test_write_stdin_can_terminate_session(tmp_path): manager = ExecSessionManager() exec_tool = ExecTool(working_dir=str(tmp_path), timeout=30, session_manager=manager) stdin_tool = WriteStdinTool(manager=manager) - command = _python_command( - "import time; print('ready', flush=True); time.sleep(30)" - ) + command = _waiting_shell_command("ready") initial = await exec_tool.execute(command=command, yield_time_ms=100) sid = _session_id(initial) waited = await stdin_tool.execute( session_id=sid, wait_for="ready", - wait_timeout_ms=3000, + wait_timeout_ms=1000, yield_time_ms=0, ) result = await stdin_tool.execute( @@ -234,9 +256,7 @@ def test_write_stdin_accepts_max_output_tokens_alias(tmp_path): manager = ExecSessionManager() exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager) stdin_tool = WriteStdinTool(manager=manager) - command = _python_command( - "import time; print('A' * 2000, flush=True); time.sleep(5)" - ) + command = _waiting_shell_command("A" * 2000) initial = await exec_tool.execute(command=command, yield_time_ms=0) sid = _session_id(initial) @@ -261,12 +281,12 @@ def test_write_stdin_preserves_completed_session_output_until_polled(tmp_path): stdin_tool = WriteStdinTool(manager=manager) command = _python_command( "import time; print('ready', flush=True); " - "time.sleep(1.0); print('done', flush=True)" + "time.sleep(0.1); print('done', flush=True)" ) - initial = await exec_tool.execute(command=command, yield_time_ms=300) + initial = await exec_tool.execute(command=command, yield_time_ms=50) sid = _session_id(initial) - await asyncio.sleep(1.2) + await asyncio.wait_for(manager._sessions[sid].process.wait(), timeout=2) final = await stdin_tool.execute(session_id=sid, chars="", yield_time_ms=0) return initial, final @@ -282,17 +302,14 @@ def test_write_stdin_can_wait_for_expected_output(tmp_path): manager = ExecSessionManager() exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager) stdin_tool = WriteStdinTool(manager=manager) - command = _python_command( - "import time; print('booting', flush=True); " - "time.sleep(0.4); print('ready', flush=True); time.sleep(5)" - ) + command = _waiting_shell_command("booting", delayed="ready") initial = await exec_tool.execute(command=command, yield_time_ms=100) sid = _session_id(initial) waited = await stdin_tool.execute( session_id=sid, wait_for="ready", - wait_timeout_ms=3000, + wait_timeout_ms=1000, yield_time_ms=0, ) cleanup = await stdin_tool.execute(session_id=sid, terminate=True, yield_time_ms=0) @@ -312,9 +329,7 @@ def test_write_stdin_wait_for_reports_timeout_without_killing_session(tmp_path): manager = ExecSessionManager() exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager) stdin_tool = WriteStdinTool(manager=manager) - command = _python_command( - "import time; print('booting', flush=True); time.sleep(5)" - ) + command = _waiting_shell_command("booting") initial = await exec_tool.execute(command=command, yield_time_ms=100) sid = _session_id(initial) @@ -365,9 +380,7 @@ def test_list_exec_sessions_reports_running_commands(tmp_path): exec_tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager) list_tool = ListExecSessionsTool(manager=manager) stdin_tool = WriteStdinTool(manager=manager) - command = _python_command( - "import time; print('ready', flush=True); time.sleep(5)" - ) + command = _waiting_shell_command("ready") initial = await exec_tool.execute(command=command, yield_time_ms=500) sid = _session_id(initial) diff --git a/tests/tools/test_mcp_probe.py b/tests/tools/test_mcp_probe.py index 20e30719..c64a1172 100644 --- a/tests/tools/test_mcp_probe.py +++ b/tests/tools/test_mcp_probe.py @@ -49,9 +49,23 @@ async def test_probe_returns_false_for_closed_port(): @pytest.mark.asyncio -async def test_probe_uses_default_port_for_http(): - """When no port in URL, should default to 80 (will fail -> False).""" +async def test_probe_uses_default_port_for_http(monkeypatch: pytest.MonkeyPatch): + """When no port is present, probe the validated address on port 80.""" + attempts: list[tuple[str, int]] = [] + + monkeypatch.setattr( + "nanobot.agent.tools.mcp.resolve_url_target", + lambda _url: (True, "", ("93.184.216.34",)), + ) + + async def _open_connection(host: str, port: int): + attempts.append((host, port)) + raise ConnectionRefusedError + + monkeypatch.setattr("nanobot.agent.tools.mcp.asyncio.open_connection", _open_connection) + assert await _probe_http_url("http://unreachable-host.test/mcp") is False + assert attempts == [("93.184.216.34", 80)] @pytest.mark.asyncio diff --git a/tests/tools/test_mcp_tool.py b/tests/tools/test_mcp_tool.py index e81a9d04..ccc2abeb 100644 --- a/tests/tools/test_mcp_tool.py +++ b/tests/tools/test_mcp_tool.py @@ -811,6 +811,17 @@ async def test_connect_mcp_servers_env_proxy_adds_proxy_mounts_and_keeps_pinned_ monkeypatch.setenv("NO_PROXY", "localhost,127.0.0.1,::1") monkeypatch.setattr(mcp_mod, "validate_url_target", _validate) monkeypatch.setattr(mcp_mod, "_probe_http_url", _reachable) + monkeypatch.setattr( + mcp_mod, + "PinnedDNSAsyncTransport", + lambda: httpx.MockTransport(lambda request: httpx.Response(200, request=request)), + ) + monkeypatch.setattr( + "nanobot.security.network.httpx.AsyncHTTPTransport", + lambda **_kwargs: httpx.MockTransport( + lambda request: httpx.Response(200, request=request) + ), + ) monkeypatch.setattr(mcp_mod.httpx, "AsyncClient", FakeAsyncClient) monkeypatch.setattr(sys.modules["mcp.client.sse"], "sse_client", _capturing_sse_client) monkeypatch.setattr( @@ -832,6 +843,17 @@ async def test_connect_mcp_servers_env_proxy_adds_proxy_mounts_and_keeps_pinned_ def test_mcp_http_clients_no_proxy_env_keeps_pinned_direct_route(monkeypatch): monkeypatch.setenv("HTTPS_PROXY", "http://proxy.example:8080") monkeypatch.setenv("NO_PROXY", "mcp.example.com") + monkeypatch.setattr( + mcp_mod, + "PinnedDNSAsyncTransport", + lambda: httpx.MockTransport(lambda request: httpx.Response(200, request=request)), + ) + monkeypatch.setattr( + "nanobot.security.network.httpx.AsyncHTTPTransport", + lambda **_kwargs: httpx.MockTransport( + lambda request: httpx.Response(200, request=request) + ), + ) kwargs = mcp_mod._pinned_transport_kwargs() @@ -989,6 +1011,11 @@ async def test_connect_mcp_servers_streamable_http_uses_finite_timeout( monkeypatch.setattr(mcp_mod, "validate_url_target", _validate) monkeypatch.setattr(mcp_mod, "_probe_http_url", _reachable) + monkeypatch.setattr( + mcp_mod, + "PinnedDNSAsyncTransport", + lambda: httpx.MockTransport(lambda request: httpx.Response(200, request=request)), + ) monkeypatch.setattr( sys.modules["mcp.client.streamable_http"], "streamable_http_client", diff --git a/tests/tools/test_web_fetch_security.py b/tests/tools/test_web_fetch_security.py index ab9bbe9f..560522b7 100644 --- a/tests/tools/test_web_fetch_security.py +++ b/tests/tools/test_web_fetch_security.py @@ -78,6 +78,11 @@ def _patch_web_fetch_fake_client(monkeypatch: pytest.MonkeyPatch) -> list[dict]: return FakeJinaResponse() monkeypatch.setattr(web_module.httpx, "AsyncClient", FakeClient) + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) + monkeypatch.setattr( + "nanobot.security.network.httpx.AsyncHTTPTransport", + lambda **_kwargs: object(), + ) return client_kwargs @@ -121,27 +126,12 @@ async def test_web_fetch_blocks_localhost_even_in_full_workspace_scope(tmp_path) @pytest.mark.asyncio -async def test_web_fetch_result_contains_untrusted_flag(): +async def test_web_fetch_result_contains_untrusted_flag(monkeypatch: pytest.MonkeyPatch): """When fetch succeeds, result JSON must include untrusted=True and the banner.""" tool = WebFetchTool() + _patch_web_fetch_fake_client(monkeypatch) - fake_html = "
Hello world
" - - - class FakeResponse: - status_code = 200 - url = "https://example.com/page" - text = fake_html - headers = {"content-type": "text/html"} - is_redirect = False - def raise_for_status(self): pass - def json(self): return {} - - async def _fake_get(self, url, **kwargs): - return FakeResponse() - - with patch("nanobot.security.network.socket.getaddrinfo", _fake_resolve_public), \ - patch("httpx.AsyncClient.get", _fake_get): + with patch("nanobot.security.network.socket.getaddrinfo", _fake_resolve_public): result = await tool.execute(url="https://example.com/page") data = json.loads(result) @@ -237,6 +227,11 @@ async def test_web_fetch_env_proxy_adds_proxy_mounts_and_keeps_pinned_transport( def test_web_fetch_no_proxy_env_keeps_pinned_direct_route(monkeypatch): monkeypatch.setenv("HTTPS_PROXY", "http://proxy.example:8080") monkeypatch.setenv("NO_PROXY", "example.com") + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) + monkeypatch.setattr( + "nanobot.security.network.httpx.AsyncHTTPTransport", + lambda **_kwargs: object(), + ) kwargs = web_module._fetch_client_kwargs(None, 15.0) @@ -265,6 +260,16 @@ async def test_web_fetch_does_not_fallback_after_pinned_dns_rebind_rejection(mon monkeypatch.setattr(tool, "_fetch_jina", _unexpected_jina) monkeypatch.setattr(tool, "_fetch_readability", _unexpected_readability) + class FailTransport(httpx.AsyncBaseTransport): + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + raise AssertionError("rebound target must be rejected before transport") + + monkeypatch.setattr( + web_module, + "_pinned_dns_transport", + lambda: PinnedDNSAsyncTransport(inner=FailTransport()), + ) + with patch("nanobot.security.network.socket.getaddrinfo", _rebinding_resolver): result = await tool.execute(url="http://evil.example/page") @@ -330,6 +335,7 @@ async def test_web_fetch_can_skip_jina_and_use_custom_user_agent(monkeypatch): monkeypatch.setattr(tool, "_fetch_jina", _fail_jina) monkeypatch.setattr(tool, "_extract_readable_html", lambda html, mode: "Hello world") monkeypatch.setattr("nanobot.agent.tools.web.httpx.AsyncClient", FakeClient) + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) with patch("nanobot.security.network.socket.getaddrinfo", _fake_resolve_public): result = await tool.execute(url="https://example.com/page") @@ -373,6 +379,7 @@ async def test_web_fetch_falls_back_when_readability_dependency_is_missing(monke monkeypatch.setattr(tool, "_extract_readable_html", _missing_readability) monkeypatch.setattr("nanobot.agent.tools.web.httpx.AsyncClient", FakeClient) + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) with patch("nanobot.security.network.socket.getaddrinfo", _fake_resolve_public): result = await tool._fetch_readability("https://example.com/page", "markdown", 5000) @@ -430,6 +437,7 @@ async def test_web_fetch_blocks_private_redirect_before_readability_request(monk return FakeRedirectResponse() monkeypatch.setattr(web_module.httpx, "AsyncClient", FakeClient) + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) def resolve_public_start_only(hostname, port, family=0, type_=0): if hostname == "attacker.example": @@ -475,6 +483,7 @@ async def test_web_fetch_blocks_private_redirect_before_returning_image(monkeypa super().__init__(*args, transport=transport, **kwargs) monkeypatch.setattr("nanobot.agent.tools.web.httpx.AsyncClient", TransportAsyncClient) + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) def resolve_public_start_only(hostname, port, family=0, type_=0): if hostname == "example.com": @@ -515,6 +524,7 @@ async def test_web_fetch_does_not_request_private_redirect_target(monkeypatch): super().__init__(*args, **kwargs) monkeypatch.setattr(web_module.httpx, "AsyncClient", TransportAsyncClient) + monkeypatch.setattr(web_module, "_pinned_dns_transport", lambda: object()) def resolve_public_start_only(hostname, port, family=0, type_=0): if hostname == "attacker.example": diff --git a/tests/tools/test_web_fetch_url_sanitization.py b/tests/tools/test_web_fetch_url_sanitization.py index 8a24338c..bba12194 100644 --- a/tests/tools/test_web_fetch_url_sanitization.py +++ b/tests/tools/test_web_fetch_url_sanitization.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +from contextlib import contextmanager from unittest.mock import patch import pytest @@ -41,9 +42,14 @@ class FakeClient: return FakeResponse() -def _patch_env(): - return patch("nanobot.security.network.socket.getaddrinfo", _fake_resolve_public), \ - patch("nanobot.agent.tools.web.httpx.AsyncClient", FakeClient) +@contextmanager +def _patched_web_fetch(): + with ( + patch("nanobot.security.network.socket.getaddrinfo", _fake_resolve_public), + patch("nanobot.agent.tools.web.httpx.AsyncClient", FakeClient), + patch("nanobot.agent.tools.web._pinned_dns_transport", lambda: object()), + ): + yield # --- urlparse / _validate_url level tests --- @@ -77,7 +83,7 @@ def test_backtick_url_produces_empty_scheme_in_urlparse(): @pytest.mark.asyncio async def test_execute_strips_backticks_and_succeeds(): tool = WebFetchTool() - with _patch_env()[0], _patch_env()[1]: + with _patched_web_fetch(): result = await tool.execute(url="`https://example.com/page`") data = json.loads(result) assert "error" not in data, f"unexpected error: {data}" @@ -86,7 +92,7 @@ async def test_execute_strips_backticks_and_succeeds(): @pytest.mark.asyncio async def test_execute_strips_double_quotes_and_succeeds(): tool = WebFetchTool() - with _patch_env()[0], _patch_env()[1]: + with _patched_web_fetch(): result = await tool.execute(url='"https://example.com/page"') data = json.loads(result) assert "error" not in data, f"unexpected error: {data}" @@ -95,7 +101,7 @@ async def test_execute_strips_double_quotes_and_succeeds(): @pytest.mark.asyncio async def test_execute_strips_single_quotes_and_succeeds(): tool = WebFetchTool() - with _patch_env()[0], _patch_env()[1]: + with _patched_web_fetch(): result = await tool.execute(url="'https://example.com/page'") data = json.loads(result) assert "error" not in data, f"unexpected error: {data}" @@ -104,7 +110,7 @@ async def test_execute_strips_single_quotes_and_succeeds(): @pytest.mark.asyncio async def test_execute_strips_space_and_backticks(): tool = WebFetchTool() - with _patch_env()[0], _patch_env()[1]: + with _patched_web_fetch(): result = await tool.execute(url=" `https://example.com/page` ") data = json.loads(result) assert "error" not in data, f"unexpected error: {data}" @@ -113,7 +119,7 @@ async def test_execute_strips_space_and_backticks(): @pytest.mark.asyncio async def test_execute_strips_mixed_markdown_and_quotes(): tool = WebFetchTool() - with _patch_env()[0], _patch_env()[1]: + with _patched_web_fetch(): result = await tool.execute(url='"`https://example.com/page`"') data = json.loads(result) assert "error" not in data, f"unexpected error: {data}" @@ -122,7 +128,7 @@ async def test_execute_strips_mixed_markdown_and_quotes(): @pytest.mark.asyncio async def test_execute_keeps_case_insensitive_http_scheme(): tool = WebFetchTool() - with _patch_env()[0], _patch_env()[1]: + with _patched_web_fetch(): result = await tool.execute(url="HTTPS://example.com/page") data = json.loads(result) assert "error" not in data, f"unexpected error: {data}" diff --git a/tests/utils/test_abbreviate_path.py b/tests/utils/test_abbreviate_path.py index 573ca0a9..6b794f3c 100644 --- a/tests/utils/test_abbreviate_path.py +++ b/tests/utils/test_abbreviate_path.py @@ -1,6 +1,7 @@ """Tests for abbreviate_path utility.""" import os + from nanobot.utils.path import abbreviate_path @@ -9,7 +10,6 @@ class TestAbbreviatePathShort: assert abbreviate_path("/home/user/file.py") == "/home/user/file.py" def test_exact_max_len_unchanged(self): - path = "/a/b/c" # 7 chars assert abbreviate_path("/a/b/c", max_len=7) == "/a/b/c" def test_basename_only(self): diff --git a/tests/utils/test_searchusage.py b/tests/utils/test_searchusage.py index 205ccd91..10c3a58f 100644 --- a/tests/utils/test_searchusage.py +++ b/tests/utils/test_searchusage.py @@ -2,16 +2,16 @@ from __future__ import annotations -import pytest from unittest.mock import AsyncMock, MagicMock, patch +import pytest + +from nanobot.utils.helpers import build_status_content from nanobot.utils.searchusage import ( SearchUsageInfo, _parse_tavily_usage, fetch_search_usage, ) -from nanobot.utils.helpers import build_status_content - # --------------------------------------------------------------------------- # SearchUsageInfo.format() tests diff --git a/webui/src/tests/diff-syntax-highlight.integration.test.tsx b/webui/src/tests/diff-syntax-highlight.integration.test.tsx index 3a721650..bf4efcc8 100644 --- a/webui/src/tests/diff-syntax-highlight.integration.test.tsx +++ b/webui/src/tests/diff-syntax-highlight.integration.test.tsx @@ -40,7 +40,13 @@ describe("DiffSyntaxHighlight with Prism", () => { , ); - const highlighted = await screen.findByTestId("syntax-highlighted-diff-hunk"); + // Full-suite workers can keep the first Prism grammar import busy for more + // than Testing Library's one-second default, especially on Windows. + const highlighted = await screen.findByTestId( + "syntax-highlighted-diff-hunk", + {}, + { timeout: 10_000 }, + ); await waitFor( () => { const tokens = highlighted.querySelectorAll