mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 00:13:14 +00:00
fix: injection protections and concurrency improvements
This commit is contained in:
1 parent
b144265094
commit
99484e02c8
12 files changed
+452
-84
No files matched your search
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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}]}
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
|
||||
Reference in new issue
Block a user