diff --git a/src/kernel/_client.py b/src/kernel/_client.py index 497cc5d5..9a040e34 100644 --- a/src/kernel/_client.py +++ b/src/kernel/_client.py @@ -41,6 +41,7 @@ strip_direct_vm_auth, rewrite_direct_vm_options, browser_routing_config_from_env, + should_retry_stale_direct_vm_auth, maybe_evict_browser_route_from_response, maybe_populate_browser_route_cache_from_response, ) @@ -353,6 +354,13 @@ def _prepare_options(self, options: Any) -> Any: def _prepare_request(self, request: httpx.Request) -> None: strip_direct_vm_auth(request, cache=self.browser_route_cache) + @override + def _should_retry(self, response: httpx.Response) -> bool: + if should_retry_stale_direct_vm_auth(response): + maybe_evict_browser_route_from_response(response, cache=self.browser_route_cache) + return True + return super()._should_retry(response) + @override def _process_response( self, @@ -722,6 +730,13 @@ async def _prepare_options(self, options: Any) -> Any: async def _prepare_request(self, request: httpx.Request) -> None: strip_direct_vm_auth(request, cache=self.browser_route_cache) + @override + def _should_retry(self, response: httpx.Response) -> bool: + if should_retry_stale_direct_vm_auth(response): + maybe_evict_browser_route_from_response(response, cache=self.browser_route_cache) + return True + return super()._should_retry(response) + @override async def _process_response( self, diff --git a/src/kernel/lib/browser_routing/routing.py b/src/kernel/lib/browser_routing/routing.py index 7f28d726..1db91027 100644 --- a/src/kernel/lib/browser_routing/routing.py +++ b/src/kernel/lib/browser_routing/routing.py @@ -44,7 +44,7 @@ def browser_routing_config_from_env() -> BrowserRoutingConfig: # Path prefixes eligible for direct-to-VM routing. "telemetry/stream" is # the live SSE endpoint (VM); "telemetry/events" is a historical read # served by the control plane (S2) and must NOT be here. - return BrowserRoutingConfig(subresources=("curl", "telemetry/stream")) + return BrowserRoutingConfig(subresources=("curl", "telemetry/stream", "computer", "playwright")) if raw.strip() == "": return BrowserRoutingConfig() @@ -110,14 +110,18 @@ def maybe_populate_browser_route_cache_from_response(response: httpx.Response, * def maybe_evict_browser_route_from_response(response: httpx.Response, *, cache: BrowserRouteCache) -> None: - if not response.is_success: + if response.is_success: + session_id = _session_id_to_evict_from_response(response) + if session_id: + cache.delete(session_id) return - session_id = _session_id_to_evict_from_response(response) - if not session_id: + if not is_stale_direct_vm_auth_response(response): return - cache.delete(session_id) + session_id = _session_id_from_direct_vm_response(response, cache=cache) + if session_id: + cache.delete(session_id) def populate_browser_route_cache_from_value(value: object, *, cache: BrowserRouteCache) -> None: @@ -161,6 +165,24 @@ def _session_id_to_evict_from_response(response: httpx.Response) -> str | None: return None +def _session_id_from_direct_vm_response(response: httpx.Response, *, cache: BrowserRouteCache) -> str | None: + raw = str(response.request.url) + for route in cache.values(): + if raw.startswith(route.base_url.rstrip("/") + "/"): + return route.session_id + return None + + +def is_stale_direct_vm_auth_response(response: httpx.Response) -> bool: + if response.status_code not in {401, 403}: + return False + return bool(response.request.url.params.get("jwt")) + + +def should_retry_stale_direct_vm_auth(response: httpx.Response) -> bool: + return is_stale_direct_vm_auth_response(response) + + def _session_id_from_browser_delete_path(path: str) -> str | None: match = _BROWSER_DELETE_BY_ID_PATH.match(path) if match is None: diff --git a/tests/test_browser_routing.py b/tests/test_browser_routing.py index 3538c221..c10f7b0b 100644 --- a/tests/test_browser_routing.py +++ b/tests/test_browser_routing.py @@ -390,9 +390,14 @@ def test_browser_route_from_browser_requires_base_url_and_jwt() -> None: assert browser_route_from_browser({**_fake_browser(), "cdp_ws_url": None}) is None -def test_browser_routing_config_from_env_defaults_to_curl(monkeypatch: pytest.MonkeyPatch) -> None: +def test_browser_routing_config_from_env_defaults(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("KERNEL_BROWSER_ROUTING_SUBRESOURCES", raising=False) - assert browser_routing_config_from_env().subresources == ("curl", "telemetry/stream") + assert browser_routing_config_from_env().subresources == ( + "curl", + "telemetry/stream", + "computer", + "playwright", + ) def test_direct_vm_routing_allowlist_segment_boundary() -> None: @@ -401,13 +406,16 @@ def test_direct_vm_routing_allowlist_segment_boundary() -> None: # stream-prefixed-but-different path is not matched. from kernel.lib.browser_routing.routing import _matches_direct_vm_prefix - prefixes = ("curl", "telemetry/stream") + prefixes = ("curl", "telemetry/stream", "computer", "playwright") assert _matches_direct_vm_prefix("telemetry/stream", prefixes) is True assert _matches_direct_vm_prefix("telemetry/stream/x", prefixes) is True assert _matches_direct_vm_prefix("telemetry/events", prefixes) is False assert _matches_direct_vm_prefix("telemetry/streaming-config", prefixes) is False assert _matches_direct_vm_prefix("telemetry", prefixes) is False assert _matches_direct_vm_prefix("curl/raw", prefixes) is True + assert _matches_direct_vm_prefix("computer/screenshot", prefixes) is True + assert _matches_direct_vm_prefix("playwright/execute", prefixes) is True + assert _matches_direct_vm_prefix("process/exec", prefixes) is False assert _matches_direct_vm_prefix("fs/read", prefixes) is False @@ -424,10 +432,8 @@ def test_rewrite_direct_vm_options_keeps_telemetry_events_on_control_plane() -> ) cache = BrowserRouteCache() - cache.set( - BrowserRoute(session_id="sess-1", base_url="http://browser-session.test/browser/kernel", jwt="token-abc") - ) - config = BrowserRoutingConfig(subresources=("curl", "telemetry/stream")) + cache.set(BrowserRoute(session_id="sess-1", base_url="http://browser-session.test/browser/kernel", jwt="token-abc")) + config = BrowserRoutingConfig(subresources=("curl", "telemetry/stream", "computer", "playwright")) events = rewrite_direct_vm_options( FinalRequestOptions(method="get", url="/browsers/sess-1/telemetry/events"), cache=cache, config=config @@ -439,7 +445,117 @@ def test_rewrite_direct_vm_options_keeps_telemetry_events_on_control_plane() -> ) assert str(stream.url).startswith("http://browser-session.test/browser/kernel/telemetry/stream") + screenshot = rewrite_direct_vm_options( + FinalRequestOptions(method="post", url="/browsers/sess-1/computer/screenshot"), cache=cache, config=config + ) + assert str(screenshot.url).startswith("http://browser-session.test/browser/kernel/computer/screenshot") + + execute = rewrite_direct_vm_options( + FinalRequestOptions(method="post", url="/browsers/sess-1/playwright/execute"), cache=cache, config=config + ) + assert str(execute.url).startswith("http://browser-session.test/browser/kernel/playwright/execute") + + process = rewrite_direct_vm_options( + FinalRequestOptions(method="post", url="/browsers/sess-1/process/exec"), cache=cache, config=config + ) + assert process.url == "/browsers/sess-1/process/exec" + + fs_read = rewrite_direct_vm_options( + FinalRequestOptions(method="get", url="/browsers/sess-1/fs/read_file"), cache=cache, config=config + ) + assert fs_read.url == "/browsers/sess-1/fs/read_file" + def test_browser_routing_config_from_env_empty_string_disables_routing(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("KERNEL_BROWSER_ROUTING_SUBRESOURCES", "") assert browser_routing_config_from_env().subresources == () + + +@respx.mock +def test_computer_screenshot_and_playwright_execute_route_to_vm_by_default( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("KERNEL_BROWSER_ROUTING_SUBRESOURCES", raising=False) + screenshot = respx.post("http://browser-session.test/browser/kernel/computer/screenshot").mock( + return_value=httpx.Response(200, content=b"png", headers={"content-type": "image/png"}) + ) + execute = respx.post("http://browser-session.test/browser/kernel/playwright/execute").mock( + return_value=httpx.Response(200, json={"success": True}) + ) + with Kernel(base_url=base_url, api_key=api_key, _strict_response_validation=True) as client: + _cache_browser(client) + client.browsers.computer.capture_screenshot("sess-1") + out = client.browsers.playwright.execute("sess-1", code="return 1") + + assert screenshot.called + screenshot_req = cast(httpx.Request, cast(Any, screenshot.calls[0]).request) + assert screenshot_req.url.params.get("jwt") == "token-abc" + assert screenshot_req.headers.get("Authorization") is None + assert execute.called + execute_req = cast(httpx.Request, cast(Any, execute.calls[0]).request) + assert execute_req.url.params.get("jwt") == "token-abc" + assert execute_req.headers.get("Authorization") is None + assert out.success is True + + +@respx.mock +def test_process_fs_and_telemetry_events_stay_on_api_origin_by_default( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("KERNEL_BROWSER_ROUTING_SUBRESOURCES", raising=False) + process = respx.post(f"{base_url}/browsers/sess-1/process/exec").mock( + return_value=httpx.Response(200, json={"exit_code": 0, "stdout_b64": "", "stderr_b64": ""}) + ) + fs_read = respx.get(f"{base_url}/browsers/sess-1/fs/read_file").mock( + return_value=httpx.Response(200, content=b"x", headers={"content-type": "application/octet-stream"}) + ) + events = respx.get(f"{base_url}/browsers/sess-1/telemetry/events").mock(return_value=httpx.Response(200, json=[])) + with Kernel(base_url=base_url, api_key=api_key, _strict_response_validation=True) as client: + _cache_browser(client) + client.browsers.process.exec("sess-1", command="echo") + client.browsers.fs.read_file("sess-1", path="/tmp/x") + client.browsers.telemetry.events("sess-1") + + assert process.called + assert fs_read.called + assert events.called + + +@respx.mock +def test_stale_direct_vm_jwt_evicts_cache_and_retries_control_plane( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("KERNEL_BROWSER_ROUTING_SUBRESOURCES", raising=False) + + def _skip_retry_sleep(_self: object, **_kwargs: object) -> None: + return None + + monkeypatch.setattr("kernel._base_client.SyncAPIClient._sleep_for_retry", _skip_retry_sleep) + vm = respx.post("http://browser-session.test/browser/kernel/computer/screenshot").mock( + return_value=httpx.Response(401, text="Invalid JWT") + ) + api = respx.post(f"{base_url}/browsers/sess-1/computer/screenshot").mock( + return_value=httpx.Response(200, content=b"png", headers={"content-type": "image/png"}) + ) + with Kernel(base_url=base_url, api_key=api_key, _strict_response_validation=True) as client: + _cache_browser(client) + client.browsers.computer.capture_screenshot("sess-1") + assert client.browser_route_cache.get("sess-1") is None + + assert vm.called + assert api.called + api_req = cast(httpx.Request, cast(Any, api.calls[0]).request) + assert api_req.headers.get("Authorization") == f"Bearer {api_key}" + + +def test_stale_direct_vm_auth_retry_does_not_require_cached_route() -> None: + from kernel.lib.browser_routing.routing import should_retry_stale_direct_vm_auth + + request = httpx.Request( + "POST", + "http://browser-session.test/browser/kernel/computer/screenshot?jwt=token-abc", + ) + response = httpx.Response(401, text="Invalid JWT", request=request) + empty = BrowserRouteCache() + assert should_retry_stale_direct_vm_auth(response) is True + assert empty.get("sess-1") is None