Files
Heimgeist/backend/tests/test_workflow_api.py

163 lines
6.7 KiB
Python

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, after_sequence=0, follow=True)
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.")