diff --git a/docs/langchain.md b/docs/langchain.md index 0a849414..dcb29525 100644 --- a/docs/langchain.md +++ b/docs/langchain.md @@ -20,7 +20,7 @@ schemas, sync-native with `asyncio.to_thread` async bridges): | tool | what it does | |---|---| -| `sandbox_exec` | run a shell command, cut off after 5 minutes with the reply saying so; stdout/stderr/exit code; disk state persists across calls | +| `sandbox_exec` | run a shell command, cut off after 5 minutes with the reply saying so; the budget is one wall clock over the claim, the dial and the command, so a first call that waits on a slow or redirecting cluster still answers inside it; stdout/stderr/exit code; disk state persists across calls | | `sandbox_write_file` | write a text file (atomic on the guest; the parent directory must exist) | | `sandbox_read_file` | read a text file | | `sandbox_list_dir` | list a directory as JSON | diff --git a/docs/sdk-python.md b/docs/sdk-python.md index 3b8f1f5f..bddb5318 100644 --- a/docs/sdk-python.md +++ b/docs/sdk-python.md @@ -73,7 +73,8 @@ tenant token (resource-creating verbs only; operator surfaces answer it 403). On a cluster every node shares the root token and the same tenants set. `timeout` bounds every control-plane request and the data-plane dial and upgrade; a guest stream then lives until the guest ends it, as in the -Go SDK. +Go SDK. A caller with a wall clock of its own passes `deadline` to a claim +(below) so the redirect walk cannot outlive it. **Clusters need nothing extra**: dial any node. On a warm miss the entry node answers with a redirect and `new` follows it transparently; the @@ -101,6 +102,7 @@ sb = client.new("ghcr.io/cocoonstack/sandbox/rt:24.04", | `volumes` | bare names or `{name, mount?, mode?}` mappings | `None` | attach and mount up to eight unique catalog dataset disks; an omitted mount defaults to `/volumes/`; `mode` is `"ro"` (default) or `"rw"` — `"rw"` requires the catalog entry's `writable: true`; accepted by `Client.new` and `Template.new` | | `mount` | bool | `True` | mount every requested volume. `False` attaches the devices and leaves the mounting to the workload; a mapping carrying `mount` is then a `TypeError` | | `ttl_seconds` | int | server default 5m | sandbox TTL, server-capped at 24h. The node reaps the sandbox after the TTL even if the client vanishes | +| `deadline` | `time.monotonic()` value | `None` | keyword-only wall clock over the whole claim, redirect targets included: every request is bounded by whatever is left and a spent deadline raises `TimeoutError`. Without it each redirect candidate gets a fresh `timeout`, so a cluster with unreachable peers can cost a multiple of it. Also on `Template.new` and `Checkpoint.new` | `new` returns when the sandbox's silkd answers: a warm hit is milliseconds, a cold key can take the full boot. A volume claim may consume an ordinary warm diff --git a/sdk/langchain/cocoonsandbox_langchain/toolkit.py b/sdk/langchain/cocoonsandbox_langchain/toolkit.py index be74a151..32bad383 100644 --- a/sdk/langchain/cocoonsandbox_langchain/toolkit.py +++ b/sdk/langchain/cocoonsandbox_langchain/toolkit.py @@ -114,19 +114,19 @@ def close(self) -> None: if sb is not None: sb.close() - def sandbox(self) -> Sandbox: - """The claimed sandbox, claiming (or branching) on first use.""" + def sandbox(self, deadline: float | None = None) -> Sandbox: + """The claimed sandbox, claiming (or branching) on first use within the caller's deadline.""" with self._lock: if self._closed: raise RuntimeError("toolkit is closed") if self._sb is None: - self._sb = self._claim() + self._sb = self._claim(deadline) return self._sb - def _claim(self) -> Sandbox: + def _claim(self, deadline: float | None = None) -> Sandbox: if self._from_checkpoint: - return self._client.checkpoint(self._from_checkpoint).new(ttl_seconds=self._ttl) - return self._client.new(self._template, net=self._net, ttl_seconds=self._ttl) + return self._client.checkpoint(self._from_checkpoint).new(ttl_seconds=self._ttl, deadline=deadline) + return self._client.new(self._template, net=self._net, ttl_seconds=self._ttl, deadline=deadline) def _tool(self, name: str, description: str, schema: type[BaseModel], func: Callable[..., str]) -> StructuredTool: async def arun(**kwargs): @@ -142,7 +142,7 @@ def _exec(self, command: str, cwd: str = "") -> str: errs: list[bytes] = [] tail = "" try: - sb = self.sandbox() + sb = self.sandbox(deadline) timeout = deadline - time.monotonic() if timeout <= 0: raise TimeoutError("the claim used the whole call budget") diff --git a/sdk/langchain/tests/test_toolkit.py b/sdk/langchain/tests/test_toolkit.py index d28b5b21..73cde3c0 100644 --- a/sdk/langchain/tests/test_toolkit.py +++ b/sdk/langchain/tests/test_toolkit.py @@ -2,10 +2,12 @@ shaping, lazy claim, and double-close safety — no real node.""" import asyncio +import time import pytest from cocoonsandbox_langchain import CocoonToolkit +from cocoonsandbox_langchain.toolkit import CALL_TIMEOUT def test_tools_shape(monkeypatch): @@ -41,6 +43,30 @@ def test_file_tools_round_trip(monkeypatch): assert "a.txt" in list_dir.invoke({"path": "/w"}) +def test_exec_claims_inside_the_call_budget(monkeypatch): + kit = CocoonToolkit("127.0.0.1:1") + seen = [] + + def claim(deadline=None): + seen.append(deadline) + return FakeSandbox() + + monkeypatch.setattr(kit, "_claim", claim) + kit.get_tools()[0].invoke({"command": "echo hi"}) + assert len(seen) == 1 and seen[0] is not None + assert 0 < seen[0] - time.monotonic() <= CALL_TIMEOUT, seen[0] + + +def test_exec_reports_a_claim_that_outlives_the_budget(monkeypatch): + kit = CocoonToolkit("127.0.0.1:1") + + def claim(deadline=None): + raise TimeoutError("claim timed out") + + monkeypatch.setattr(kit, "_claim", claim) + assert kit.get_tools()[0].invoke({"command": "echo hi"}) == f"cut off after {CALL_TIMEOUT}s" + + def test_close_releases_once(monkeypatch): kit, fake = hooked(monkeypatch) kit.get_tools()[0].invoke({"command": "x"}) @@ -89,5 +115,5 @@ def close(self): def hooked(monkeypatch): kit = CocoonToolkit("127.0.0.1:1") fake = FakeSandbox() - monkeypatch.setattr(kit, "_claim", lambda: fake) + monkeypatch.setattr(kit, "_claim", lambda deadline=None: fake) return kit, fake diff --git a/sdk/python/cocoonsandbox/checkpoint.py b/sdk/python/cocoonsandbox/checkpoint.py index 414d51e1..3a1afe84 100644 --- a/sdk/python/cocoonsandbox/checkpoint.py +++ b/sdk/python/cocoonsandbox/checkpoint.py @@ -20,10 +20,12 @@ def __init__(self, client: Client, addr: str, rec: dict[str, Any]) -> None: self.sandbox_id = rec.get("sandbox_id", "") self.created_at = rec.get("created_at", "") - def new(self, ttl_seconds: int = 0) -> Sandbox: + def new(self, ttl_seconds: int = 0, *, deadline: float | None = None) -> Sandbox: """Claims from the checkpoint, following redirects with one origin fallback.""" claim = {"ttl_seconds": ttl_seconds} if ttl_seconds else {} - return self._client._claim_from(self._addr, claim, f"/v1/checkpoints/{self.id}/claim", "claim checkpoint") + return self._client._claim_from( + self._addr, claim, f"/v1/checkpoints/{self.id}/claim", "claim checkpoint", deadline=deadline + ) def delete(self) -> None: """Deletes the checkpoint with eventual peer cleanup bounded by checkpoint_ttl_hours.""" diff --git a/sdk/python/cocoonsandbox/client.py b/sdk/python/cocoonsandbox/client.py index 85ad3be7..c59d6db2 100644 --- a/sdk/python/cocoonsandbox/client.py +++ b/sdk/python/cocoonsandbox/client.py @@ -7,6 +7,7 @@ import queue import ssl import threading +import time import urllib.error import urllib.parse import urllib.request @@ -14,7 +15,7 @@ from typing import Any, TypeVar, cast from .checkpoint import Checkpoint -from .conn import Conn, dial_agent +from .conn import Conn, dial_agent, remaining_timeout from .endpoint import _endpoint_url from .errors import APIError from .sandbox import Sandbox @@ -57,10 +58,12 @@ def new( claim_ref: str = "", volumes: list[str | Mapping[str, str]] | None = None, mount: bool = True, + *, + deadline: float | None = None, ) -> Sandbox: - """Claims a sandbox; a warm hit is milliseconds.""" + """Claims a sandbox; a warm hit is milliseconds. deadline is a time.monotonic() bound over the whole claim.""" claim = _claim_body(template, net, size, ttl_seconds, volumes, mount, claim_ref) - return self._claim_from(self.addr, claim) + return self._claim_from(self.addr, claim, deadline=deadline) def delete_template(self, template: str, net: str = "", size: str = "") -> None: """Removes a promoted template by name; on a cluster the delete follows gossip to the owner node (one hop).""" @@ -130,8 +133,16 @@ def info(self) -> dict[str, Any]: """The node's pool/claim counters, as served by GET /v1/info.""" return self._request(self.addr, "GET", "/v1/info", None, "info") - def _claim_from(self, addr: str, claim: dict[str, Any], path: str = "/v1/claim", verb: str = "claim") -> Sandbox: - reply = self._post_json(addr, path, claim, verb) + def _claim_from( + self, + addr: str, + claim: dict[str, Any], + path: str = "/v1/claim", + verb: str = "claim", + *, + deadline: float | None = None, + ) -> Sandbox: + reply = self._post_json(addr, path, claim, verb, deadline=deadline) redirect = reply.get("redirect") or [] if not redirect: return self._handle_from(addr, reply) @@ -140,7 +151,7 @@ def _claim_from(self, addr: str, claim: dict[str, Any], path: str = "/v1/claim", claim["require_promoted"] = True def post(peer: str) -> dict[str, Any]: - return self._post_json(peer, path, claim, verb) + return self._post_json(peer, path, claim, verb, deadline=deadline) owner, reply = _redirect_fallback(addr, redirect, post, verb) return self._handle_from(owner, reply) @@ -168,8 +179,10 @@ def _handle_from(self, dialed: str, reply: dict[str, Any]) -> Sandbox: volumes=reply.get("volumes") or [], ) - def _post_json(self, addr: str, path: str, body: dict[str, Any], verb: str) -> dict[str, Any]: - return self._request(addr, "POST", path, body, verb) + def _post_json( + self, addr: str, path: str, body: dict[str, Any], verb: str, *, deadline: float | None = None + ) -> dict[str, Any]: + return self._request(addr, "POST", path, body, verb, deadline=deadline) def _request( self, @@ -180,7 +193,9 @@ def _request( verb: str, bearer: str = "", timeout: float = 0.0, + deadline: float | None = None, ) -> dict[str, Any]: + timeout = remaining_timeout(timeout or self.timeout, deadline, verb) data = json.dumps(body).encode() if body is not None else None url = addr + path if "://" in addr else f"{self._scheme}://{addr}{path}" req = urllib.request.Request(url, data=data, method=method) @@ -190,7 +205,7 @@ def _request( if token: req.add_header("Authorization", f"Bearer {token}") try: - with self._open(req, timeout or self.timeout) as resp: + with self._open(req, timeout) as resp: raw = resp.read() except urllib.error.HTTPError as exc: try: @@ -198,10 +213,11 @@ def _request( except (OSError, http.client.HTTPException) as read_exc: detail = str(read_exc) raise APIError(verb, exc.code, detail) from None - except urllib.error.URLError as exc: - raise APIError(verb, 0, str(exc.reason)) from None - except (OSError, http.client.HTTPException) as exc: - raise APIError(verb, 0, str(exc)) from None + except (urllib.error.URLError, OSError, http.client.HTTPException) as exc: + if deadline is not None and time.monotonic() >= deadline: + raise TimeoutError(f"{verb} timed out") from None + detail = str(exc.reason) if isinstance(exc, urllib.error.URLError) else str(exc) + raise APIError(verb, 0, detail) from None if not raw: return {} try: diff --git a/sdk/python/cocoonsandbox/conn.py b/sdk/python/cocoonsandbox/conn.py index d1f21596..cb420def 100644 --- a/sdk/python/cocoonsandbox/conn.py +++ b/sdk/python/cocoonsandbox/conn.py @@ -160,7 +160,7 @@ def dial_agent( host = endpoint.hostname port = endpoint.port or (443 if endpoint.scheme == "https" else 80) try: - sock = socket.create_connection((host, port), timeout=_remaining_timeout(timeout, deadline)) + sock = socket.create_connection((host, port), timeout=remaining_timeout(timeout, deadline)) except OSError as exc: raise ProtocolError(f"dial {addr}: {exc}") from exc sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) @@ -169,7 +169,7 @@ def dial_agent( if endpoint.scheme == "https": context = ssl_context or ssl.create_default_context() sock = context.wrap_socket(sock, server_hostname=host, do_handshake_on_connect=False) - sock.settimeout(_remaining_timeout(timeout, deadline)) + sock.settimeout(remaining_timeout(timeout, deadline)) try: sock.do_handshake() except OSError as exc: @@ -182,16 +182,16 @@ def dial_agent( f"Authorization: Bearer {token}\r\n" "\r\n" ) - sock.settimeout(_remaining_timeout(timeout, deadline)) + sock.settimeout(remaining_timeout(timeout, deadline)) sock.sendall(request.encode()) reader = sock.makefile("rb") - sock.settimeout(_remaining_timeout(timeout, deadline)) + sock.settimeout(remaining_timeout(timeout, deadline)) status = reader.readline(1024).decode(errors="replace") parts = status.split(" ", 2) code = int(parts[1]) if len(parts) > 1 and parts[1].isdigit() else 0 body_len = 0 while True: - sock.settimeout(_remaining_timeout(timeout, deadline)) + sock.settimeout(remaining_timeout(timeout, deadline)) header = reader.readline(4096) if header in (b"\r\n", b"\n", b""): break @@ -204,10 +204,10 @@ def dial_agent( except ValueError as exc: raise ProtocolError("invalid content-length in upgrade reply") from exc if code != 101: - sock.settimeout(_remaining_timeout(timeout, deadline)) + sock.settimeout(remaining_timeout(timeout, deadline)) body = reader.read(min(body_len, MAX_FRAME)).decode(errors="replace") if body_len else "" raise APIError("agent upgrade", code, body.strip() or status.strip()) - _remaining_timeout(timeout, deadline) + remaining_timeout(timeout, deadline) sock.settimeout(None) return Conn(sock, reader) except Exception: @@ -217,10 +217,11 @@ def dial_agent( raise -def _remaining_timeout(timeout: float, deadline: float | None) -> float: +def remaining_timeout(timeout: float, deadline: float | None, what: str = "agent dial") -> float: + """What is left of timeout under an optional absolute monotonic deadline; raises once it is spent.""" if deadline is None: return timeout remaining = deadline - time.monotonic() if remaining <= 0: - raise TimeoutError("agent dial timed out") + raise TimeoutError(f"{what} timed out") return min(timeout, remaining) diff --git a/sdk/python/cocoonsandbox/template.py b/sdk/python/cocoonsandbox/template.py index 5c49aef4..d367e403 100644 --- a/sdk/python/cocoonsandbox/template.py +++ b/sdk/python/cocoonsandbox/template.py @@ -23,7 +23,12 @@ def __init__(self, client: Client, addr: str, name: str, net: str, size: str, co self.content_digest = content_digest def new( - self, ttl_seconds: int = 0, volumes: list[str | Mapping[str, str]] | None = None, mount: bool = True + self, + ttl_seconds: int = 0, + volumes: list[str | Mapping[str, str]] | None = None, + mount: bool = True, + *, + deadline: float | None = None, ) -> Sandbox: """Claims the template, following placement when volumes require it. mount=False attaches the volumes without mounting them.""" @@ -32,9 +37,9 @@ def new( claim = _claim_body(self.name, self.net, self.size, ttl_seconds, volumes, mount) if volumes: - return self._client._claim_from(self._addr, claim) + return self._client._claim_from(self._addr, claim, deadline=deadline) claim["no_redirect"] = True - reply = self._client._post_json(self._addr, "/v1/claim", claim, "claim") + reply = self._client._post_json(self._addr, "/v1/claim", claim, "claim", deadline=deadline) return self._client._handle_from(self._addr, reply) def delete(self) -> None: diff --git a/sdk/python/tests/conftest.py b/sdk/python/tests/conftest.py index 6101223b..a28f8b9b 100644 --- a/sdk/python/tests/conftest.py +++ b/sdk/python/tests/conftest.py @@ -50,6 +50,15 @@ def spawn(routes): server.shutdown() +@pytest.fixture +def black_hole(): + sock = socket.socket() + sock.bind(("127.0.0.1", 0)) + sock.listen(8) + yield f"127.0.0.1:{sock.getsockname()[1]}" + sock.close() + + @pytest.fixture def dead_addr(): sock = socket.socket() diff --git a/sdk/python/tests/test_client.py b/sdk/python/tests/test_client.py index 44020ee9..2fbbc36e 100644 --- a/sdk/python/tests/test_client.py +++ b/sdk/python/tests/test_client.py @@ -2,6 +2,7 @@ import json import threading +import time from http.server import BaseHTTPRequestHandler, HTTPServer import pytest @@ -306,6 +307,29 @@ def test_checkpoint_listing_binds_handles(node): assert branch.id == "sb_branch" +def test_claim_refuses_a_spent_deadline(node): + seen = recording_claim({"id": "sb_1", "token": "tok"}) + with pytest.raises(TimeoutError): + Client(node).new("rt:24.04", deadline=time.monotonic() - 1) + assert seen == [], "a spent deadline still reached the node" + + +def test_claim_deadline_bounds_the_redirect_walk(node, black_hole): + FakeNode.routes[("POST", "/v1/claim")] = lambda body, path: (200, {"redirect": [black_hole, black_hole]}) + client = Client(node, timeout=2.0) + started = time.monotonic() + with pytest.raises(TimeoutError): + client.new("rt:24.04", deadline=started + 0.3) + elapsed = time.monotonic() - started + assert elapsed < 1.5, elapsed + + +def test_claim_deadline_bounds_the_entry_node(black_hole): + client = Client(black_hole, timeout=2.0) + with pytest.raises(TimeoutError): + client.new("rt:24.04", deadline=time.monotonic() + 0.3) + + def recording_claim(reply): seen = []