from pathlib import Path from types import SimpleNamespace from datetime import datetime import json import tempfile import unittest from unittest.mock import AsyncMock, patch from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from backend.agent import api from backend.agent.models import WorkflowDefinition, WorkflowRun from backend.agent.schemas import WorkflowRunRequest from backend.agent.workflows.builtins import seed_builtin_workflows from backend.database import Base from backend.models import ChatMessage, ChatSession class WorkflowApiTests(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self): self.temp = tempfile.TemporaryDirectory() self.engine = create_engine( f"sqlite:///{Path(self.temp.name) / 'api.db'}", connect_args={"check_same_thread": False}, ) Base.metadata.create_all(self.engine) self.Session = sessionmaker(bind=self.engine) db = self.Session() seed_builtin_workflows(db) direct = db.query(WorkflowDefinition).filter_by(slug="input-output").one() chat = ChatSession(session_id="regenerate-chat") db.add(chat) db.flush() user = ChatMessage(session_pk=chat.id, role="user", content="original") db.add(user) db.flush() assistant = ChatMessage( session_pk=chat.id, role="assistant", content="old answer", workflow_id=direct.id, workflow_revision_id=direct.current_revision_id, ) db.add(assistant) db.flush() db.commit() self.user_message_id = user.message_id self.assistant_message_id = assistant.message_id self.direct_workflow_id = direct.id self.direct_revision_id = direct.current_revision_id db.close() async def asyncTearDown(self): self.engine.dispose() self.temp.cleanup() async def test_regeneration_reuses_user_row_and_prunes_old_answer(self): started = [] fake_runtime = SimpleNamespace(start=started.append) request = WorkflowRunRequest( session_id="regenerate-chat", message="edited", model="test-model", selection_mode="direct", workflow_revision_id=self.direct_revision_id, regenerate_index=0, ) with patch.object(api, "SessionLocal", self.Session), patch.object(api, "runtime", fake_runtime): response = await api.create_workflow_run(request) self.assertEqual(started, [response["run_id"]]) db = self.Session() rows = db.query(ChatMessage).order_by(ChatMessage.id.asc()).all() self.assertEqual(len(rows), 1) self.assertEqual(rows[0].message_id, self.user_message_id) self.assertEqual(rows[0].content, "edited") run = db.query(WorkflowRun).filter_by(id=response["run_id"]).one() self.assertEqual(run.user_message_id, self.user_message_id) db.close() async def test_run_resolves_previous_assistant_and_message_url_targets(self): fake_runtime = SimpleNamespace(start=lambda _run_id: None) request = WorkflowRunRequest( session_id="regenerate-chat", message="Save https://example.com/source for later.", model="test-model", selection_mode="explicit", workflow_id=self.direct_workflow_id, ) with patch.object(api, "SessionLocal", self.Session), patch.object(api, "runtime", fake_runtime): response = await api.create_workflow_run(request) db = self.Session() run = db.query(WorkflowRun).filter_by(id=response["run_id"]).one() run_values = json.loads(run.inputs_json)["run"] self.assertEqual(run_values["target_message_id"], self.assistant_message_id) self.assertEqual(run_values["target_url"], "https://example.com/source") db.close() async def test_automatic_run_persists_router_web_search_inputs(self): db = self.Session() web = db.query(WorkflowDefinition).filter_by(slug="web-answer").one() web_id = web.id db.close() started = [] fake_runtime = SimpleNamespace(start=started.append) router_result = { "workflow_id": web_id, "workflow_slug": "web-answer", "confidence": 0.95, "reason": "current information", "inputs": {"web_search_query": "Iran news today", "web_search_queries": ["Iran news today", "Iran latest news"]}, "fallback": False, } request = WorkflowRunRequest( session_id=None, message="what happened in iran today", model="test-model", selection_mode="automatic", ) with ( patch.object(api, "SessionLocal", self.Session), patch.object(api, "runtime", fake_runtime), patch.object(api, "select_workflow", new=AsyncMock(return_value=router_result)), ): response = await api.create_workflow_run(request) self.assertEqual(started, [response["run_id"]]) db = self.Session() run = db.query(WorkflowRun).filter_by(id=response["run_id"]).one() run_values = json.loads(run.inputs_json)["run"] self.assertEqual(run_values["web_search_query"], "Iran news today") self.assertEqual(run_values["web_search_queries"], ["Iran news today", "Iran latest news"]) db.close() async def test_interrupted_run_event_stream_emits_terminal_event(self): db = self.Session() workflow = db.query(WorkflowDefinition).filter_by(slug="research").one() run = WorkflowRun( workflow_id=workflow.id, workflow_revision_id=workflow.current_revision_id, status="interrupted", selection_mode="automatic", inputs_json=json.dumps({"input": {"prompt": "x"}, "run": {}}), error_json=json.dumps({"message": "Backend restarted before the workflow completed."}), finished_at=datetime.utcnow(), ) db.add(run) db.commit() run_id = run.id db.close() with patch.object(api, "SessionLocal", self.Session): response = api.get_workflow_events(run_id) body = b"" async for chunk in response.body_iterator: body += chunk.encode("utf-8") if isinstance(chunk, str) else chunk events = [json.loads(line) for line in body.decode("utf-8").splitlines() if line.strip()] interrupted = next(event for event in events if event["type"] == "run_interrupted") self.assertEqual(interrupted["payload"]["message"], "Backend restarted before the workflow completed.")