mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-03 09:12:55 +00:00
Add document I/O to workflows and surface run artifacts
Let workflow runs consume and produce documents end to end: bridge uploaded attachments into run-scoped artifacts so nodes receive the input documents (with a per-run cap and server-computed size/sha256, and the run row pre-created so produced artifacts are authorized during the run); emit the run id to the client and add a builder panel that lists, previews, and downloads a run's artifacts; and allow attaching documents to a Preview run via the existing upload flow. Also fixes issues a compliance workflow surfaced: attachment ownership now keys on the raw identity instead of a sanitized one (the sanitized form could not be read back and could collide across users); workflow code nodes read prior state from a state.json data file instead of templating it into the program, so untrusted document content can never be interpolated into executed code; structured node output wrapped in code fences is recovered; and the live speech-to-text ownership check compares the raw identity.
This commit is contained in:
1 parent
396eb92ab0
commit
5a87fa6258
20 files changed
+1186
-47
No files matched your search
@@ -1,6 +1,6 @@
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, Generator, Optional
|
||||
from typing import Any, Dict, Generator, List, Optional
|
||||
|
||||
from application.agents.base import BaseAgent
|
||||
from application.agents.workflows.schemas import (
|
||||
@@ -22,6 +22,10 @@ from application.storage.db.session import db_readonly, db_session
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Per-run cap on attachments staged as run-scoped artifacts; the remainder is
|
||||
# dropped (the per-user artifact quota is only a best-effort soft cap).
|
||||
_MAX_INPUT_DOCUMENTS = 25
|
||||
|
||||
|
||||
class WorkflowAgent(BaseAgent):
|
||||
"""A specialized agent that executes predefined workflows."""
|
||||
@@ -39,6 +43,7 @@ class WorkflowAgent(BaseAgent):
|
||||
self.workflow_owner = workflow_owner
|
||||
self._workflow_data = workflow
|
||||
self._engine: Optional[WorkflowEngine] = None
|
||||
self._run_persisted = False
|
||||
|
||||
@log_activity()
|
||||
def gen(
|
||||
@@ -54,8 +59,15 @@ class WorkflowAgent(BaseAgent):
|
||||
yield {"type": "error", "error": "Failed to load workflow configuration."}
|
||||
return
|
||||
self._engine = WorkflowEngine(graph, self)
|
||||
yield from self._engine.execute({}, query)
|
||||
self._save_workflow_run(query)
|
||||
|
||||
owner_id = self._resolve_owner_id()
|
||||
pg_workflow_id = self._precreate_workflow_run(owner_id, query)
|
||||
self._run_persisted = pg_workflow_id is not None
|
||||
|
||||
input_documents = self._bridge_attachments(owner_id, persisted=self._run_persisted)
|
||||
|
||||
yield from self._engine.execute({"input_documents": input_documents}, query)
|
||||
self._finalize_workflow_run(owner_id, pg_workflow_id, query)
|
||||
|
||||
def _load_workflow_graph(self) -> Optional[WorkflowGraph]:
|
||||
if self._workflow_data:
|
||||
@@ -180,12 +192,126 @@ class WorkflowAgent(BaseAgent):
|
||||
logger.error(f"Failed to load workflow from database: {e}")
|
||||
return None
|
||||
|
||||
def _save_workflow_run(self, query: str) -> None:
|
||||
if not self._engine:
|
||||
return
|
||||
def _resolve_owner_id(self) -> Optional[str]:
|
||||
"""Resolve the run owner from the explicit workflow owner or the token ``sub``."""
|
||||
owner_id = self.workflow_owner
|
||||
if not owner_id and isinstance(self.decoded_token, dict):
|
||||
owner_id = self.decoded_token.get("sub")
|
||||
return owner_id
|
||||
|
||||
def _resolve_owned_workflow_pg_id(
|
||||
self, conn: Any, owner_id: Optional[str]
|
||||
) -> Optional[str]:
|
||||
"""Return the owned workflow's PG id, or None for an unowned/draft id."""
|
||||
if not self.workflow_id or not owner_id:
|
||||
return None
|
||||
wf_repo = WorkflowsRepository(conn)
|
||||
if looks_like_uuid(self.workflow_id):
|
||||
workflow_row = wf_repo.get(self.workflow_id, owner_id)
|
||||
else:
|
||||
workflow_row = wf_repo.get_by_legacy_id(self.workflow_id, owner_id)
|
||||
return str(workflow_row["id"]) if workflow_row is not None else None
|
||||
|
||||
def _precreate_workflow_run(self, owner_id: Optional[str], query: str) -> Optional[str]:
|
||||
"""Insert the run row up front so run-scoped artifacts are authz-reachable mid-run."""
|
||||
if not self._engine or not self.workflow_id or not owner_id:
|
||||
return None
|
||||
try:
|
||||
with db_session() as conn:
|
||||
pg_workflow_id = self._resolve_owned_workflow_pg_id(conn, owner_id)
|
||||
if pg_workflow_id is None:
|
||||
return None
|
||||
WorkflowRunsRepository(conn).create(
|
||||
pg_workflow_id,
|
||||
owner_id,
|
||||
ExecutionStatus.RUNNING.value,
|
||||
run_id=self._engine.workflow_run_id,
|
||||
inputs={"query": query},
|
||||
started_at=datetime.now(timezone.utc),
|
||||
)
|
||||
return pg_workflow_id
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to pre-create workflow run: {e}")
|
||||
return None
|
||||
|
||||
def _bridge_attachments(
|
||||
self, owner_id: Optional[str], *, persisted: bool
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Stage uploaded attachments as run-scoped artifacts the nodes can read.
|
||||
|
||||
Bytes are read server-side from each attachment's ``upload_path`` and
|
||||
re-persisted through ``persist_new_artifact`` (size/sha256/storage key all
|
||||
derived server-side); only the resulting references enter the run state.
|
||||
"""
|
||||
if not self._engine or not self.attachments or not owner_id:
|
||||
return []
|
||||
# Without a persisted run row the artifacts would be orphaned (no authz
|
||||
# parent), so skip the bridge for unowned/draft ids.
|
||||
if not persisted:
|
||||
return []
|
||||
from application.sandbox.artifacts_capture import persist_new_artifact
|
||||
from application.storage.storage_creator import StorageCreator
|
||||
|
||||
storage = StorageCreator.get_storage()
|
||||
if len(self.attachments) > _MAX_INPUT_DOCUMENTS:
|
||||
dropped = len(self.attachments) - _MAX_INPUT_DOCUMENTS
|
||||
logger.warning(
|
||||
"Workflow run input documents exceed cap (%d); dropping %d attachment(s)",
|
||||
_MAX_INPUT_DOCUMENTS,
|
||||
dropped,
|
||||
)
|
||||
refs: List[Dict[str, Any]] = []
|
||||
for index, attachment in enumerate(self.attachments[:_MAX_INPUT_DOCUMENTS]):
|
||||
upload_path = attachment.get("upload_path") or attachment.get("path")
|
||||
if not upload_path:
|
||||
continue
|
||||
filename = attachment.get("filename") or "attachment"
|
||||
mime_type = attachment.get("mime_type") or "application/octet-stream"
|
||||
attachment_id = attachment.get("id", index)
|
||||
try:
|
||||
data = storage.get_file(upload_path).read()
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"Failed to read attachment %s for workflow run: %s",
|
||||
attachment_id,
|
||||
type(exc).__name__,
|
||||
)
|
||||
continue
|
||||
try:
|
||||
ref = persist_new_artifact(
|
||||
user_id=owner_id,
|
||||
kind="file",
|
||||
data=data,
|
||||
filename=filename,
|
||||
mime_type=mime_type,
|
||||
title=filename,
|
||||
workflow_run_id=self._engine.workflow_run_id,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"Failed to persist attachment %s artifact: %s",
|
||||
attachment_id,
|
||||
type(exc).__name__,
|
||||
)
|
||||
continue
|
||||
if ref is None:
|
||||
continue
|
||||
refs.append(
|
||||
{
|
||||
"artifact_id": ref["artifact_id"],
|
||||
"ref": ref.get("ref"),
|
||||
"filename": ref["filename"],
|
||||
"mime_type": ref["mime_type"],
|
||||
}
|
||||
)
|
||||
return refs
|
||||
|
||||
def _finalize_workflow_run(
|
||||
self, owner_id: Optional[str], pg_workflow_id: Optional[str], query: str
|
||||
) -> None:
|
||||
"""Write the run's terminal status/result; upsert the row if pre-creation was skipped."""
|
||||
if not self._engine:
|
||||
return
|
||||
try:
|
||||
run = WorkflowRun(
|
||||
workflow_id=self.workflow_id or "unknown",
|
||||
@@ -197,32 +323,44 @@ class WorkflowAgent(BaseAgent):
|
||||
created_at=datetime.now(timezone.utc),
|
||||
completed_at=datetime.now(timezone.utc),
|
||||
)
|
||||
steps_json = [step.model_dump(mode="json") for step in run.steps]
|
||||
|
||||
if not self.workflow_id or not owner_id:
|
||||
return
|
||||
with db_session() as conn:
|
||||
wf_repo = WorkflowsRepository(conn)
|
||||
if looks_like_uuid(self.workflow_id):
|
||||
workflow_row = wf_repo.get(self.workflow_id, owner_id)
|
||||
else:
|
||||
workflow_row = wf_repo.get_by_legacy_id(
|
||||
self.workflow_id, owner_id,
|
||||
if pg_workflow_id is None:
|
||||
pg_workflow_id = self._resolve_owned_workflow_pg_id(conn, owner_id)
|
||||
if pg_workflow_id is None:
|
||||
return
|
||||
runs_repo = WorkflowRunsRepository(conn)
|
||||
updated = False
|
||||
if self._run_persisted:
|
||||
updated = runs_repo.finalize(
|
||||
self._engine.workflow_run_id,
|
||||
owner_id,
|
||||
run.status.value,
|
||||
result=run.outputs,
|
||||
steps=steps_json,
|
||||
ended_at=run.completed_at,
|
||||
)
|
||||
if not updated:
|
||||
logger.warning(
|
||||
"Workflow run %s finalize matched no row; "
|
||||
"recovering via insert so terminal data is not lost",
|
||||
self._engine.workflow_run_id,
|
||||
)
|
||||
if not self._run_persisted or not updated:
|
||||
runs_repo.create(
|
||||
pg_workflow_id,
|
||||
owner_id,
|
||||
run.status.value,
|
||||
run_id=self._engine.workflow_run_id,
|
||||
inputs=run.inputs,
|
||||
result=run.outputs,
|
||||
steps=steps_json,
|
||||
started_at=run.created_at,
|
||||
ended_at=run.completed_at,
|
||||
)
|
||||
if workflow_row is None:
|
||||
return
|
||||
WorkflowRunsRepository(conn).create(
|
||||
str(workflow_row["id"]),
|
||||
owner_id,
|
||||
run.status.value,
|
||||
# Persist under the engine's run id so any run-scoped
|
||||
# artifacts produced by code nodes resolve their parent.
|
||||
run_id=self._engine.workflow_run_id,
|
||||
inputs=run.inputs,
|
||||
result=run.outputs,
|
||||
steps=[step.model_dump(mode="json") for step in run.steps],
|
||||
started_at=run.created_at,
|
||||
ended_at=run.completed_at,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save workflow run: {e}")
|
||||
|
||||
|
||||
@@ -71,6 +71,11 @@ class WorkflowEngine:
|
||||
) -> Generator[Dict[str, str], None, None]:
|
||||
self._initialize_state(initial_inputs, query)
|
||||
|
||||
# Surface the run id up front so the client can list this run's
|
||||
# artifacts (GET /api/artifacts?workflow_run_id=) once it has been
|
||||
# persisted; the same id parents every artifact produced by code nodes.
|
||||
yield {"type": "workflow_run", "workflow_run_id": self.workflow_run_id}
|
||||
|
||||
start_node = self.graph.get_start_node()
|
||||
if not start_node:
|
||||
yield {"type": "error", "error": "No start node found in workflow."}
|
||||
@@ -348,6 +353,9 @@ class WorkflowEngine:
|
||||
code = config.code or ""
|
||||
if not code.strip():
|
||||
raise ValueError(f'Code node "{node.title}" has no code to execute.')
|
||||
# Code nodes are NEVER Jinja-rendered: state is untrusted (document-derived)
|
||||
# so interpolating it into the program would be code injection. Prior state is
|
||||
# passed as DATA via ``state.json`` (read below), never templated into code.
|
||||
|
||||
user_id = self._resolve_user_id()
|
||||
if not user_id:
|
||||
@@ -361,6 +369,12 @@ class WorkflowEngine:
|
||||
manager.open(session_id)
|
||||
try:
|
||||
loaded = self._materialize_code_inputs(manager, session_id, config.inputs, user_id)
|
||||
# Stage prior state as DATA the node code reads with
|
||||
# ``json.load(open("state.json"))`` -- e.g. ``state["decision"]``. The
|
||||
# file lands at the workspace root, which is the kernel cwd, so a
|
||||
# relative open resolves it. State is never templated into the program.
|
||||
state_json = json.dumps(self._json_safe_state(), default=str).encode("utf-8")
|
||||
manager.put_file(session_id, "state.json", state_json)
|
||||
pre_signatures = snapshot_signatures(manager, session_id)
|
||||
result = manager.exec(session_id, code, timeout=timeout)
|
||||
artifacts = capture_artifacts(
|
||||
@@ -478,6 +492,18 @@ class WorkflowEngine:
|
||||
"""Sanitize the run id into the sandbox-gateway charset for the session key."""
|
||||
return _SESSION_ID_RE.sub("-", str(self.workflow_run_id)) or str(uuid.uuid4())
|
||||
|
||||
def _json_safe_state(self) -> Dict[str, Any]:
|
||||
"""Project ``self.state`` to a JSON-safe dict (the code node reads it from state.json)."""
|
||||
projection: Dict[str, Any] = {}
|
||||
for key, value in self.state.items():
|
||||
if not isinstance(key, str):
|
||||
continue
|
||||
normalized_key = key.strip()
|
||||
if not normalized_key:
|
||||
continue
|
||||
projection[normalized_key] = value
|
||||
return projection
|
||||
|
||||
def _resolve_user_id(self) -> Optional[str]:
|
||||
"""Resolve the run's owner for artifact ownership/quota accounting."""
|
||||
user_id = getattr(self.agent, "user", None)
|
||||
@@ -548,10 +574,36 @@ class WorkflowEngine:
|
||||
try:
|
||||
return True, json.loads(normalized_response)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(
|
||||
"Workflow agent returned structured output that was not valid JSON"
|
||||
)
|
||||
return False, None
|
||||
pass
|
||||
|
||||
# Some models wrap structured output in a ```json ... ``` fence or add
|
||||
# prose around it; recover the JSON object/array before giving up so a
|
||||
# well-formed-but-fenced response still validates.
|
||||
candidate = self._strip_json_fence(normalized_response)
|
||||
if candidate is not None:
|
||||
try:
|
||||
return True, json.loads(candidate)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
logger.warning(
|
||||
"Workflow agent returned structured output that was not valid JSON"
|
||||
)
|
||||
return False, None
|
||||
|
||||
@staticmethod
|
||||
def _strip_json_fence(text: str) -> Optional[str]:
|
||||
"""Extract the JSON payload from a fenced/prose-wrapped response, or None."""
|
||||
fence = re.search(r"```(?:json)?\s*(.*?)\s*```", text, re.DOTALL)
|
||||
if fence:
|
||||
return fence.group(1).strip()
|
||||
# Fall back to the outermost {...} or [...] span.
|
||||
for open_ch, close_ch in (("{", "}"), ("[", "]")):
|
||||
start = text.find(open_ch)
|
||||
end = text.rfind(close_ch)
|
||||
if start != -1 and end > start:
|
||||
return text[start : end + 1]
|
||||
return None
|
||||
|
||||
def _normalize_node_json_schema(
|
||||
self, schema: Optional[Dict[str, Any]], node_title: str
|
||||
|
||||
@@ -806,6 +806,13 @@ class StreamProcessor:
|
||||
self.agent_config["workflow"] = self.data["workflow"]
|
||||
if isinstance(self.decoded_token, dict):
|
||||
self.agent_config["workflow_owner"] = self.decoded_token.get("sub")
|
||||
# A saved workflow id alongside the embedded graph (builder
|
||||
# Preview) lets the run persist a ``workflow_runs`` row so its
|
||||
# artifacts are listable + authz'd; ownership is re-checked on
|
||||
# save, so a forged id for another user's workflow never persists.
|
||||
preview_workflow_id = self.data.get("workflow_id")
|
||||
if preview_workflow_id:
|
||||
self.agent_config["workflow_id"] = str(preview_workflow_id)
|
||||
|
||||
self.agent_config.update(
|
||||
{
|
||||
@@ -1658,6 +1665,12 @@ class StreamProcessor:
|
||||
agent_kwargs["workflow_id"] = workflow_config
|
||||
elif isinstance(workflow_config, dict):
|
||||
agent_kwargs["workflow"] = workflow_config
|
||||
# Embedded-graph Preview run that names a saved workflow: run the
|
||||
# canvas graph but persist the run under the saved id so artifacts
|
||||
# parent to a real, ownership-checked ``workflow_runs`` row.
|
||||
saved_workflow_id = self.agent_config.get("workflow_id")
|
||||
if saved_workflow_id:
|
||||
agent_kwargs["workflow_id"] = saved_workflow_id
|
||||
workflow_owner = self.agent_config.get("workflow_owner")
|
||||
if workflow_owner:
|
||||
agent_kwargs["workflow_owner"] = workflow_owner
|
||||
|
||||
@@ -47,8 +47,14 @@ def _resolve_authenticated_user():
|
||||
decoded_token = getattr(request, "decoded_token", None)
|
||||
api_key = request.form.get("api_key") or request.args.get("api_key")
|
||||
|
||||
# Return the RAW user identity (not safe_filename'd): the attachment row's
|
||||
# ``user_id`` is an identity key that must match the raw ``sub`` used by
|
||||
# /stream (initial_user_id) and the artifact authz gate. safe_filename is
|
||||
# applied only to the storage path component at write time. (An email-style
|
||||
# sub like ``a@b.com`` was previously sanitized to ``abcom``, so the
|
||||
# uploaded attachment became unreadable by its own owner on /stream.)
|
||||
if decoded_token:
|
||||
return safe_filename(decoded_token.get("sub"))
|
||||
return decoded_token.get("sub")
|
||||
|
||||
if api_key:
|
||||
with db_readonly() as conn:
|
||||
@@ -57,7 +63,7 @@ def _resolve_authenticated_user():
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Invalid API key"}), 401
|
||||
)
|
||||
return safe_filename(agent.get("user_id"))
|
||||
return agent.get("user_id")
|
||||
|
||||
return None
|
||||
|
||||
@@ -162,7 +168,12 @@ class StoreAttachment(Resource):
|
||||
attachment_id = uuid.uuid4()
|
||||
original_filename = safe_filename(os.path.basename(file.filename))
|
||||
_enforce_uploaded_audio_size_limit(file, original_filename)
|
||||
relative_path = f"{settings.UPLOAD_FOLDER}/{user}/attachments/{str(attachment_id)}/{original_filename}"
|
||||
# safe_filename only the path component; the DB user_id stays raw.
|
||||
path_user = safe_filename(user)
|
||||
relative_path = (
|
||||
f"{settings.UPLOAD_FOLDER}/{path_user}/attachments/"
|
||||
f"{str(attachment_id)}/{original_filename}"
|
||||
)
|
||||
|
||||
metadata = storage.save_file(file, relative_path)
|
||||
file_info = {
|
||||
@@ -422,7 +433,10 @@ class LiveSpeechToTextChunk(Resource):
|
||||
404,
|
||||
)
|
||||
|
||||
if safe_filename(str(session_state.get("user", ""))) != auth_user:
|
||||
# The stored ``user`` is the RAW sub (see _resolve_authenticated_user), and
|
||||
# auth_user is also raw, so compare raw-vs-raw -- never safe_filename'd, or an
|
||||
# email-style sub (a@b.com) would 403 the owner on their own session.
|
||||
if str(session_state.get("user", "")) != auth_user:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Forbidden"}),
|
||||
403,
|
||||
@@ -592,7 +606,9 @@ class LiveSpeechToTextFinish(Resource):
|
||||
404,
|
||||
)
|
||||
|
||||
if safe_filename(str(session_state.get("user", ""))) != auth_user:
|
||||
# Stored ``user`` and auth_user are both RAW subs; compare raw-vs-raw so an
|
||||
# email-style sub owner is not 403'd on their own finish call.
|
||||
if str(session_state.get("user", "")) != auth_user:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Forbidden"}),
|
||||
403,
|
||||
|
||||
@@ -6,6 +6,7 @@ written once after workflow execution completes and never updated.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import Connection, text
|
||||
@@ -59,6 +60,33 @@ class WorkflowRunsRepository:
|
||||
res = self._conn.execute(stmt)
|
||||
return row_to_dict(res.fetchone())
|
||||
|
||||
def finalize(
|
||||
self,
|
||||
run_id: str,
|
||||
user_id: str,
|
||||
status: str,
|
||||
*,
|
||||
result: dict | None = None,
|
||||
steps: list | None = None,
|
||||
ended_at: Optional[datetime] = None,
|
||||
) -> bool:
|
||||
"""Update a pre-created run row with its terminal status/result; owner-scoped."""
|
||||
values: dict = {"status": status}
|
||||
if result is not None:
|
||||
values["result"] = result
|
||||
if steps is not None:
|
||||
values["steps"] = steps
|
||||
if ended_at is not None:
|
||||
values["ended_at"] = ended_at
|
||||
stmt = (
|
||||
workflow_runs_table.update()
|
||||
.where(workflow_runs_table.c.id == run_id)
|
||||
.where(workflow_runs_table.c.user_id == user_id)
|
||||
.values(**values)
|
||||
)
|
||||
res = self._conn.execute(stmt)
|
||||
return res.rowcount > 0
|
||||
|
||||
def get(self, run_id: str) -> Optional[dict]:
|
||||
res = self._conn.execute(
|
||||
text("SELECT * FROM workflow_runs WHERE id = CAST(:id AS uuid)"),
|
||||
|
||||
@@ -2895,6 +2895,7 @@ function WorkflowBuilderInner() {
|
||||
className="bg-card w-full max-w-none p-0 sm:max-w-[600px] md:max-w-[700px] lg:max-w-[800px]"
|
||||
>
|
||||
<WorkflowPreview
|
||||
workflowId={workflowId}
|
||||
workflowData={{
|
||||
name: workflowName,
|
||||
description: workflowDescription,
|
||||
|
||||
@@ -4,6 +4,7 @@ import {
|
||||
Circle,
|
||||
Code2,
|
||||
Database,
|
||||
FileBox,
|
||||
Flag,
|
||||
GitBranch,
|
||||
Loader2,
|
||||
@@ -25,6 +26,7 @@ import ConversationBubble from '../../conversation/ConversationBubble';
|
||||
import { Query } from '../../conversation/conversationModels';
|
||||
import { AppDispatch } from '../../store';
|
||||
import { WorkflowEdge, WorkflowNode } from '../types/workflow';
|
||||
import WorkflowRunArtifacts from './WorkflowRunArtifacts';
|
||||
import {
|
||||
addQuery,
|
||||
fetchWorkflowPreviewAnswer,
|
||||
@@ -48,6 +50,9 @@ interface WorkflowData {
|
||||
|
||||
interface WorkflowPreviewProps {
|
||||
workflowData: WorkflowData;
|
||||
// Saved workflow id (when the draft has been persisted); enables run-artifact
|
||||
// listing by persisting a ``workflow_runs`` row for the preview run.
|
||||
workflowId?: string | null;
|
||||
}
|
||||
|
||||
const NODE_ICONS: Record<string, React.ReactNode> = {
|
||||
@@ -241,6 +246,54 @@ function ExecutionDetails({
|
||||
);
|
||||
}
|
||||
|
||||
function RunArtifactsSection({
|
||||
workflowRunId,
|
||||
isOpen,
|
||||
onToggle,
|
||||
}: {
|
||||
workflowRunId: string;
|
||||
isOpen: boolean;
|
||||
onToggle: () => void;
|
||||
}) {
|
||||
return (
|
||||
<div className="mb-4 flex w-full flex-col flex-wrap items-start self-start lg:flex-nowrap">
|
||||
<div className="my-2 flex flex-row items-center justify-center gap-3">
|
||||
<div className="flex h-[26px] w-[30px] items-center justify-center">
|
||||
<FileBox className="h-5 w-5 text-gray-600 dark:text-gray-400" />
|
||||
</div>
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
onClick={onToggle}
|
||||
className="h-auto gap-2 px-0 py-0 hover:bg-transparent"
|
||||
>
|
||||
<p className="text-base font-semibold">Artifacts</p>
|
||||
<img
|
||||
src={ChevronDownIcon}
|
||||
alt="ChevronDown"
|
||||
className={cn(
|
||||
'h-4 w-4 transform transition-transform duration-200 dark:invert',
|
||||
isOpen ? 'rotate-180' : '',
|
||||
)}
|
||||
/>
|
||||
</Button>
|
||||
</div>
|
||||
<div
|
||||
className={cn(
|
||||
'ml-3 grid w-full transition-all duration-300 ease-in-out',
|
||||
isOpen ? 'grid-rows-[1fr] opacity-100' : 'grid-rows-[0fr] opacity-0',
|
||||
)}
|
||||
>
|
||||
<div className="overflow-hidden">
|
||||
<div className="max-h-[480px] overflow-y-auto pr-2">
|
||||
{isOpen && <WorkflowRunArtifacts workflowRunId={workflowRunId} />}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
function WorkflowMiniMap({
|
||||
nodes,
|
||||
activeNodeId,
|
||||
@@ -376,6 +429,7 @@ function WorkflowMiniMap({
|
||||
|
||||
export default function WorkflowPreview({
|
||||
workflowData,
|
||||
workflowId,
|
||||
}: WorkflowPreviewProps) {
|
||||
const dispatch = useDispatch<AppDispatch>();
|
||||
|
||||
@@ -386,6 +440,9 @@ export default function WorkflowPreview({
|
||||
|
||||
const [lastQueryReturnedErr, setLastQueryReturnedErr] = useState(false);
|
||||
const [openDetailsIndex, setOpenDetailsIndex] = useState<number | null>(null);
|
||||
const [openArtifactsIndex, setOpenArtifactsIndex] = useState<number | null>(
|
||||
null,
|
||||
);
|
||||
|
||||
const fetchStream = useRef<{ abort: () => void } | null>(null);
|
||||
const stepRefs = useRef<Map<string, HTMLDivElement>>(new Map());
|
||||
@@ -414,11 +471,12 @@ export default function WorkflowPreview({
|
||||
question,
|
||||
workflowData,
|
||||
indx: index,
|
||||
workflowId,
|
||||
}),
|
||||
);
|
||||
fetchStream.current = promise;
|
||||
},
|
||||
[dispatch, workflowData],
|
||||
[dispatch, workflowData, workflowId],
|
||||
);
|
||||
|
||||
const handleQuestion = useCallback(
|
||||
@@ -597,6 +655,20 @@ export default function WorkflowPreview({
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* Run artifacts (only once a persisted run id is known
|
||||
and the run is no longer streaming) */}
|
||||
{query.workflowRunId && !isStreamingLastQuery && (
|
||||
<RunArtifactsSection
|
||||
workflowRunId={query.workflowRunId}
|
||||
isOpen={openArtifactsIndex === index}
|
||||
onToggle={() =>
|
||||
setOpenArtifactsIndex(
|
||||
openArtifactsIndex === index ? null : index,
|
||||
)
|
||||
}
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* Response bubble */}
|
||||
{(query.response ||
|
||||
shouldShowThought ||
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
import { ChevronLeft, FileBox } from 'lucide-react';
|
||||
import { useCallback, useEffect, useState } from 'react';
|
||||
import { useSelector } from 'react-redux';
|
||||
|
||||
import userService from '../../api/services/userService';
|
||||
import DocumentArtifactView from '../../components/DocumentArtifactView';
|
||||
import Spinner from '../../components/Spinner';
|
||||
import {
|
||||
isDocumentArtifact,
|
||||
type DocumentArtifact,
|
||||
} from '../../components/artifactViewUtils';
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { selectToken } from '../../preferences/preferenceSlice';
|
||||
|
||||
interface RunArtifactSummary {
|
||||
id: string;
|
||||
kind: string | null;
|
||||
title: string | null;
|
||||
current_version: number | null;
|
||||
}
|
||||
|
||||
interface WorkflowRunArtifactsProps {
|
||||
workflowRunId: string;
|
||||
}
|
||||
|
||||
/** List a workflow run's produced artifacts, with click-through preview + download. */
|
||||
export default function WorkflowRunArtifacts({
|
||||
workflowRunId,
|
||||
}: WorkflowRunArtifactsProps) {
|
||||
const token = useSelector(selectToken);
|
||||
const [artifacts, setArtifacts] = useState<RunArtifactSummary[] | null>(null);
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
const [selectedId, setSelectedId] = useState<string | null>(null);
|
||||
|
||||
const [detail, setDetail] = useState<DocumentArtifact | null>(null);
|
||||
const [detailLoading, setDetailLoading] = useState(false);
|
||||
const [detailError, setDetailError] = useState<string | null>(null);
|
||||
|
||||
const loadList = useCallback(() => {
|
||||
if (!workflowRunId) return;
|
||||
let cancelled = false;
|
||||
setLoading(true);
|
||||
setError(null);
|
||||
userService
|
||||
.listWorkflowRunArtifacts(workflowRunId, token)
|
||||
.then(async (res: Response) => {
|
||||
if (cancelled) return;
|
||||
if (!res.ok) {
|
||||
// The run row is missing/unauthorized (e.g. an unsaved-draft preview):
|
||||
// show an informational empty state rather than a hard error.
|
||||
setArtifacts([]);
|
||||
setLoading(false);
|
||||
return;
|
||||
}
|
||||
const data = await res.json().catch(() => null);
|
||||
if (cancelled) return;
|
||||
setArtifacts(data?.success ? (data.artifacts ?? []) : []);
|
||||
setLoading(false);
|
||||
})
|
||||
.catch(() => {
|
||||
if (cancelled) return;
|
||||
setError('Failed to load artifacts');
|
||||
setLoading(false);
|
||||
});
|
||||
return () => {
|
||||
cancelled = true;
|
||||
};
|
||||
}, [workflowRunId, token]);
|
||||
|
||||
useEffect(() => {
|
||||
const cleanup = loadList();
|
||||
return cleanup;
|
||||
}, [loadList]);
|
||||
|
||||
const fetchDetail = useCallback(
|
||||
(artifactId: string) => {
|
||||
let cancelled = false;
|
||||
setDetailLoading(true);
|
||||
setDetailError(null);
|
||||
userService
|
||||
.getDocumentArtifact(artifactId, token)
|
||||
.then(async (res: Response) => {
|
||||
if (cancelled) return;
|
||||
if (!res.ok) {
|
||||
setDetailError('Failed to load artifact');
|
||||
setDetailLoading(false);
|
||||
return;
|
||||
}
|
||||
const data = await res.json().catch(() => null);
|
||||
if (cancelled) return;
|
||||
if (data?.success && isDocumentArtifact(data.artifact)) {
|
||||
setDetail(data.artifact);
|
||||
setDetailLoading(false);
|
||||
} else {
|
||||
setDetailError('This artifact cannot be previewed');
|
||||
setDetailLoading(false);
|
||||
}
|
||||
})
|
||||
.catch(() => {
|
||||
if (cancelled) return;
|
||||
setDetailError('Failed to load artifact');
|
||||
setDetailLoading(false);
|
||||
});
|
||||
return () => {
|
||||
cancelled = true;
|
||||
};
|
||||
},
|
||||
[token],
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
if (!selectedId) {
|
||||
setDetail(null);
|
||||
setDetailError(null);
|
||||
return;
|
||||
}
|
||||
const cleanup = fetchDetail(selectedId);
|
||||
return cleanup;
|
||||
}, [selectedId, fetchDetail]);
|
||||
|
||||
if (loading) {
|
||||
return (
|
||||
<div className="flex items-center gap-2 px-3 py-4 text-sm text-gray-500 dark:text-gray-400">
|
||||
<Spinner size="small" /> Loading artifacts...
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (error) {
|
||||
return (
|
||||
<div className="flex items-center justify-between gap-2 px-3 py-3 text-sm text-red-500">
|
||||
<span>{error}</span>
|
||||
<Button type="button" variant="outline" size="sm" onClick={loadList}>
|
||||
Retry
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (!artifacts || artifacts.length === 0) {
|
||||
return (
|
||||
<div className="px-3 py-3 text-sm text-gray-500 dark:text-gray-400">
|
||||
No artifacts produced by this run.
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (selectedId) {
|
||||
return (
|
||||
<div className="flex h-full min-h-0 flex-col">
|
||||
<div className="mb-2 flex items-center">
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="gap-1 px-2"
|
||||
onClick={() => setSelectedId(null)}
|
||||
>
|
||||
<ChevronLeft className="h-4 w-4" />
|
||||
Back to artifacts
|
||||
</Button>
|
||||
</div>
|
||||
<div className="min-h-0 flex-1 overflow-hidden">
|
||||
{detailLoading ? (
|
||||
<div className="flex h-full items-center justify-center">
|
||||
<Spinner />
|
||||
</div>
|
||||
) : detailError ? (
|
||||
<div className="flex h-full items-center justify-center">
|
||||
<p className="text-sm text-red-500">{detailError}</p>
|
||||
</div>
|
||||
) : detail ? (
|
||||
<DocumentArtifactView
|
||||
artifact={detail}
|
||||
onRefresh={() => fetchDetail(selectedId)}
|
||||
/>
|
||||
) : null}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<ul className="space-y-2">
|
||||
{artifacts.map((artifact) => (
|
||||
<li key={artifact.id}>
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
onClick={() => setSelectedId(artifact.id)}
|
||||
className="h-auto w-full justify-start gap-3 px-3 py-2 text-left"
|
||||
>
|
||||
<FileBox className="h-4 w-4 shrink-0 text-gray-500 dark:text-gray-400" />
|
||||
<div className="min-w-0 flex-1">
|
||||
<div className="truncate text-sm font-medium text-gray-900 dark:text-white">
|
||||
{artifact.title || `Artifact ${artifact.id.slice(0, 8)}`}
|
||||
</div>
|
||||
<div className="truncate text-xs text-gray-500 dark:text-gray-400">
|
||||
{artifact.kind || 'file'}
|
||||
{artifact.current_version != null
|
||||
? ` · v${artifact.current_version}`
|
||||
: ''}
|
||||
</div>
|
||||
</div>
|
||||
</Button>
|
||||
</li>
|
||||
))}
|
||||
</ul>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
import { describe, expect, it } from 'vitest';
|
||||
|
||||
import { Attachment } from '../../upload/uploadSlice';
|
||||
import reducer, {
|
||||
addQuery,
|
||||
collectCompletedAttachmentIds,
|
||||
setWorkflowRunId,
|
||||
} from './workflowPreviewSlice';
|
||||
|
||||
const seedState = () => reducer(undefined, { type: '@@INIT' });
|
||||
|
||||
const att = (over: Partial<Attachment>): Attachment => ({
|
||||
id: 'a1',
|
||||
fileName: 'f.pdf',
|
||||
progress: 100,
|
||||
status: 'completed',
|
||||
taskId: 't1',
|
||||
...over,
|
||||
});
|
||||
|
||||
describe('collectCompletedAttachmentIds', () => {
|
||||
it('returns ids of completed attachments only', () => {
|
||||
const ids = collectCompletedAttachmentIds([
|
||||
att({ id: 'done', status: 'completed' }),
|
||||
att({ id: 'busy', status: 'processing' }),
|
||||
att({ id: 'up', status: 'uploading' }),
|
||||
att({ id: 'bad', status: 'failed' }),
|
||||
]);
|
||||
expect(ids).toEqual(['done']);
|
||||
});
|
||||
|
||||
it('drops completed rows with no server id and returns [] when none', () => {
|
||||
expect(
|
||||
collectCompletedAttachmentIds([att({ id: '', status: 'completed' })]),
|
||||
).toEqual([]);
|
||||
expect(collectCompletedAttachmentIds([])).toEqual([]);
|
||||
});
|
||||
});
|
||||
|
||||
describe('setWorkflowRunId', () => {
|
||||
it('stores the run id on the addressed query', () => {
|
||||
let state = seedState();
|
||||
state = reducer(state, addQuery({ prompt: 'run it' }));
|
||||
state = reducer(
|
||||
state,
|
||||
setWorkflowRunId({ index: 0, workflowRunId: 'run-1' }),
|
||||
);
|
||||
expect(state.queries[0].workflowRunId).toBe('run-1');
|
||||
});
|
||||
|
||||
it('ignores an out-of-range index without throwing', () => {
|
||||
let state = seedState();
|
||||
state = reducer(state, addQuery({ prompt: 'q' }));
|
||||
state = reducer(
|
||||
state,
|
||||
setWorkflowRunId({ index: 5, workflowRunId: 'run-x' }),
|
||||
);
|
||||
expect(state.queries[0].workflowRunId).toBeUndefined();
|
||||
});
|
||||
});
|
||||
@@ -2,6 +2,7 @@ import { createAsyncThunk, createSlice, PayloadAction } from '@reduxjs/toolkit';
|
||||
|
||||
import conversationService from '../../api/services/conversationService';
|
||||
import { Query, Status } from '../../conversation/conversationModels';
|
||||
import { Attachment, clearAttachments } from '../../upload/uploadSlice';
|
||||
import { WorkflowEdge, WorkflowNode } from '../types/workflow';
|
||||
|
||||
export interface WorkflowExecutionStep {
|
||||
@@ -26,6 +27,9 @@ interface WorkflowData {
|
||||
|
||||
export interface WorkflowQuery extends Query {
|
||||
executionSteps?: WorkflowExecutionStep[];
|
||||
// The run id for this query's execution; lets the artifacts panel list the
|
||||
// run's produced artifacts (GET /api/artifacts?workflow_run_id=).
|
||||
workflowRunId?: string;
|
||||
}
|
||||
|
||||
export interface WorkflowPreviewState {
|
||||
@@ -56,6 +60,16 @@ interface ThunkState {
|
||||
token: string | null;
|
||||
};
|
||||
workflowPreview: WorkflowPreviewState;
|
||||
upload: { attachments: Attachment[] };
|
||||
}
|
||||
|
||||
/** Server-assigned ids of completed attachments, for the run's ``attachments``. */
|
||||
export function collectCompletedAttachmentIds(
|
||||
attachments: Attachment[],
|
||||
): string[] {
|
||||
return attachments
|
||||
.filter((att) => att.status === 'completed' && att.id)
|
||||
.map((att) => att.id);
|
||||
}
|
||||
|
||||
export const fetchWorkflowPreviewAnswer = createAsyncThunk<
|
||||
@@ -64,11 +78,15 @@ export const fetchWorkflowPreviewAnswer = createAsyncThunk<
|
||||
question: string;
|
||||
workflowData: WorkflowData;
|
||||
indx?: number;
|
||||
workflowId?: string | null;
|
||||
},
|
||||
{ state: ThunkState }
|
||||
>(
|
||||
'workflowPreview/fetchAnswer',
|
||||
async ({ question, workflowData, indx }, { dispatch, getState }) => {
|
||||
async (
|
||||
{ question, workflowData, indx, workflowId },
|
||||
{ dispatch, getState },
|
||||
) => {
|
||||
if (abortController) abortController.abort();
|
||||
abortController = new AbortController();
|
||||
const { signal } = abortController;
|
||||
@@ -76,9 +94,20 @@ export const fetchWorkflowPreviewAnswer = createAsyncThunk<
|
||||
const state = getState();
|
||||
|
||||
if (state.preference) {
|
||||
// Doc-driven workflows read uploaded files as run-scoped artifacts; pass
|
||||
// the completed attachment ids so the backend bridges them (user-scoped).
|
||||
const attachmentIds = collectCompletedAttachmentIds(
|
||||
state.upload.attachments,
|
||||
);
|
||||
if (attachmentIds.length > 0) dispatch(clearAttachments());
|
||||
|
||||
// A saved workflow id lets the run persist a ``workflow_runs`` row so its
|
||||
// produced artifacts are listable + authz'd; omitted for unsaved drafts.
|
||||
const payload = {
|
||||
question,
|
||||
workflow: workflowData,
|
||||
...(workflowId ? { workflow_id: workflowId } : {}),
|
||||
...(attachmentIds.length > 0 ? { attachments: attachmentIds } : {}),
|
||||
save_conversation: false,
|
||||
visibility: 'hidden',
|
||||
};
|
||||
@@ -117,6 +146,13 @@ export const fetchWorkflowPreviewAnswer = createAsyncThunk<
|
||||
|
||||
if (data.type === 'end') {
|
||||
dispatch(workflowPreviewSlice.actions.setStatus('idle'));
|
||||
} else if (data.type === 'workflow_run') {
|
||||
dispatch(
|
||||
setWorkflowRunId({
|
||||
index: targetIndex,
|
||||
workflowRunId: data.workflow_run_id,
|
||||
}),
|
||||
);
|
||||
} else if (data.type === 'thought') {
|
||||
dispatch(
|
||||
updateThought({
|
||||
@@ -366,6 +402,15 @@ export const workflowPreviewSlice = createSlice({
|
||||
state.executionSteps.push(updatedStep);
|
||||
}
|
||||
},
|
||||
setWorkflowRunId(
|
||||
state,
|
||||
action: PayloadAction<{ index: number; workflowRunId: string }>,
|
||||
) {
|
||||
const { index, workflowRunId } = action.payload;
|
||||
if (state.queries[index]) {
|
||||
state.queries[index].workflowRunId = workflowRunId;
|
||||
}
|
||||
},
|
||||
setActiveNodeId(state, action: PayloadAction<string | null>) {
|
||||
state.activeNodeId = action.payload;
|
||||
},
|
||||
@@ -440,6 +485,7 @@ export const {
|
||||
updateStreamingSource,
|
||||
updateToolCall,
|
||||
updateExecutionStep,
|
||||
setWorkflowRunId,
|
||||
setActiveNodeId,
|
||||
setStatus,
|
||||
raiseError,
|
||||
|
||||
@@ -112,6 +112,8 @@ const endpoints = {
|
||||
GET_ARTIFACT: (artifactId: string) => `/api/artifact/${artifactId}`,
|
||||
GET_DOCUMENT_ARTIFACT: (artifactId: string) =>
|
||||
`/api/artifacts/${artifactId}`,
|
||||
LIST_WORKFLOW_RUN_ARTIFACTS: (workflowRunId: string) =>
|
||||
`/api/artifacts?workflow_run_id=${encodeURIComponent(workflowRunId)}`,
|
||||
DOWNLOAD_ARTIFACT: (artifactId: string, version?: number) =>
|
||||
version != null
|
||||
? `/api/artifacts/${artifactId}/download?version=${version}`
|
||||
|
||||
@@ -327,6 +327,14 @@ const userService = {
|
||||
token: string | null,
|
||||
): Promise<Response> =>
|
||||
apiClient.get(endpoints.USER.GET_DOCUMENT_ARTIFACT(artifactId), token),
|
||||
listWorkflowRunArtifacts: (
|
||||
workflowRunId: string,
|
||||
token: string | null,
|
||||
): Promise<Response> =>
|
||||
apiClient.get(
|
||||
endpoints.USER.LIST_WORKFLOW_RUN_ARTIFACTS(workflowRunId),
|
||||
token,
|
||||
),
|
||||
downloadArtifact: (
|
||||
artifactId: string,
|
||||
token: string | null,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Tests for application/agents/workflow_agent.py graph loading and saving.
|
||||
|
||||
Tests _parse_embedded_workflow, _load_from_database, and _save_workflow_run
|
||||
Tests _parse_embedded_workflow, _load_from_database, and _finalize_workflow_run
|
||||
against the ephemeral ``pg_conn`` fixture. Agent construction is bypassed via
|
||||
``__new__`` so we avoid the BaseAgent's LLM/tool wiring.
|
||||
"""
|
||||
@@ -20,6 +20,7 @@ def _make_agent(*, workflow_id=None, workflow=None, workflow_owner=None,
|
||||
agent.workflow_owner = workflow_owner
|
||||
agent._workflow_data = workflow
|
||||
agent._engine = None
|
||||
agent._run_persisted = False
|
||||
agent.decoded_token = decoded_token or {}
|
||||
return agent
|
||||
|
||||
@@ -184,7 +185,7 @@ class TestSaveWorkflowRun:
|
||||
def test_returns_when_no_engine(self):
|
||||
agent = _make_agent(workflow_id="x", workflow_owner="u")
|
||||
# _engine is None
|
||||
agent._save_workflow_run("query")
|
||||
agent._finalize_workflow_run(agent.workflow_owner, None, "query")
|
||||
# should not raise
|
||||
|
||||
def test_returns_when_no_workflow_id(self):
|
||||
@@ -193,7 +194,7 @@ class TestSaveWorkflowRun:
|
||||
agent._engine.execution_log = []
|
||||
agent._engine.state = {}
|
||||
agent._engine.get_execution_summary.return_value = []
|
||||
agent._save_workflow_run("query")
|
||||
agent._finalize_workflow_run(agent.workflow_owner, None, "query")
|
||||
|
||||
def test_returns_when_workflow_missing_in_db(self, pg_conn):
|
||||
agent = _make_agent(
|
||||
@@ -206,7 +207,7 @@ class TestSaveWorkflowRun:
|
||||
agent._engine.get_execution_summary.return_value = []
|
||||
with _patch_db(pg_conn):
|
||||
# Should just return None since workflow not found in DB
|
||||
agent._save_workflow_run("q")
|
||||
agent._finalize_workflow_run(agent.workflow_owner, None, "q")
|
||||
|
||||
def test_creates_run_row(self, pg_conn):
|
||||
from application.storage.db.repositories.workflows import (
|
||||
@@ -233,7 +234,7 @@ class TestSaveWorkflowRun:
|
||||
agent._engine.workflow_run_id = run_id
|
||||
|
||||
with _patch_db(pg_conn):
|
||||
agent._save_workflow_run("my query")
|
||||
agent._finalize_workflow_run(agent.workflow_owner, None, "my query")
|
||||
|
||||
runs = WorkflowRunsRepository(pg_conn).list_for_workflow(str(wf["id"]))
|
||||
assert len(runs) >= 1
|
||||
@@ -255,7 +256,7 @@ class TestSaveWorkflowRun:
|
||||
"application.agents.workflow_agent.db_session", _broken
|
||||
):
|
||||
# Should not raise
|
||||
agent._save_workflow_run("q")
|
||||
agent._finalize_workflow_run(agent.workflow_owner, None, "q")
|
||||
|
||||
|
||||
class TestDetermineRunStatus:
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""Tests for the Postgres dual-write path inside WorkflowAgent._save_workflow_run.
|
||||
"""Tests for the Postgres write path inside WorkflowAgent._finalize_workflow_run.
|
||||
|
||||
Specifically verifies the inner ``_pg_write`` closure that:
|
||||
1. Calls WorkflowsRepository.get_by_legacy_id() to resolve the Mongo workflow id.
|
||||
@@ -64,7 +64,7 @@ def _stub_mongo(agent, insert_id=None):
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _save_workflow_run — PG dual-write logic
|
||||
# _finalize_workflow_run — PG write logic
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ patched persistence helper (no live gateway / DB / storage), plus the
|
||||
serialization round-trip and CEL branching on an artifact reference.
|
||||
"""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
@@ -61,19 +62,25 @@ class _Result:
|
||||
|
||||
|
||||
class _FakeManager:
|
||||
"""Records open/close and returns a fixed exec result; no real sandbox."""
|
||||
"""Records open/close/put_file/exec and returns a fixed result; no real sandbox."""
|
||||
|
||||
def __init__(self, result):
|
||||
self._result = result
|
||||
self.opened = []
|
||||
self.closed = []
|
||||
self.put_files = []
|
||||
self.last_code = None
|
||||
|
||||
def open(self, session_id, ttl=None):
|
||||
self.opened.append(session_id)
|
||||
return session_id
|
||||
|
||||
def put_file(self, session_id, dest_path, data):
|
||||
self.put_files.append((dest_path, data))
|
||||
|
||||
def exec(self, session_id, code, timeout=None):
|
||||
self.last_timeout = timeout
|
||||
self.last_code = code
|
||||
return self._result
|
||||
|
||||
def close(self, session_id):
|
||||
@@ -141,6 +148,40 @@ def test_code_node_no_artifacts_still_writes_status(patch_sandbox):
|
||||
assert engine.state["out"] == {"artifacts": [], "status": "ok"}
|
||||
|
||||
|
||||
def test_code_node_reads_prior_state_from_state_json(patch_sandbox):
|
||||
# Prior state is staged as DATA in state.json (workspace root = kernel cwd) so
|
||||
# node code reads it with json.load(open("state.json")) -- e.g. state["decision"].
|
||||
engine = _engine()
|
||||
engine.state["decision"] = {"pass": True, "score": 7}
|
||||
node = _code_node(
|
||||
output_variable="out",
|
||||
code="import json\nd = json.load(open('state.json'))\nprint(d['decision'])\n",
|
||||
)
|
||||
|
||||
list(engine._execute_code_node(node))
|
||||
|
||||
manager = patch_sandbox["manager_holder"]["manager"]
|
||||
staged = dict(manager.put_files)
|
||||
assert "state.json" in staged
|
||||
payload = json.loads(staged["state.json"].decode("utf-8"))
|
||||
assert payload["decision"] == {"pass": True, "score": 7}
|
||||
|
||||
|
||||
def test_code_node_literal_braces_passed_verbatim_not_templated(patch_sandbox):
|
||||
# Proves code nodes are NOT Jinja-rendered: a literal ``{{ ... }}`` in the code
|
||||
# reaches exec() byte-for-byte (no injection path that interpolates state).
|
||||
engine = _engine()
|
||||
engine.state["decision"] = "INJECTED"
|
||||
literal = "x = '{{ agent.decision }}'\nprint(x)\n"
|
||||
node = _code_node(output_variable="out", code=literal)
|
||||
|
||||
list(engine._execute_code_node(node))
|
||||
|
||||
manager = patch_sandbox["manager_holder"]["manager"]
|
||||
assert manager.last_code == literal
|
||||
assert "INJECTED" not in manager.last_code
|
||||
|
||||
|
||||
def test_code_node_failure_raises(patch_sandbox):
|
||||
engine = _engine()
|
||||
patch_sandbox["result"] = _Result(ok=False, error_name="ValueError", error_value="boom")
|
||||
|
||||
@@ -65,6 +65,21 @@ class TestExecuteLoop:
|
||||
events = list(engine.execute({}, "query"))
|
||||
assert any(e.get("type") == "error" and "start node" in e.get("error", "") for e in events)
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_emits_workflow_run_id_first(self):
|
||||
"""A ``workflow_run`` event carrying the run id precedes any step event."""
|
||||
nodes = [
|
||||
_make_node("n1", NodeType.START, "Start"),
|
||||
_make_node("n2", NodeType.END, "End", config={"config": {}}),
|
||||
]
|
||||
edges = [_make_edge("e1", "n1", "n2")]
|
||||
engine = WorkflowEngine(_make_graph(nodes, edges), _make_agent())
|
||||
events = list(engine.execute({}, "hello"))
|
||||
run_events = [e for e in events if e.get("type") == "workflow_run"]
|
||||
assert len(run_events) == 1
|
||||
assert run_events[0]["workflow_run_id"] == engine.workflow_run_id
|
||||
assert events[0]["type"] == "workflow_run"
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_start_to_end(self):
|
||||
nodes = [
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
"""Engine robustness: fenced/prose structured output recovery in _parse_structured_output."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from application.agents.workflows.workflow_engine import WorkflowEngine
|
||||
|
||||
|
||||
def _engine() -> WorkflowEngine:
|
||||
return WorkflowEngine.__new__(WorkflowEngine)
|
||||
|
||||
|
||||
def test_parse_structured_output_bare_json():
|
||||
ok, val = _engine()._parse_structured_output('{"x": true}')
|
||||
assert ok and val == {"x": True}
|
||||
|
||||
|
||||
def test_parse_structured_output_json_fence():
|
||||
ok, val = _engine()._parse_structured_output('```json\n{"a": 1, "b": null}\n```')
|
||||
assert ok and val == {"a": 1, "b": None}
|
||||
|
||||
|
||||
def test_parse_structured_output_unlabelled_fence():
|
||||
ok, val = _engine()._parse_structured_output("text\n```\n{\"y\": [1, 2]}\n```\nmore")
|
||||
assert ok and val == {"y": [1, 2]}
|
||||
|
||||
|
||||
def test_parse_structured_output_prose_wrapped_object():
|
||||
ok, val = _engine()._parse_structured_output('Result: {"z": "v"} done')
|
||||
assert ok and val == {"z": "v"}
|
||||
|
||||
|
||||
def test_parse_structured_output_non_json():
|
||||
ok, val = _engine()._parse_structured_output("not json at all")
|
||||
assert ok is False and val is None
|
||||
|
||||
|
||||
def test_parse_structured_output_empty():
|
||||
ok, val = _engine()._parse_structured_output(" ")
|
||||
assert ok is False and val is None
|
||||
|
||||
|
||||
def test_strip_json_fence_prefers_fence_over_braces():
|
||||
out = WorkflowEngine._strip_json_fence('lead {"outer": 1} ```json\n{"in": 2}\n```')
|
||||
assert out == '{"in": 2}'
|
||||
@@ -0,0 +1,280 @@
|
||||
"""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)
|
||||
assert storage.get_file(v["storage_path"]).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
|
||||
@@ -108,6 +108,41 @@ class TestStreamProcessorAgentConfiguration:
|
||||
assert isinstance(processor.agent_config, dict)
|
||||
assert processor.agent_id is None
|
||||
|
||||
def test_embedded_workflow_without_saved_id(self):
|
||||
"""A preview run with no saved workflow id carries no ``workflow_id``."""
|
||||
from application.api.answer.services.stream_processor import StreamProcessor
|
||||
|
||||
request_data = {
|
||||
"question": "Test",
|
||||
"workflow": {"nodes": [], "edges": []},
|
||||
}
|
||||
processor = StreamProcessor(request_data, {"sub": "user_123"})
|
||||
processor._configure_agent()
|
||||
|
||||
assert processor.agent_config["agent_type"] == "workflow"
|
||||
assert processor.agent_config["workflow"] == {"nodes": [], "edges": []}
|
||||
assert "workflow_id" not in processor.agent_config
|
||||
|
||||
def test_embedded_workflow_with_saved_id_persists_run(self):
|
||||
"""A saved workflow id alongside the embedded graph is captured so the
|
||||
run can persist a ``workflow_runs`` row for artifact listing."""
|
||||
from application.api.answer.services.stream_processor import StreamProcessor
|
||||
|
||||
request_data = {
|
||||
"question": "Test",
|
||||
"workflow": {"nodes": [], "edges": []},
|
||||
"workflow_id": "11111111-1111-1111-1111-111111111111",
|
||||
}
|
||||
processor = StreamProcessor(request_data, {"sub": "user_123"})
|
||||
processor._configure_agent()
|
||||
|
||||
assert processor.agent_config["agent_type"] == "workflow"
|
||||
assert (
|
||||
processor.agent_config["workflow_id"]
|
||||
== "11111111-1111-1111-1111-111111111111"
|
||||
)
|
||||
assert processor.agent_config["workflow_owner"] == "user_123"
|
||||
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
|
||||
@@ -486,6 +486,18 @@ class TestResolveAuthenticatedUser:
|
||||
assert result is not None
|
||||
assert "jwt_user" in result
|
||||
|
||||
def test_returns_raw_email_sub_unsanitized(self, flask_app):
|
||||
# Regression: an email-style sub must be returned verbatim so the
|
||||
# attachment row's user_id matches the raw sub /stream queries with.
|
||||
# Previously safe_filename stripped "@"/"." (a@b.com -> abcom), making
|
||||
# an uploaded attachment unreadable by its own owner on /stream.
|
||||
from application.api.user.attachments.routes import _resolve_authenticated_user
|
||||
|
||||
app = Flask(__name__)
|
||||
with app.test_request_context("/api/store_attachment", method="POST"):
|
||||
request.decoded_token = {"sub": "alex@arc53.com"}
|
||||
assert _resolve_authenticated_user() == "alex@arc53.com"
|
||||
|
||||
def test_returns_user_from_valid_api_key_form(self, flask_app):
|
||||
from application.api.user.attachments.routes import _resolve_authenticated_user
|
||||
|
||||
@@ -1439,6 +1451,70 @@ class TestLiveSpeechToTextAdditional:
|
||||
response = finish_resource.post()
|
||||
assert _get_response_status(response) == 403
|
||||
|
||||
@patch("application.api.user.attachments.routes.STTCreator.create_stt")
|
||||
@patch("application.api.user.attachments.routes.get_redis_instance")
|
||||
def test_live_stt_email_sub_owner_round_trip(
|
||||
self, mock_get_redis, mock_create_stt, flask_app
|
||||
):
|
||||
# Regression: an email-style sub (a@b.com) is stored RAW on the session, so
|
||||
# the owner must pass the raw-vs-raw ownership gate on chunk + finish. Before
|
||||
# the fix, safe_filename(stored_user) != raw auth_user 403'd the owner.
|
||||
from application.api.user.attachments.routes import (
|
||||
LiveSpeechToTextChunk,
|
||||
LiveSpeechToTextFinish,
|
||||
LiveSpeechToTextStart,
|
||||
)
|
||||
|
||||
app = Flask(__name__)
|
||||
fake_redis = FakeRedis()
|
||||
mock_get_redis.return_value = fake_redis
|
||||
owner_sub = "alex@arc53.com"
|
||||
|
||||
start_resource = LiveSpeechToTextStart()
|
||||
with app.test_request_context(
|
||||
"/api/stt/live/start",
|
||||
method="POST",
|
||||
json={"language": "en"},
|
||||
):
|
||||
request.decoded_token = {"sub": owner_sub}
|
||||
start_response = start_resource.post()
|
||||
session_id = _get_response_json(start_response)["session_id"]
|
||||
|
||||
mock_stt = MagicMock()
|
||||
mock_stt.transcribe.return_value = {
|
||||
"text": "hello this is a longer test phrase for transcript stabilization today now",
|
||||
"language": "en",
|
||||
"duration_s": 1.0,
|
||||
"segments": [],
|
||||
"provider": "openai",
|
||||
}
|
||||
mock_create_stt.return_value = mock_stt
|
||||
|
||||
chunk_resource = LiveSpeechToTextChunk()
|
||||
with app.test_request_context(
|
||||
"/api/stt/live/chunk",
|
||||
method="POST",
|
||||
data={
|
||||
"session_id": session_id,
|
||||
"chunk_index": "0",
|
||||
"file": (io.BytesIO(b"chunk-0"), "chunk-0.wav"),
|
||||
},
|
||||
content_type="multipart/form-data",
|
||||
):
|
||||
request.decoded_token = {"sub": owner_sub}
|
||||
chunk_response = chunk_resource.post()
|
||||
assert _get_response_status(chunk_response) == 200
|
||||
|
||||
finish_resource = LiveSpeechToTextFinish()
|
||||
with app.test_request_context(
|
||||
"/api/stt/live/finish",
|
||||
method="POST",
|
||||
json={"session_id": session_id},
|
||||
):
|
||||
request.decoded_token = {"sub": owner_sub}
|
||||
finish_response = finish_resource.post()
|
||||
assert _get_response_status(finish_response) == 200
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestServeImage:
|
||||
|
||||
Reference in new issue
Block a user