Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/langchain.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 |
Expand Down
4 changes: 3 additions & 1 deletion docs/sdk-python.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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/<name>`; `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
Expand Down
14 changes: 7 additions & 7 deletions sdk/langchain/cocoonsandbox_langchain/toolkit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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")
Expand Down
28 changes: 27 additions & 1 deletion sdk/langchain/tests/test_toolkit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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"})
Expand Down Expand Up @@ -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
6 changes: 4 additions & 2 deletions sdk/python/cocoonsandbox/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
42 changes: 29 additions & 13 deletions sdk/python/cocoonsandbox/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,14 +7,15 @@
import queue
import ssl
import threading
import time
import urllib.error
import urllib.parse
import urllib.request
from collections.abc import Callable, Iterable, Mapping, Sequence
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
Expand Down Expand Up @@ -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)."""
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -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)
Expand All @@ -190,18 +205,19 @@ 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:
detail = _error_message(exc.read())
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:
Expand Down
19 changes: 10 additions & 9 deletions sdk/python/cocoonsandbox/conn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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)
11 changes: 8 additions & 3 deletions sdk/python/cocoonsandbox/template.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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:
Expand Down
9 changes: 9 additions & 0 deletions sdk/python/tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
24 changes: 24 additions & 0 deletions sdk/python/tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import json
import threading
import time
from http.server import BaseHTTPRequestHandler, HTTPServer

import pytest
Expand Down Expand Up @@ -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 = []

Expand Down