diff --git a/application/agents/tools/artifact_generator.py b/application/agents/tools/artifact_generator.py index 977af80b..76545cca 100644 --- a/application/agents/tools/artifact_generator.py +++ b/application/agents/tools/artifact_generator.py @@ -268,13 +268,22 @@ _RENDERERS: Dict[str, str] = { "import json\n" "from openpyxl import Workbook\n" "spec = json.load(open({spec_path!r}))\n" + # Formula-injection guard: spec content is model / prompt-injection + # controlled, so neutralize string cells openpyxl would treat as a live + # formula (leading =,+,-,@ or a control char) by quote-prefixing them. + "def _safe_cell(c):\n" + " if c is None:\n" + " return ''\n" + " if isinstance(c, str) and c[:1] in ('=', '+', '-', '@', chr(9), chr(13), chr(10)):\n" + " return \"'\" + c\n" + " return c\n" "wb = Workbook()\n" "wb.remove(wb.active)\n" "for idx, sheet in enumerate(spec.get('sheets', [])):\n" " name = str(sheet.get('name') or ('Sheet%d' % (idx + 1)))[:31]\n" " ws = wb.create_sheet(title=name)\n" " for row in (sheet.get('rows') or []):\n" - " ws.append([('' if c is None else c) for c in row])\n" + " ws.append([_safe_cell(c) for c in row])\n" "if not wb.sheetnames:\n" " wb.create_sheet(title='Sheet1')\n" "wb.save({out_path!r})\n" diff --git a/application/agents/tools/attachment_bridge.py b/application/agents/tools/attachment_bridge.py index c3887cf8..9fbbfceb 100644 --- a/application/agents/tools/attachment_bridge.py +++ b/application/agents/tools/attachment_bridge.py @@ -12,6 +12,7 @@ from __future__ import annotations import logging from typing import Any, Dict, List, Optional +from application.core.settings import settings from application.sandbox.artifacts_capture import QuotaExceeded, persist_new_artifact from application.storage.db.repositories.artifacts import ArtifactsRepository from application.storage.db.repositories.attachments import AttachmentsRepository @@ -109,11 +110,30 @@ def bridge_attachment( raise AttachmentBridgeError(f"attachment {attachment_id} has no stored content.") filename = attachment.get("filename") or "attachment" mime_type = attachment.get("mime_type") or "application/octet-stream" + # Reject oversize attachments BEFORE buffering them: the authoritative ``size`` + # column lets us avoid pulling a multi-hundred-MB file fully into worker memory, + # and the bounded read below backstops a missing/lying ``size``. + max_bytes = int(getattr(settings, "ARTIFACT_MAX_BYTES", 0) or 0) + declared_size = attachment.get("size") + if max_bytes and isinstance(declared_size, (int, float)) and declared_size > max_bytes: + raise AttachmentBridgeError( + f"attachment {attachment_id} exceeds the {max_bytes}-byte artifact size limit." + ) try: - data = StorageCreator.get_storage().get_file(upload_path).read() + file_obj = StorageCreator.get_storage().get_file(upload_path) + try: + data = file_obj.read(max_bytes + 1) if max_bytes else file_obj.read() + finally: + close = getattr(file_obj, "close", None) + if callable(close): + close() except Exception as exc: logger.exception("attachment_bridge: failed to read attachment bytes") raise AttachmentBridgeError(f"failed to read attachment {attachment_id}.") from exc + if max_bytes and len(data) > max_bytes: + raise AttachmentBridgeError( + f"attachment {attachment_id} exceeds the {max_bytes}-byte artifact size limit." + ) try: ref = persist_new_artifact( user_id=user_id, diff --git a/application/sandbox/daytona.py b/application/sandbox/daytona.py index 524838a6..100225fc 100644 --- a/application/sandbox/daytona.py +++ b/application/sandbox/daytona.py @@ -89,6 +89,11 @@ class DaytonaSandbox(CodeSandbox): self._max_sandboxes = max_sandboxes self._handles: Dict[str, _Handle] = {} self._lock = threading.Lock() + # Session ids with a create/reattach in flight; a concurrent open() for the + # same id waits here and reuses the result, so two threads never create (and + # pay for) two cloud sandboxes for one session. + self._creating: set = set() + self._create_cv = threading.Condition(self._lock) # -- Helpers --------------------------------------------------------- @@ -128,41 +133,50 @@ class DaytonaSandbox(CodeSandbox): ``max_sandboxes`` so a flood of sessions cannot run up unbounded paid resources. """ self._validate_session_id(session_id) - with self._lock: + # Wait out any in-flight create/reattach for this session, then reuse its + # handle; otherwise claim the in-flight slot so we are the sole creator. + with self._create_cv: + while session_id in self._creating: + self._create_cv.wait() existing = self._handles.get(session_id) - if existing is not None: - return existing.sandbox_id - - # Cross-restart reattach: an earlier process may have created (and labelled) - # a sandbox for this session that is still live in the cloud. Reuse it - # instead of leaking it behind a brand-new create. - reattached = self._reattach_existing(session_id) - if reattached is not None: - self._prime(reattached) - return reattached.sandbox_id - - with self._lock: - if len(self._handles) >= self._max_sandboxes: - raise RuntimeError( - f"Daytona sandbox cap reached ({self._max_sandboxes} live); refusing to create another" - ) - - sandbox = self._create_sandbox(session_id) - # Crash-safe register: if anything between create and registration raises, - # delete the just-created sandbox so it cannot orphan as a paid resource. + if existing is not None: + return existing.sandbox_id + self._creating.add(session_id) try: - sandbox_id = sandbox.id - handle = _Handle(sandbox, sandbox_id, _WORKSPACE_ROOT) + # Cross-restart reattach: an earlier process may have created (and labelled) + # a sandbox for this session that is still live in the cloud. Reuse it + # instead of leaking it behind a brand-new create. + reattached = self._reattach_existing(session_id) + if reattached is not None: + self._prime(reattached) + return reattached.sandbox_id + with self._lock: - self._handles[session_id] = handle - except Exception: + if len(self._handles) >= self._max_sandboxes: + raise RuntimeError( + f"Daytona sandbox cap reached ({self._max_sandboxes} live); refusing to create another" + ) + + sandbox = self._create_sandbox(session_id) + # Crash-safe register: if anything between create and registration raises, + # delete the just-created sandbox so it cannot orphan as a paid resource. try: - self._client.delete(sandbox) - except Exception as del_exc: # noqa: BLE001 - cleanup is best-effort - logger.warning("Failed to delete orphaned Daytona sandbox during open: %s", del_exc) - raise - self._prime(handle) - return sandbox_id + sandbox_id = sandbox.id + handle = _Handle(sandbox, sandbox_id, _WORKSPACE_ROOT) + with self._lock: + self._handles[session_id] = handle + except Exception: + try: + self._client.delete(sandbox) + except Exception as del_exc: # noqa: BLE001 - cleanup is best-effort + logger.warning("Failed to delete orphaned Daytona sandbox during open: %s", del_exc) + raise + self._prime(handle) + return sandbox_id + finally: + with self._create_cv: + self._creating.discard(session_id) + self._create_cv.notify_all() def _reattach_existing(self, session_id: str) -> Optional["_Handle"]: """Find a live cloud sandbox labelled for ``session_id`` and rebuild a handle from it. diff --git a/application/sandbox/jupyter_gateway.py b/application/sandbox/jupyter_gateway.py index f055a7ea..413a09bd 100644 --- a/application/sandbox/jupyter_gateway.py +++ b/application/sandbox/jupyter_gateway.py @@ -82,6 +82,11 @@ class JupyterKernelGatewaySandbox(CodeSandbox): self._max_file_bytes = max_file_bytes self._kernels: Dict[str, _Kernel] = {} self._lock = threading.Lock() + # Session ids with a create in flight; a second open() for the same id + # waits on this CV and reuses the result instead of double-creating a + # kernel (the idempotency contract SandboxManager.open relies on). + self._creating: set = set() + self._create_cv = threading.Condition(self._lock) # -- HTTP helpers ---------------------------------------------------- @@ -114,27 +119,39 @@ class JupyterKernelGatewaySandbox(CodeSandbox): # -- Lifecycle ------------------------------------------------------- def open(self, session_id: str) -> str: - """Start a fresh kernel for ``session_id`` and prime its workspace cwd.""" - self._validate_session_id(session_id) - with self._lock: - existing = self._kernels.get(session_id) - if existing is not None: - return existing.kernel_id + """Start a fresh kernel for ``session_id`` and prime its workspace cwd. - resp = requests.post( - f"{self._base_url}/api/kernels", - headers=self._headers(), - data=json.dumps({"name": self._kernel_name}), - timeout=self._http_timeout, - ) - resp.raise_for_status() - kernel_id = resp.json()["id"] - workspace = f"{_WORKSPACE_ROOT}/{session_id}" - kernel = _Kernel(kernel_id, workspace) - with self._lock: - self._kernels[session_id] = kernel - self._prime(kernel) - return kernel_id + Idempotent under concurrency: if another thread is already creating a kernel + for this session, wait for it and reuse the result rather than POSTing a + second kernel that would orphan on the gateway. + """ + self._validate_session_id(session_id) + with self._create_cv: + while session_id in self._creating: + self._create_cv.wait() + existing = self._kernels.get(session_id) + if existing is not None: + return existing.kernel_id + self._creating.add(session_id) + try: + resp = requests.post( + f"{self._base_url}/api/kernels", + headers=self._headers(), + data=json.dumps({"name": self._kernel_name}), + timeout=self._http_timeout, + ) + resp.raise_for_status() + kernel_id = resp.json()["id"] + workspace = f"{_WORKSPACE_ROOT}/{session_id}" + kernel = _Kernel(kernel_id, workspace) + with self._lock: + self._kernels[session_id] = kernel + self._prime(kernel) + return kernel_id + finally: + with self._create_cv: + self._creating.discard(session_id) + self._create_cv.notify_all() def attach(self, session_id: str) -> str: """Reattach to a still-running kernel for ``session_id``; open a cold one if gone.""" @@ -344,7 +361,7 @@ class JupyterKernelGatewaySandbox(CodeSandbox): if msg_type == "stream": if not truncated: text = content.get("text", "") - buffered += len(text) + buffered += len(text.encode("utf-8", "ignore")) if buffered > self._max_output_bytes: truncated = True self._interrupt_and_drain(ws, msg_id, kernel_id) # runaway output: stop and drain @@ -354,7 +371,17 @@ class JupyterKernelGatewaySandbox(CodeSandbox): else: stdout_parts.append(text) elif msg_type in ("execute_result", "display_data"): - self._capture_rich(result, msg_type, content) + # Rich outputs (results/display_data/plots) count against the SAME + # byte budget as streams: untrusted code can emit a huge DataFrame / + # IPython.display.HTML / many large images that the backend would + # otherwise buffer unbounded. Drop and truncate once over budget. + if not truncated: + buffered += self._rich_payload_bytes(content) + if buffered > self._max_output_bytes: + truncated = True + self._interrupt_and_drain(ws, msg_id, kernel_id) + break + self._capture_rich(result, msg_type, content) elif msg_type == "error": result.status = "error" result.exit_code = 1 @@ -389,6 +416,20 @@ class JupyterKernelGatewaySandbox(CodeSandbox): result.error_value = value result.exit_code = -1 + @staticmethod + def _rich_payload_bytes(content: dict) -> int: + """Approximate the serialized byte size of a rich-output bundle's data payloads.""" + total = 0 + for value in (content.get("data") or {}).values(): + if isinstance(value, str): + total += len(value.encode("utf-8", "ignore")) + else: + try: + total += len(json.dumps(value, default=str).encode("utf-8", "ignore")) + except Exception: + total += len(str(value).encode("utf-8", "ignore")) + return total + @staticmethod def _capture_rich(result: ExecResult, msg_type: str, content: dict) -> None: """Sort a rich output into results/display_data and pull out any image plots.""" diff --git a/application/sandbox/manager.py b/application/sandbox/manager.py index edd4c623..a88c5bf2 100644 --- a/application/sandbox/manager.py +++ b/application/sandbox/manager.py @@ -92,17 +92,28 @@ class SandboxManager: if session is not None and session.ready: session.last_access = now return session.handle - # A placeholder for this id means another thread is mid cold-open; refresh - # it and re-reserve below. Both threads call the (idempotent) backend.open, - # which dedupes by session id, so no cap overshoot and no orphaned runtime. reaped = self._reap_locked(now) - evicted = self._make_room_locked() - self._sessions[session_id] = _Session( - session_id=session_id, - ttl=self._clamp_ttl(ttl), - created_at=now, - last_access=now, - ) + if session is None: + # Genuinely new key: evict an LRU-idle victim if at capacity (may + # raise SandboxCapacityError), then RESERVE a placeholder slot so + # concurrent opens can't overshoot the cap. + evicted = self._make_room_locked() + self._sessions[session_id] = _Session( + session_id=session_id, + ttl=self._clamp_ttl(ttl), + created_at=now, + last_access=now, + ) + else: + # A not-yet-ready placeholder for THIS id already occupies a cap slot + # (another thread is mid cold-open). Overwriting the same key adds no + # slot, so do NOT call _make_room_locked here -- it would wrongly evict + # an innocent LRU-idle session or raise SandboxCapacityError. Refresh + # in place; both threads call the (idempotent) backend.open and the + # finalize step re-checks before binding the handle. + evicted = None + session.last_access = now + session.ttl = self._clamp_ttl(ttl) # Cold backend open and victim teardown run OUTSIDE the lock. for sid, handle in reaped: diff --git a/application/services/artifact_resource_service.py b/application/services/artifact_resource_service.py index 3ef88398..b9d86968 100644 --- a/application/services/artifact_resource_service.py +++ b/application/services/artifact_resource_service.py @@ -93,8 +93,12 @@ def _is_texty(mime_type: str) -> bool: return any(mime.endswith(s) for s in _TEXT_MIME_SUFFIXES) -def _resolve_principal(api_key: Optional[str]) -> Optional[str]: - """Resolve a Bearer api_key to its owning ``user_id``; None when unresolvable.""" +def _resolve_agent(api_key: Optional[str]) -> Optional[dict]: + """Resolve a Bearer api_key to its owning agent row (carries ``id`` + ``user_id``). + + Returns the whole agent so callers can scope artifact visibility to that agent's + conversations, not the owner's entire corpus. None when unresolvable. + """ if not api_key: return None try: @@ -103,7 +107,9 @@ def _resolve_principal(api_key: Optional[str]) -> Optional[str]: except Exception: logger.exception("artifact resource: principal resolution failed") return None - return agent.get("user_id") if agent else None + if not agent or not agent.get("user_id") or not agent.get("id"): + return None + return agent def _resource_uri(artifact_id: str, version: int) -> str: @@ -122,12 +128,14 @@ def list_artifact_resources(api_key: Optional[str]) -> List[Resource]: middleware's ``on_read_resource`` intercepts every ``artifact://`` read and streams the real bytes via :func:`read_artifact_resource`. """ - user_id = _resolve_principal(api_key) - if not user_id: + agent = _resolve_agent(api_key) + if not agent: return [] try: with db_readonly() as conn: - rows = ArtifactsRepository(conn).list_artifacts(user_id=user_id) + rows = ArtifactsRepository(conn).list_artifacts_for_agent( + str(agent["id"]), str(agent["user_id"]) + ) except Exception: logger.exception("artifact resource: list failed") return [] @@ -178,9 +186,11 @@ def read_artifact_resource(api_key: Optional[str], uri: str) -> ArtifactReadResu if not looks_like_uuid(artifact_id): raise ResourceNotFound(f"artifact {artifact_id} not found") - user_id = _resolve_principal(api_key) - if not user_id: + agent = _resolve_agent(api_key) + if not agent: raise ResourceDenied("unauthenticated") + user_id = str(agent["user_id"]) + agent_id = str(agent["id"]) try: with db_readonly() as conn: @@ -188,9 +198,13 @@ def read_artifact_resource(api_key: Optional[str], uri: str) -> ArtifactReadResu artifact = repo.get_artifact(artifact_id) if artifact is None: raise ResourceNotFound(f"artifact {artifact_id} not found") - # Ownership is the authz point: never serve another principal's artifact - # over MCP, regardless of conversation/workflow parent sharing. - if str(artifact.get("user_id")) != str(user_id): + # Ownership is the first authz point: never serve another principal's + # artifact over MCP, regardless of conversation/workflow parent sharing. + if str(artifact.get("user_id")) != user_id: + raise ResourceDenied("forbidden") + # Agent scope is the second: a per-agent key only reads artifacts from + # its own conversations, not the owner's other agents / workflow runs. + if not repo.artifact_in_agent_scope(artifact_id, agent_id): raise ResourceDenied("forbidden") version_row = repo.get_version(artifact_id, version) except (DataError, DBAPIError) as exc: diff --git a/application/storage/db/repositories/artifacts.py b/application/storage/db/repositories/artifacts.py index e703d3b3..edbfbb98 100644 --- a/application/storage/db/repositories/artifacts.py +++ b/application/storage/db/repositories/artifacts.py @@ -184,6 +184,36 @@ class ArtifactsRepository: ) return [_artifact_to_dict(r) for r in result.fetchall()] + def list_artifacts_for_agent(self, agent_id: str, user_id: str) -> list[dict]: + """List artifacts whose parent conversation belongs to ``agent_id`` (owner-scoped). + + Scopes a per-agent api-key's artifact visibility to the conversations that + agent produced, so the key cannot enumerate the owner's whole corpus. + """ + result = self._conn.execute( + text( + "SELECT a.* FROM artifacts a " + "JOIN conversations c ON a.conversation_id = c.id " + "WHERE c.agent_id = CAST(:agent_id AS uuid) AND a.user_id = :user_id " + "ORDER BY a.created_at DESC, a.id DESC" + ), + {"agent_id": str(agent_id), "user_id": user_id}, + ) + return [_artifact_to_dict(r) for r in result.fetchall()] + + def artifact_in_agent_scope(self, artifact_id: str, agent_id: str) -> bool: + """True if ``artifact_id``'s parent conversation belongs to ``agent_id``.""" + result = self._conn.execute( + text( + "SELECT 1 FROM artifacts a " + "JOIN conversations c ON a.conversation_id = c.id " + "WHERE a.id = CAST(:id AS uuid) AND c.agent_id = CAST(:agent_id AS uuid) " + "LIMIT 1" + ), + {"id": artifact_id, "agent_id": str(agent_id)}, + ) + return result.fetchone() is not None + def position_in_parent( self, artifact_id: str, diff --git a/tests/agents/tools/test_artifact_generator_unit.py b/tests/agents/tools/test_artifact_generator_unit.py index 116f5c1c..b2abe0a9 100644 --- a/tests/agents/tools/test_artifact_generator_unit.py +++ b/tests/agents/tools/test_artifact_generator_unit.py @@ -153,6 +153,35 @@ def test_renderer_keeps_injection_text_as_literal_content(): assert prs.slides[0].shapes.title.text == payload +def test_spreadsheet_renderer_neutralizes_formula_injection(): + # Spec content is model / prompt-injection controlled. A string cell starting + # with a formula trigger must be stored as TEXT (data_type 's'), never as a + # live formula; genuine numbers stay numeric. + from openpyxl import load_workbook + + spec = { + "sheets": [ + { + "name": "S", + "rows": [ + ['=HYPERLINK("http://evil","x")', "+1+1", "@cmd", "-danger"], + ["safe text", -5, 3.14], + ], + } + ] + } + wb = load_workbook(_render_in_process("spreadsheet", spec)) + ws = wb["S"] + # None of the trigger cells are formulas; each is a quote-prefixed string. + for coord in ("A1", "B1", "C1", "D1"): + assert ws[coord].data_type == "s", coord + assert str(ws[coord].value).startswith("'") + # Benign text is unchanged; numbers stay numeric (no spurious quoting). + assert ws["A2"].value == "safe text" and ws["A2"].data_type == "s" + assert ws["B2"].value == -5 and ws["B2"].data_type == "n" + assert ws["C2"].value == 3.14 and ws["C2"].data_type == "n" + + def test_pdf_renderer_escapes_markup_and_does_not_execute(): payload = "not bold & \"\" '''os.system('x')'''" spec = {"title": payload, "blocks": [{"type": "paragraph", "text": payload}]} diff --git a/tests/agents/tools/test_attachment_bridge.py b/tests/agents/tools/test_attachment_bridge.py index 4d7f595b..3f4559f2 100644 --- a/tests/agents/tools/test_attachment_bridge.py +++ b/tests/agents/tools/test_attachment_bridge.py @@ -34,8 +34,12 @@ class _FakeFile: def __init__(self, data: bytes) -> None: self._data = data - def read(self) -> bytes: - return self._data + def read(self, size: int = -1) -> bytes: + # Mirror BinaryIO.read(n): the bridge does a bounded read to cap memory. + return self._data if size is None or size < 0 else self._data[:size] + + def close(self) -> None: + pass class _FakeStorage: @@ -224,6 +228,22 @@ def test_bridge_missing_upload_path_errors(monkeypatch): bridge_attachment(att, user_id=USER, conversation_id=CONV) +@pytest.mark.unit +def test_bridge_rejects_oversize_attachment_before_reading(monkeypatch): + # An oversize attachment is rejected via its authoritative ``size`` column, + # BEFORE the bytes are buffered into worker memory (a memory-DoS guard). + from application.core.settings import settings + + _FakeArtifactsRepo.bridged = {} + storage, calls = _patch_bridge(monkeypatch) + att = _attachment() + att["size"] = int(settings.ARTIFACT_MAX_BYTES) + 1 + with pytest.raises(AttachmentBridgeError, match="size limit"): + bridge_attachment(att, user_id=USER, conversation_id=CONV) + assert calls == [] # never persisted + assert storage.requested == [] # never even read the bytes + + # --------------------------------------------------------------------------- # code_executor wiring: fallback fires, stages bytes, succeeds # --------------------------------------------------------------------------- diff --git a/tests/sandbox/test_jupyter_gateway_isolation.py b/tests/sandbox/test_jupyter_gateway_isolation.py index 987d2cdb..90808074 100644 --- a/tests/sandbox/test_jupyter_gateway_isolation.py +++ b/tests/sandbox/test_jupyter_gateway_isolation.py @@ -145,3 +145,102 @@ def test_prime_creates_workspace_mode_0700(tmp_path, monkeypatch): assert kernel.initialized assert stat.S_IMODE(os.stat(root).st_mode) == 0o700 assert stat.S_IMODE(os.stat(workspace).st_mode) == 0o700 + + +# -- Output cap: rich outputs count against the byte budget -------------------- + + +def test_rich_payload_bytes_sums_all_data(): + content = {"data": {"text/html": "x" * 100, "text/plain": "y" * 50, "application/json": {"a": 1}}} + n = JupyterKernelGatewaySandbox._rich_payload_bytes(content) + assert n >= 150 # html + plain + serialized json all counted + + +def _frame(msg_id, msg_type, content): + return json.dumps({"parent_header": {"msg_id": msg_id}, "msg_type": msg_type, "content": content}) + + +class _FakeWS: + def __init__(self, frames): + self._frames = list(frames) + + def settimeout(self, _t): + pass + + def recv(self): + if self._frames: + return self._frames.pop(0) + import websocket + + raise websocket.WebSocketConnectionClosedException() + + +def test_collect_caps_oversize_rich_output(monkeypatch): + # A huge execute_result must NOT be buffered: once it would exceed the byte + # budget the bundle is dropped and the result is marked truncated. + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused", max_output_bytes=1000) + monkeypatch.setattr(sb, "_interrupt_and_drain", lambda *a, **k: None) + msg_id = "m1" + frames = [_frame(msg_id, "execute_result", {"data": {"text/html": "H" * 5000}})] + result = sb._collect(_FakeWS(frames), msg_id, timeout=5, kernel_id="k1") + assert result.results == [] # over-budget bundle dropped, not materialized + assert "[output truncated" in result.stderr + + +def test_collect_keeps_small_rich_output(monkeypatch): + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused", max_output_bytes=10000) + monkeypatch.setattr(sb, "_interrupt_and_drain", lambda *a, **k: None) + msg_id = "m2" + frames = [ + _frame(msg_id, "execute_result", {"data": {"text/plain": "small"}}), + _frame(msg_id, "execute_reply", {"status": "ok", "execution_count": 1}), + _frame(msg_id, "status", {"execution_state": "idle"}), + ] + result = sb._collect(_FakeWS(frames), msg_id, timeout=5, kernel_id="k2") + assert len(result.results) == 1 + assert "[output truncated" not in (result.stderr or "") + + +# -- open() is idempotent under concurrency ----------------------------------- + + +def test_concurrent_open_creates_one_kernel(monkeypatch): + # Two threads opening the same session must POST exactly one kernel; a second + # would orphan on the gateway. The CV guard serializes per-session creation. + import threading + + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + posts = {"n": 0} + post_lock = threading.Lock() + start = threading.Event() + + class _Resp: + def raise_for_status(self): + pass + + def json(self): + return {"id": f"kernel-{posts['n']}"} + + def _fake_post(url, **kwargs): + with post_lock: + posts["n"] += 1 + start.wait(timeout=2) # hold both threads at the POST to force the race + return _Resp() + + monkeypatch.setattr(jupyter_gateway.requests, "post", _fake_post) + monkeypatch.setattr(sb, "_prime", lambda kernel: None) + + results = {} + + def _open(i): + results[i] = sb.open("session-x") + + threads = [threading.Thread(target=_open, args=(i,)) for i in range(2)] + for t in threads: + t.start() + start.set() + for t in threads: + t.join(timeout=5) + + assert posts["n"] == 1, "concurrent open() created more than one kernel" + assert results[0] == results[1] # both callers got the same kernel id diff --git a/tests/sandbox/test_sandbox_manager.py b/tests/sandbox/test_sandbox_manager.py index ce30636e..7aaa1a4f 100644 --- a/tests/sandbox/test_sandbox_manager.py +++ b/tests/sandbox/test_sandbox_manager.py @@ -1,6 +1,7 @@ """Unit tests for SandboxManager and SandboxCreator using an in-memory backend.""" import threading +import time from typing import Dict, List import pytest @@ -256,6 +257,57 @@ def test_reuse_existing_session_never_evicts_at_cap(backend): assert mgr.session_count() == 1 +def test_concurrent_open_same_id_does_not_evict_innocent(): + """A 2nd open() of an id mid cold-open reuses its placeholder, never evicts another session. + + The session_id is derived from conversation/run id, so two concurrent requests on + one conversation both call open(same_id). The second must NOT see the first's + not-yet-ready placeholder as a reason to make room (evict an innocent LRU-idle + session or raise SandboxCapacityError) -- overwriting the same key adds no slot. + """ + state = {"entries": 0} + both_in_backend = threading.Event() + release = threading.Event() + guard = threading.Lock() + + class _BlockingSameId(FakeBackend): + def open(self, session_id: str) -> str: + if session_id == "slow": + with guard: + state["entries"] += 1 + if state["entries"] >= 2: + both_in_backend.set() + assert release.wait(timeout=5) + return super().open(session_id) + + backend = _BlockingSameId() + mgr = SandboxManager(backend, max_ttl=600, max_sessions=2) + mgr.open("keep") # a ready, idle session occupying one of the two slots + + t1 = threading.Thread(target=lambda: mgr.open("slow")) + t1.start() + # Wait until t1 has registered the "slow" placeholder and is blocked in backend.open + # (registry is now {keep, slow} = at cap). + deadline = time.monotonic() + 5 + while mgr.session_count() < 2 and time.monotonic() < deadline: + time.sleep(0.01) + assert mgr.session_count() == 2 + + t2 = threading.Thread(target=lambda: mgr.open("slow")) + t2.start() + try: + # Both opens have passed the lock section into backend.open. With the bug, t2 + # would have evicted "keep" in its lock section before reaching here. + assert both_in_backend.wait(timeout=5) + assert mgr.has_session("keep"), "concurrent same-id open evicted an innocent session" + assert "keep" not in backend.torn_down + finally: + release.set() + t1.join(timeout=5) + t2.join(timeout=5) + assert mgr.has_session("slow") and mgr.has_session("keep") + + # --------------------------------------------------------------------------- # Idle reaper # --------------------------------------------------------------------------- diff --git a/tests/services/test_artifact_resource_service.py b/tests/services/test_artifact_resource_service.py index a4e9ed06..1df1dfaf 100644 --- a/tests/services/test_artifact_resource_service.py +++ b/tests/services/test_artifact_resource_service.py @@ -30,6 +30,8 @@ STRANGER = "stranger-2" ART_TEXT = "11111111-1111-4111-8111-111111111111" ART_BIN = "22222222-2222-4222-8222-222222222222" ART_FOREIGN = "33333333-3333-4333-8333-333333333333" +# Owned by OWNER but produced by a DIFFERENT agent — must be invisible to owner-key. +ART_OTHER_AGENT = "44444444-4444-4444-8444-444444444444" @contextmanager @@ -38,9 +40,12 @@ def _fake_conn(): class _FakeAgents: - """Stub AgentsRepository: maps api_key -> agent row (or None).""" + """Stub AgentsRepository: maps api_key -> agent row (id + user_id) or None.""" - _MAP = {"owner-key": {"user_id": OWNER}, "stranger-key": {"user_id": STRANGER}} + _MAP = { + "owner-key": {"id": "agent-owner", "user_id": OWNER}, + "stranger-key": {"id": "agent-stranger", "user_id": STRANGER}, + } def __init__(self, conn): pass @@ -58,8 +63,16 @@ class _FakeArtifacts: def __init__(self, conn): pass - def list_artifacts(self, user_id=None, conversation_id=None, workflow_run_id=None): - return [a for a in self.artifacts.values() if a["user_id"] == user_id] + def list_artifacts_for_agent(self, agent_id, user_id): + return [ + a + for a in self.artifacts.values() + if a.get("agent_id") == agent_id and a["user_id"] == user_id + ] + + def artifact_in_agent_scope(self, artifact_id, agent_id): + art = self.artifacts.get(artifact_id) + return art is not None and art.get("agent_id") == agent_id def get_artifact(self, artifact_id): return self.artifacts.get(artifact_id) @@ -81,20 +94,30 @@ class _FakeStorage: def _wire(monkeypatch): """Point the service's DB/storage seams at the in-memory fakes.""" _FakeArtifacts.artifacts = { - ART_TEXT: {"id": ART_TEXT, "user_id": OWNER, "kind": "data", "title": "notes", "current_version": 2}, - ART_BIN: {"id": ART_BIN, "user_id": OWNER, "kind": "image", "title": "chart", "current_version": 1}, + ART_TEXT: {"id": ART_TEXT, "user_id": OWNER, "agent_id": "agent-owner", "kind": "data", "title": "notes", "current_version": 2}, + ART_BIN: {"id": ART_BIN, "user_id": OWNER, "agent_id": "agent-owner", "kind": "image", "title": "chart", "current_version": 1}, ART_FOREIGN: { "id": ART_FOREIGN, "user_id": STRANGER, + "agent_id": "agent-stranger", "kind": "data", "title": "secret", "current_version": 1, }, + ART_OTHER_AGENT: { + "id": ART_OTHER_AGENT, + "user_id": OWNER, + "agent_id": "agent-other", + "kind": "data", + "title": "other-agent", + "current_version": 1, + }, } _FakeArtifacts.versions = { (ART_TEXT, 2): {"mime_type": "text/csv", "storage_path": "k/text.csv", "preview_text": None}, (ART_BIN, 1): {"mime_type": "image/png", "storage_path": "k/chart.png", "preview_text": None}, (ART_FOREIGN, 1): {"mime_type": "text/plain", "storage_path": "k/secret.txt", "preview_text": None}, + (ART_OTHER_AGENT, 1): {"mime_type": "text/plain", "storage_path": "k/other.txt", "preview_text": None}, } monkeypatch.setattr(svc, "db_readonly", _fake_conn) monkeypatch.setattr(svc, "AgentsRepository", _FakeAgents) @@ -156,6 +179,12 @@ class TestReadArtifactResource: res = svc.read_artifact_resource("owner-key", f"artifact://{ART_TEXT}/v2") assert res.text == "cached preview" + def test_read_denies_owner_artifact_from_another_agent(self): + # Owned by the same user but produced by a different agent: a per-agent key + # is scoped like its search and must NOT read the owner's other-agent corpus. + with pytest.raises(svc.ResourceDenied): + svc.read_artifact_resource("owner-key", f"artifact://{ART_OTHER_AGENT}/v1") + def test_foreign_owner_is_denied(self): with pytest.raises(svc.ResourceDenied): svc.read_artifact_resource("owner-key", f"artifact://{ART_FOREIGN}/v1")