163 lines
6.7 KiB
Python
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.")
|