fix: injection protections and concurrency improvements

This commit is contained in:
Alex committed 2026-06-29 19:56:10 +02:00
1 parent b144265094
commit 99484e02c8
12 files changed
+452 -84

No files matched your search

+10 -1
View File
@@ -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"
+21 -1
View File
@@ -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,
+45 -31
View File
@@ -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.
+63 -22
View File
@@ -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."""
+21 -10
View File
@@ -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:
@@ -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:
@@ -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,
@@ -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 = "<b>not bold</b> & \"</para>\" '''os.system('x')'''"
spec = {"title": payload, "blocks": [{"type": "paragraph", "text": payload}]}
+22 -2
View File
@@ -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
# ---------------------------------------------------------------------------
@@ -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
+52
View File
@@ -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
# ---------------------------------------------------------------------------
@@ -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")