mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 18:13:03 +00:00
Wrap a storage read in a context manager (close the file handle), add a comment on the expected SandboxCapacityError in the churn test, and unify the artifacts-routes test on a single import style. Also annotate the intentional per-session sandbox workspace path with # nosec B108 (controlled dir inside the runner container, not an insecure shared temp file).
282 lines
10 KiB
Python
282 lines
10 KiB
Python
"""Workflow input-document bridge: uploaded attachments become run-scoped artifacts.
|
|
|
|
The agent pre-creates the ``workflow_runs`` row, re-persists each attachment's bytes
|
|
through the canonical artifact path (server-side size/sha256/storage key), and passes
|
|
the resulting references into the run as ``initial_inputs["input_documents"]`` so nodes
|
|
can read ``agent.input_documents``.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import io
|
|
import uuid
|
|
|
|
import pytest
|
|
from sqlalchemy import text
|
|
|
|
from application.agents.workflow_agent import WorkflowAgent, _MAX_INPUT_DOCUMENTS
|
|
from application.agents.workflows.workflow_engine import WorkflowEngine
|
|
from application.storage.db.repositories.artifacts import ArtifactsRepository
|
|
from application.storage.db.repositories.workflow_runs import WorkflowRunsRepository
|
|
from application.storage.local import LocalStorage
|
|
from application.storage.storage_creator import StorageCreator
|
|
|
|
pytestmark = pytest.mark.integration
|
|
|
|
OWNER = "user-bridge"
|
|
|
|
|
|
def _wire(pg_engine, tmp_path, monkeypatch) -> LocalStorage:
|
|
"""Point storage + the db session at the ephemeral fixtures."""
|
|
storage = LocalStorage(base_dir=str(tmp_path))
|
|
monkeypatch.setattr(StorageCreator, "_instance", storage, raising=False)
|
|
monkeypatch.setattr("application.storage.db.session.get_engine", lambda: pg_engine)
|
|
return storage
|
|
|
|
|
|
def _make_workflow(pg_engine, owner: str = OWNER) -> str:
|
|
"""Insert an owned workflow row and return its id."""
|
|
wf_id = str(uuid.uuid4())
|
|
with pg_engine.begin() as conn:
|
|
conn.execute(
|
|
text(
|
|
"INSERT INTO workflows (id, user_id, name, current_graph_version) "
|
|
"VALUES (CAST(:id AS uuid), :uid, :name, 1)"
|
|
),
|
|
{"id": wf_id, "uid": owner, "name": "Bridge WF"},
|
|
)
|
|
return wf_id
|
|
|
|
|
|
def _stage_attachment(storage: LocalStorage, data: bytes, filename: str, mime: str) -> dict:
|
|
"""Write attachment bytes to storage and return the attachment dict shape."""
|
|
upload_path = f"inputs/{OWNER}/attachments/{uuid.uuid4()}_{filename}"
|
|
storage.save_file(io.BytesIO(data), upload_path)
|
|
return {
|
|
"id": str(uuid.uuid4()),
|
|
"filename": filename,
|
|
"upload_path": upload_path,
|
|
"path": upload_path,
|
|
"mime_type": mime,
|
|
"size": len(data),
|
|
"user_id": OWNER,
|
|
}
|
|
|
|
|
|
def _agent(workflow_id, attachments, owner: str = OWNER) -> WorkflowAgent:
|
|
"""Build a WorkflowAgent without invoking the LLM-creating base __init__."""
|
|
agent = WorkflowAgent.__new__(WorkflowAgent)
|
|
agent.workflow_id = workflow_id
|
|
agent.workflow_owner = owner
|
|
agent.decoded_token = {"sub": owner}
|
|
agent.attachments = attachments
|
|
agent.chat_history = []
|
|
agent.retrieved_docs = []
|
|
agent._workflow_data = None
|
|
agent._engine = None
|
|
agent._run_persisted = False
|
|
return agent
|
|
|
|
|
|
_EMBEDDED_GRAPH = {
|
|
"name": "Draft",
|
|
"nodes": [
|
|
{"id": "n1", "type": "start", "title": "Start"},
|
|
{"id": "n2", "type": "end", "title": "End", "data": {}},
|
|
],
|
|
"edges": [{"id": "e1", "source": "n1", "target": "n2"}],
|
|
}
|
|
|
|
|
|
class _RecordingEngine(WorkflowEngine):
|
|
"""Engine that records initial_inputs and runs the run-row existence probe."""
|
|
|
|
probe = None
|
|
instances: list = []
|
|
|
|
def __init__(self, graph, agent, workflow_run_id=None):
|
|
super().__init__(graph, agent, workflow_run_id=workflow_run_id)
|
|
self.captured_inputs = None
|
|
_RecordingEngine.instances.append(self)
|
|
|
|
def execute(self, initial_inputs, query):
|
|
self.captured_inputs = initial_inputs
|
|
if _RecordingEngine.probe is not None:
|
|
_RecordingEngine.probe(self.workflow_run_id)
|
|
self._initialize_state(initial_inputs, query)
|
|
return iter(())
|
|
|
|
|
|
def _patch_engine(monkeypatch, probe=None) -> None:
|
|
"""Make ``_gen_inner`` build the recording engine and reset its capture state."""
|
|
_RecordingEngine.instances = []
|
|
_RecordingEngine.probe = probe
|
|
monkeypatch.setattr(
|
|
"application.agents.workflow_agent.WorkflowEngine", _RecordingEngine
|
|
)
|
|
|
|
|
|
def test_attachments_bridge_to_run_scoped_artifacts(pg_engine, tmp_path, monkeypatch):
|
|
"""N attachments -> N run-scoped artifacts + input_documents refs; nodes can read them."""
|
|
storage = _wire(pg_engine, tmp_path, monkeypatch)
|
|
wf_id = _make_workflow(pg_engine)
|
|
|
|
a1 = b"report-one-bytes"
|
|
a2 = b"second attachment payload"
|
|
attachments = [
|
|
_stage_attachment(storage, a1, "report.txt", "text/plain"),
|
|
_stage_attachment(storage, a2, "data.csv", "text/csv"),
|
|
]
|
|
agent = _agent(wf_id, attachments)
|
|
|
|
run_seen = {}
|
|
|
|
def _probe(run_id):
|
|
with pg_engine.connect() as conn:
|
|
run_seen["row"] = WorkflowRunsRepository(conn).get(run_id)
|
|
|
|
_patch_engine(monkeypatch, probe=_probe)
|
|
|
|
list(agent._gen_inner("summarize", log_context=None))
|
|
engine = _RecordingEngine.instances[-1]
|
|
|
|
# The run row existed BEFORE execute (so a mid-run download would authz).
|
|
assert run_seen["row"] is not None
|
|
assert run_seen["row"]["user_id"] == OWNER
|
|
|
|
# initial_inputs carried the refs into the run.
|
|
refs = engine.captured_inputs["input_documents"]
|
|
assert len(refs) == 2
|
|
assert {r["filename"] for r in refs} == {"report.txt", "data.csv"}
|
|
assert all(r["artifact_id"] for r in refs)
|
|
assert refs[0]["ref"] == "A1"
|
|
assert refs[1]["ref"] == "A2"
|
|
|
|
# N run-scoped artifacts persisted, parented to THIS run, server-computed size/sha256.
|
|
run_id = engine.workflow_run_id
|
|
with pg_engine.connect() as conn:
|
|
repo = ArtifactsRepository(conn)
|
|
by_name = {}
|
|
for ref, payload in zip(refs, (a1, a2)):
|
|
artifact = repo.get_artifact_in_parent(ref["artifact_id"], workflow_run_id=run_id)
|
|
assert artifact is not None
|
|
assert artifact["kind"] == "file"
|
|
version = repo.get_version(ref["artifact_id"], 1)
|
|
assert version["size"] == len(payload)
|
|
assert version["sha256"] == hashlib.sha256(payload).hexdigest()
|
|
by_name[version["filename"]] = version
|
|
assert set(by_name) == {"report.txt", "data.csv"}
|
|
assert by_name["report.txt"]["size"] == len(a1)
|
|
|
|
# A node/template can read agent.input_documents from the engine state.
|
|
context = engine._build_template_context()
|
|
assert context["agent"]["input_documents"] == refs
|
|
assert len(context["agent"]["input_documents"]) == 2
|
|
|
|
# The bytes round-trip from storage (never entered state).
|
|
with pg_engine.connect() as conn:
|
|
v = ArtifactsRepository(conn).get_version(refs[0]["artifact_id"], 1)
|
|
with storage.get_file(v["storage_path"]) as fh:
|
|
assert fh.read() == a1
|
|
|
|
|
|
def test_attachments_capped_per_run(pg_engine, tmp_path, monkeypatch):
|
|
"""More than the cap of attachments bridges only the cap; the rest are dropped."""
|
|
storage = _wire(pg_engine, tmp_path, monkeypatch)
|
|
wf_id = _make_workflow(pg_engine)
|
|
|
|
over = _MAX_INPUT_DOCUMENTS + 5
|
|
attachments = [
|
|
_stage_attachment(storage, f"doc-{i}".encode(), f"f{i}.txt", "text/plain")
|
|
for i in range(over)
|
|
]
|
|
agent = _agent(wf_id, attachments)
|
|
_patch_engine(monkeypatch)
|
|
|
|
list(agent._gen_inner("summarize", log_context=None))
|
|
engine = _RecordingEngine.instances[-1]
|
|
|
|
refs = engine.captured_inputs["input_documents"]
|
|
assert len(refs) == _MAX_INPUT_DOCUMENTS
|
|
|
|
run_id = engine.workflow_run_id
|
|
with pg_engine.connect() as conn:
|
|
n = conn.execute(
|
|
text(
|
|
"SELECT count(*) FROM artifacts WHERE workflow_run_id = CAST(:r AS uuid)"
|
|
),
|
|
{"r": run_id},
|
|
).scalar()
|
|
assert n == _MAX_INPUT_DOCUMENTS
|
|
|
|
|
|
def test_run_row_precreated_before_execute(pg_engine, tmp_path, monkeypatch):
|
|
"""An owned workflow pre-inserts the run row keyed by the engine run id."""
|
|
_wire(pg_engine, tmp_path, monkeypatch)
|
|
wf_id = _make_workflow(pg_engine)
|
|
agent = _agent(wf_id, [])
|
|
_patch_engine(monkeypatch)
|
|
|
|
list(agent._gen_inner("go", log_context=None))
|
|
engine = _RecordingEngine.instances[-1]
|
|
|
|
with pg_engine.connect() as conn:
|
|
run = WorkflowRunsRepository(conn).get(engine.workflow_run_id)
|
|
assert run is not None
|
|
assert run["user_id"] == OWNER
|
|
assert str(run["workflow_id"]) == wf_id
|
|
# Finalized to a terminal status after the run completes.
|
|
assert run["status"] == "completed"
|
|
assert run["ended_at"] is not None
|
|
|
|
|
|
def test_unowned_workflow_creates_no_run_row(pg_engine, tmp_path, monkeypatch):
|
|
"""A draft/unowned workflow id never persists a run row and skips the bridge."""
|
|
storage = _wire(pg_engine, tmp_path, monkeypatch)
|
|
# Embedded (draft) graph whose id is NOT an owned workflow row: the run
|
|
# executes but no run row is persisted and the bridge is skipped.
|
|
attachments = [_stage_attachment(storage, b"x", "f.txt", "text/plain")]
|
|
agent = _agent(str(uuid.uuid4()), attachments)
|
|
agent._workflow_data = _EMBEDDED_GRAPH
|
|
_patch_engine(monkeypatch)
|
|
|
|
list(agent._gen_inner("go", log_context=None))
|
|
engine = _RecordingEngine.instances[-1]
|
|
|
|
with pg_engine.connect() as conn:
|
|
run = WorkflowRunsRepository(conn).get(engine.workflow_run_id)
|
|
# No bridged artifacts either (would be orphaned without a parent row).
|
|
n = conn.execute(
|
|
text(
|
|
"SELECT count(*) FROM artifacts WHERE workflow_run_id = CAST(:r AS uuid)"
|
|
),
|
|
{"r": engine.workflow_run_id},
|
|
).scalar()
|
|
assert run is None
|
|
assert n == 0
|
|
assert engine.captured_inputs["input_documents"] == []
|
|
|
|
|
|
def test_no_attachments_run_still_works(pg_engine, tmp_path, monkeypatch):
|
|
"""A run with no attachments produces empty input_documents and no artifacts."""
|
|
_wire(pg_engine, tmp_path, monkeypatch)
|
|
wf_id = _make_workflow(pg_engine)
|
|
agent = _agent(wf_id, [])
|
|
_patch_engine(monkeypatch)
|
|
|
|
list(agent._gen_inner("go", log_context=None))
|
|
engine = _RecordingEngine.instances[-1]
|
|
|
|
assert engine.captured_inputs["input_documents"] == []
|
|
with pg_engine.connect() as conn:
|
|
run = WorkflowRunsRepository(conn).get(engine.workflow_run_id)
|
|
n = conn.execute(
|
|
text(
|
|
"SELECT count(*) FROM artifacts WHERE workflow_run_id = CAST(:r AS uuid)"
|
|
),
|
|
{"r": engine.workflow_run_id},
|
|
).scalar()
|
|
assert run is not None
|
|
assert n == 0
|