diff --git a/docs/source/guides/catalog-discovery.md b/docs/source/guides/catalog-discovery.md index f5b4e34431..d5a8862231 100644 --- a/docs/source/guides/catalog-discovery.md +++ b/docs/source/guides/catalog-discovery.md @@ -165,6 +165,9 @@ JSON Schema validation is necessary but is not the whole profile contract. The packaged schemas enforce object shape, required fields, relative-path safety, supported literals, conditional artifact/license-evidence presence, and exactly one orchestration interface. Other rules require semantic validation: +The relative-path profile permits printable UTF-8 (including spaces), but +rejects C0, DEL, and C1 control characters so untrusted metadata cannot forge +CLI or log lines. | Subject | Additional rule | |---------|-----------------| diff --git a/src/openenv/auto/_discovery.py b/src/openenv/auto/_discovery.py index 7b53fc6bcc..1f7054ded5 100644 --- a/src/openenv/auto/_discovery.py +++ b/src/openenv/auto/_discovery.py @@ -23,6 +23,7 @@ import os import re import stat +import tempfile from dataclasses import asdict, dataclass from pathlib import Path from typing import Any, Type @@ -345,9 +346,28 @@ def _default_cache_file() -> Path: shared, world-writable temporary directory. A fixed path under the shared temp dir lets another local user pre-create the cache file and redirect discovery to attacker-controlled import paths (`import_module` on a cached `client_module_path`). + Per the XDG Base Directory specification, relative `XDG_CACHE_HOME` values are + ignored so an untrusted working tree cannot supply a victim-owned cache file. + Relative `HOME` values are likewise rejected: `Path.home()` must not be + resolved against the current working directory. """ base = os.environ.get("XDG_CACHE_HOME") - root = Path(base) if base else Path.home() / ".cache" + if base and Path(base).is_absolute(): + root = Path(base) + else: + home = Path.home() + if home.is_absolute(): + root = home / ".cache" + else: + # Keep the fallback absolute and uid-scoped so a shared temp root + # cannot be turned into a fixed, cross-user planting target. + uid = os.getuid() if hasattr(os, "getuid") else os.getpid() + root = Path(tempfile.gettempdir()) / f"openenv-{uid}-cache" + if not root.is_absolute(): + raise RuntimeError( + "Cannot resolve an absolute discovery cache directory when " + "XDG_CACHE_HOME and HOME are both missing or relative" + ) return root / "openenv" / "discovery_cache.json" diff --git a/src/openenv/core/env_client.py b/src/openenv/core/env_client.py index f15df6cb2f..164f99483b 100644 --- a/src/openenv/core/env_client.py +++ b/src/openenv/core/env_client.py @@ -255,6 +255,15 @@ async def _best_effort_close(ws: ClientConnection) -> None: pass # Best effort +async def _best_effort_disconnect(ws: ClientConnection) -> None: + """Notify the server, then close the socket without propagating failures.""" + try: + await ws.send(json.dumps({"type": "close"})) + except (Exception, asyncio.CancelledError): + pass # Best effort + await _best_effort_close(ws) + + class EnvClient(ABC, Generic[ActT, ObsT, StateT]): """ Async environment client for persistent sessions. @@ -538,6 +547,13 @@ async def _connect_async(self) -> "EnvClient": self._ws = None self._ws_loop = None + # A timed-out request drops its socket immediately but closes it in the + # background so the timeout itself remains prompt. Wait for that close + # before opening a replacement: the old server-side session continues + # occupying a capacity slot until the close handshake finishes, and + # many environments allow only one session. + await self._drain_pending_close_tasks() + try: self._start_provider_if_needed() except Exception: @@ -573,24 +589,49 @@ async def _connect_async(self) -> "EnvClient": def disconnect(self) -> Any: return self._dispatch(self._disconnect_async) + def _schedule_socket_close( + self, ws: ClientConnection, *, notify_server: bool = False + ) -> asyncio.Task[None]: + """Schedule and track a socket close on the current event loop.""" + close = _best_effort_disconnect(ws) if notify_server else _best_effort_close(ws) + close_task = asyncio.create_task(close) + self._pending_close_tasks.add(close_task) + close_task.add_done_callback(self._pending_close_tasks.discard) + return close_task + async def _disconnect_async(self) -> None: """Close the WebSocket connection.""" if self._ws is not None: ws = self._ws ws_loop = self._ws_loop - same_loop = ws_loop is asyncio.get_running_loop() - try: - if same_loop: - await ws.send(json.dumps({"type": "close"})) - except Exception: - pass # Best effort - try: - if same_loop: - await ws.close() - except Exception: - pass + # Detach first so cancellation during the close handshake cannot + # leave a stale socket cached for a later operation. self._ws = None self._ws_loop = None + same_loop = ws_loop is asyncio.get_running_loop() + if same_loop: + # Track the detached socket before awaiting anything. If this + # caller is cancelled, the shielded task keeps closing and a + # later reconnect drains it before opening a replacement. + close_task = self._schedule_socket_close(ws, notify_server=True) + await asyncio.shield(close_task) + + async def _drain_pending_close_tasks(self) -> None: + """Wait for background socket closes owned by the current event loop. + + Shielding keeps cancellation of the caller from cancelling the close + tasks themselves. This matters both before reconnecting, when the old + server session must release its capacity slot, and during explicit + client shutdown. + """ + loop = asyncio.get_running_loop() + tasks = [ + task + for task in tuple(self._pending_close_tasks) + if not task.done() and task.get_loop() is loop + ] + if tasks: + await asyncio.shield(asyncio.gather(*tasks, return_exceptions=True)) async def _ensure_connected(self) -> None: """Ensure WebSocket connection is established on the current loop. @@ -641,9 +682,7 @@ async def _receive(self) -> Dict[str, Any]: # would actually block for up to 10s before its deadline was # honored. Scheduling it lets the exception propagate # immediately while the close still happens in the background. - close_task = asyncio.ensure_future(_best_effort_close(ws)) - self._pending_close_tasks.add(close_task) - close_task.add_done_callback(self._pending_close_tasks.discard) + self._schedule_socket_close(ws) raise return json.loads(raw) @@ -957,34 +996,36 @@ async def _close_async(self) -> None: If this client was created via from_docker_image() or from_env(), this will also stop and remove the associated container/process. """ - for child in list(self._child_clients): - with suppress(Exception): - await child.close() - self._child_clients.clear() - try: - # Wait out any backgrounded closes from a dropped socket (see - # `_receive()` / `_best_effort_close`) so a real close() call still - # sees the handshake through. SyncEnvClient.close() waits for - # `_close_async()` before stopping its loop, so the relevant risk - # is async-context cancellation of close itself — not `_stop_loop()`. - # Keep this gather inside the provider-teardown try/finally so a - # cancelled close cannot skip container/process cleanup. - if self._pending_close_tasks: - await asyncio.gather(*self._pending_close_tasks, return_exceptions=True) - await self._disconnect_async() + for child in list(self._child_clients): + with suppress(Exception): + await child.close() finally: + # Parent teardown must run even when a child close is cancelled. + self._child_clients.clear() try: - if self._provider is not None: - # Handle both ContainerProvider and RuntimeProvider - if hasattr(self._provider, "stop_container"): - self._provider.stop_container() - elif hasattr(self._provider, "stop"): - self._provider.stop() + try: + # A real close waits out backgrounded closes, but shield them + # from cancellation so their socket handshakes aren't + # abandoned midway. + await self._drain_pending_close_tasks() + finally: + # Run even when pending-close draining is cancelled. A client + # may already have reconnected, and that current socket must + # not remain cached or open during teardown. + await self._disconnect_async() finally: - if self._start_provider_on_connect: - self._base_url = None - self._ws_url = None + try: + if self._provider is not None: + # Handle both ContainerProvider and RuntimeProvider + if hasattr(self._provider, "stop_container"): + self._provider.stop_container() + elif hasattr(self._provider, "stop"): + self._provider.stop() + finally: + if self._start_provider_on_connect: + self._base_url = None + self._ws_url = None def _stop_provider_best_effort(self) -> None: """Stop the underlying provider directly, ignoring any errors. diff --git a/src/openenv/discovery/models.py b/src/openenv/discovery/models.py index 39d47d6489..4f25ff85e5 100644 --- a/src/openenv/discovery/models.py +++ b/src/openenv/discovery/models.py @@ -53,7 +53,10 @@ def relative_path(value: str) -> str: return value if ( not value - or "\x00" in value + or any( + ord(character) < 0x20 or 0x7F <= ord(character) <= 0x9F + for character in value + ) or "\\" in value or PurePosixPath(value).is_absolute() or any(part in ("", ".", "..") for part in value.split("/")) @@ -64,8 +67,10 @@ def relative_path(value: str) -> str: # Positive components exclude "." and ".." without lookaround, which some # JSON Schema regex engines do not support. +_CONTROL_CHARACTER_PATTERN = r"[\x00-\x1f\x7f-\x9f]" _PATH_COMPONENT_PATTERN = ( - r"(?:[^./\\\x00]|\.[^./\\\x00]|\.\.[^./\\\x00]|\.\.\.)[^/\\\x00]*" + r"(?:[^./\\\x00-\x1f\x7f-\x9f]|\.[^./\\\x00-\x1f\x7f-\x9f]" + r"|\.\.[^./\\\x00-\x1f\x7f-\x9f]|\.\.\.)[^/\\\x00-\x1f\x7f-\x9f]*" ) RelativePath = Annotated[ NonEmpty, @@ -79,6 +84,7 @@ def relative_path(value: str) -> str: rf"(?:/{_PATH_COMPONENT_PATTERN})*)$" ), }, + {"not": {"pattern": _CONTROL_CHARACTER_PATTERN}}, ] } ), diff --git a/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json b/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json index 3b4491bd06..93114d0d77 100644 --- a/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json @@ -55,7 +55,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -405,7 +410,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -472,7 +482,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -510,7 +525,12 @@ "items": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -524,7 +544,12 @@ "root": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, diff --git a/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json b/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json index 91576021ae..5c81448374 100644 --- a/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json @@ -25,7 +25,12 @@ "source": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -104,7 +109,12 @@ { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, diff --git a/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json b/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json index 0f49575919..b0102491c4 100644 --- a/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json @@ -49,7 +49,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -97,7 +102,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, diff --git a/tests/discovery/test_catalog_contract.py b/tests/discovery/test_catalog_contract.py index f975775264..d5ff0289c1 100644 --- a/tests/discovery/test_catalog_contract.py +++ b/tests/discovery/test_catalog_contract.py @@ -15,6 +15,9 @@ REVISION = "a" * 40 +ASCII_CONTROL_PATHS = [ + f"envs/control-{chr(codepoint)}" for codepoint in [*range(0x20), *range(0x7F, 0xA0)] +] @pytest.fixture @@ -92,7 +95,13 @@ def test_tool_declaration_cannot_borrow_another_revision(card): "envs//echo", "envs/echo/", "envs/./echo", - "envs/\x00echo", + *ASCII_CONTROL_PATHS, + "envs/trailing\n", + ".\n", + "..\n", + "envs/\n", + "envs/.\n", + "envs/..\n", ], ) def test_environment_locator_is_a_safe_repository_relative_path( @@ -115,12 +124,6 @@ def test_environment_locator_is_a_safe_repository_relative_path( "envs/...", "envs/a..b", "envs/with spaces", - "envs/trailing\n", - ".\n", - "..\n", - "envs/\n", - "envs/.\n", - "envs/..\n", ], ) def test_schema_and_model_preserve_valid_relative_locators(card, path, card_schema): diff --git a/tests/envs/test_discovery.py b/tests/envs/test_discovery.py index ff1d3890bd..a54a0703c2 100644 --- a/tests/envs/test_discovery.py +++ b/tests/envs/test_discovery.py @@ -16,6 +16,7 @@ import os import stat import tempfile +from pathlib import Path from unittest.mock import Mock, patch import openenv.auto._discovery as _discovery_module @@ -344,6 +345,45 @@ def test_cache_file_is_per_user_not_shared_tmp(self): assert tempfile.gettempdir() not in str(path) assert path.parent.name == "openenv" + def test_relative_xdg_cache_home_cannot_redirect_into_working_tree( + self, tmp_path, monkeypatch + ): + """A relative XDG path must not trust a cache planted in the checkout.""" + checkout = tmp_path / "untrusted-checkout" + planted = checkout / "cache" / "openenv" / "discovery_cache.json" + planted.parent.mkdir(parents=True) + planted.write_text("{}") + + monkeypatch.chdir(checkout) + monkeypatch.setenv("XDG_CACHE_HOME", "cache") + + path = _default_cache_file() + + assert path == Path.home() / ".cache" / "openenv" / "discovery_cache.json" + assert path.is_absolute() + assert path.resolve() != planted.resolve() + + def test_relative_home_cannot_redirect_into_working_tree( + self, tmp_path, monkeypatch + ): + """A relative HOME must not select a cache planted in the checkout.""" + checkout = tmp_path / "untrusted-checkout" + planted = checkout / "cache" / ".cache" / "openenv" / "discovery_cache.json" + planted.parent.mkdir(parents=True) + planted.write_text("{}") + + monkeypatch.chdir(checkout) + monkeypatch.delenv("XDG_CACHE_HOME", raising=False) + monkeypatch.setenv("HOME", "cache") + + path = _default_cache_file() + + assert path.is_absolute() + assert path.resolve() != planted.resolve() + assert "openenv" in path.parts + uid = os.getuid() if hasattr(os, "getuid") else os.getpid() + assert f"openenv-{uid}-cache" in path.parts + def test_world_writable_cache_is_not_trusted(self, tmp_path): f = tmp_path / "cache.json" f.write_text("{}") diff --git a/tests/test_core/test_generic_client.py b/tests/test_core/test_generic_client.py index 520f0b98e0..44196c5422 100644 --- a/tests/test_core/test_generic_client.py +++ b/tests/test_core/test_generic_client.py @@ -17,7 +17,6 @@ import asyncio import os -from contextlib import suppress from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -1663,17 +1662,118 @@ async def close(self): ) @pytest.mark.asyncio - async def test_close_async_cancelled_during_pending_gather_still_stops_provider( - self, - ): - """Cancelling `_close_async` while draining pending closes must still - tear down the provider and clear provider-owned URLs. + async def test_reconnect_waits_for_dropped_socket_to_release_capacity(self): + """A replacement connection must wait for the old session to close. - Regression: the pending-close `gather` used to run before the - try/finally that stops the provider. Cancellation of `_close_async` - itself propagates from `gather` even with `return_exceptions=True`, - which skipped provider teardown and leaked the container/process. + A timed-out socket is detached immediately and closed in the background. + Reconnecting before that handshake finishes races the server's session + accounting; with the default capacity of one, the retry is rejected. """ + close_started = asyncio.Event() + release_close = asyncio.Event() + + class SlowClose: + state = State.OPEN + + async def send(self, _message): + pass + + async def recv(self): + await asyncio.sleep(10) + + async def close(self): + close_started.set() + await release_close.wait() + self.state = State.CLOSED + + client = GenericEnvClient( + base_url="http://localhost:8000", message_timeout_s=0.01 + ) + dropped_ws = SlowClose() + client._ws = dropped_ws + client._ws_loop = asyncio.get_running_loop() + + with pytest.raises(asyncio.TimeoutError): + await client._send_and_receive({"type": "state"}) + await close_started.wait() + + replacement_ws = AsyncMock() + replacement_ws.state = State.OPEN + replacement_ws.recv.return_value = '{"type": "state", "data": {}}' + + async def fake_ws_connect(*args, **kwargs): + return replacement_ws + + with patch( + "openenv.core.env_client.ws_connect", side_effect=fake_ws_connect + ) as mock_connect: + reconnect = asyncio.create_task(client._connect_async()) + await asyncio.sleep(0) + assert not reconnect.done() + mock_connect.assert_not_called() + + release_close.set() + await reconnect + + mock_connect.assert_called_once() + assert dropped_ws.state == State.CLOSED + assert client._ws is replacement_ws + await client._close_async() + + @pytest.mark.asyncio + async def test_cancelled_disconnect_drains_current_socket_before_reconnect(self): + """A cancelled disconnect must keep tracking the detached socket.""" + close_started = asyncio.Event() + release_close = asyncio.Event() + + class SlowClose: + state = State.OPEN + + async def send(self, _message): + pass + + async def close(self): + close_started.set() + await release_close.wait() + self.state = State.CLOSED + + client = GenericEnvClient(base_url="http://localhost:8000") + current_ws = SlowClose() + client._ws = current_ws + client._ws_loop = asyncio.get_running_loop() + + disconnect = asyncio.create_task(client._disconnect_async()) + await asyncio.wait_for(close_started.wait(), timeout=1) + disconnect.cancel() + with pytest.raises(asyncio.CancelledError): + await disconnect + + replacement_ws = AsyncMock() + replacement_ws.state = State.OPEN + + async def fake_ws_connect(*args, **kwargs): + return replacement_ws + + with patch( + "openenv.core.env_client.ws_connect", side_effect=fake_ws_connect + ) as mock_connect: + reconnect = asyncio.create_task(client._connect_async()) + await asyncio.sleep(0) + reconnect_waited_for_close = not reconnect.done() + calls_before_close_finished = mock_connect.call_count + + release_close.set() + await asyncio.wait_for(reconnect, timeout=1) + + assert reconnect_waited_for_close + assert calls_before_close_finished == 0 + assert current_ws.state == State.CLOSED + assert client._ws is replacement_ws + await client._close_async() + + @pytest.mark.asyncio + async def test_cancelled_close_still_closes_current_and_pending_sockets(self): + """Cancellation while draining an old socket must not leak either one.""" class FakeRuntimeProvider: def __init__(self): @@ -1682,37 +1782,111 @@ def __init__(self): def stop(self): self.stopped = True + close_started = asyncio.Event() + release_close = asyncio.Event() + + class PendingSocket: + state = State.OPEN + + async def close(self): + close_started.set() + await release_close.wait() + self.state = State.CLOSED + + class CurrentSocket: + state = State.OPEN + + async def send(self, _message): + pass + + async def close(self): + self.state = State.CLOSED + provider = FakeRuntimeProvider() client = GenericEnvClient(provider=provider) client._base_url = "http://localhost:8000" client._ws_url = "ws://localhost:8000/ws" - hang_gate = asyncio.Event() - - async def hang_forever(): - await hang_gate.wait() - - pending = asyncio.create_task(hang_forever()) + dropped_ws = PendingSocket() + pending = asyncio.create_task(dropped_ws.close()) client._pending_close_tasks.add(pending) + pending.add_done_callback(client._pending_close_tasks.discard) + await close_started.wait() + + current_ws = CurrentSocket() + client._ws = current_ws + client._ws_loop = asyncio.get_running_loop() close_task = asyncio.create_task(client._close_async()) - await asyncio.sleep(0) # let close enter the pending-close gather + await asyncio.sleep(0) # let close enter the shielded pending-close drain assert not close_task.done() close_task.cancel() with pytest.raises(asyncio.CancelledError): await close_task - hang_gate.set() - with suppress(asyncio.CancelledError): - await pending + assert provider.stopped + assert client._base_url is None + assert client._ws_url is None + assert client._ws is None + assert current_ws.state == State.CLOSED + assert not pending.cancelled() + + release_close.set() + await pending + assert dropped_ws.state == State.CLOSED - assert provider.stopped, ( - "provider.stop() must run even when _close_async is cancelled " - "during the pending-close gather" - ) + @pytest.mark.asyncio + async def test_cancelled_child_close_still_tears_down_parent(self): + """Cancellation during child.close() must not skip parent teardown.""" + + class FakeRuntimeProvider: + def __init__(self): + self.stopped = False + + def stop(self): + self.stopped = True + + child_started = asyncio.Event() + release_child = asyncio.Event() + + class SlowChild: + async def close(self): + child_started.set() + await release_child.wait() + + class ParentSocket: + state = State.OPEN + + async def send(self, _message): + pass + + async def close(self): + self.state = State.CLOSED + + provider = FakeRuntimeProvider() + client = GenericEnvClient(provider=provider) + client._base_url = "http://localhost:8000" + client._ws_url = "ws://localhost:8000/ws" + client._child_clients.append(SlowChild()) + parent_ws = ParentSocket() + client._ws = parent_ws + client._ws_loop = asyncio.get_running_loop() + + close_task = asyncio.create_task(client._close_async()) + await child_started.wait() + close_task.cancel() + with pytest.raises(asyncio.CancelledError): + await close_task + + assert provider.stopped assert client._base_url is None assert client._ws_url is None + assert client._ws is None + assert parent_ws.state == State.CLOSED + assert client._child_clients == [] + + release_child.set() # ============================================================================