Files
Heimgeist/backend/tests/test_workflow_router_builtins.py

187 lines
10 KiB
Python

import json
from pathlib import Path
import tempfile
import unittest
from unittest.mock import AsyncMock, patch
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from backend.agent.models import WorkflowDefinition
from backend.agent.registry import NativeToolProvider
from backend.agent.router import select_workflow
from backend.agent.tools import register_native_tools
from backend.agent.validation import validate_workflow_graph
from backend.agent.workflows.builtins import BUILTIN_WORKFLOWS, seed_builtin_workflows
from backend.database import Base
from backend.ollama_client import OllamaChatResult
class RouterAndBuiltinTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.temp = tempfile.TemporaryDirectory()
engine = create_engine(f"sqlite:///{Path(self.temp.name) / 'router.db'}")
Base.metadata.create_all(engine)
self.Session = sessionmaker(bind=engine)
self.engine = engine
db = self.Session(); seed_builtin_workflows(db); db.close()
async def asyncTearDown(self):
self.engine.dispose(); self.temp.cleanup()
async def test_valid_selection_malformed_disabled_and_low_confidence_fallback(self):
db = self.Session()
web = db.query(WorkflowDefinition).filter_by(slug="web-answer").one()
with patch("backend.agent.router.chat_typed", new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({"workflow_id": web.id, "confidence": 0.9, "reason": "current", "inputs": {}})))):
result = await select_workflow(db, message="tell me something useful", recent_messages=[], attachments=[], library_slug=None, router_model=None, chat_model="model")
self.assertEqual(result["workflow_slug"], "web-answer")
with patch("backend.agent.router.chat_typed", new=AsyncMock(return_value=OllamaChatResult(content="not json"))):
self.assertEqual((await select_workflow(db, message="x", recent_messages=[], attachments=[], library_slug=None, router_model=None, chat_model="model"))["workflow_slug"], "input-output")
web.enabled = False; db.commit()
with patch("backend.agent.router.chat_typed", new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({"workflow_id": web.id, "confidence": 0.9, "reason": "x", "inputs": {}})))):
self.assertEqual((await select_workflow(db, message="x", recent_messages=[], attachments=[], library_slug=None, router_model=None, chat_model="model"))["workflow_slug"], "input-output")
web.enabled = True; db.commit()
with patch("backend.agent.router.chat_typed", new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({"workflow_id": web.id, "confidence": 0.2, "reason": "x", "inputs": {}})))):
self.assertEqual((await select_workflow(db, message="x", recent_messages=[], attachments=[], library_slug=None, router_model=None, chat_model="model"))["workflow_slug"], "input-output")
db.close()
async def test_fast_paths_auto_select_or_force_web(self):
db = self.Session()
web = db.query(WorkflowDefinition).filter_by(slug="web-answer").one()
combined = db.query(WorkflowDefinition).filter_by(slug="knowledge-web-answer").one()
with patch(
"backend.agent.router.chat_typed",
new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({
"workflow_id": web.id,
"confidence": 0.95,
"reason": "forced web",
"inputs": {"web_search_query": "Iran news today", "web_search_queries": ["Iran news today", "Iran latest news"]},
}))),
) as mocked:
result = await select_workflow(
db, message="hello", recent_messages=[], attachments=[],
library_slug=None, router_model=None, chat_model="model", web_search_enabled=True,
)
self.assertEqual(result["workflow_slug"], "web-answer")
self.assertEqual(result["inputs"]["web_search_query"], "Iran news today")
mocked.assert_awaited()
direct = db.query(WorkflowDefinition).filter_by(slug="input-output").one()
with patch(
"backend.agent.router.chat_typed",
new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({
"workflow_id": direct.id,
"confidence": 0.95,
"reason": "bad router choice",
"inputs": {},
}))),
):
result = await select_workflow(
db, message="what happened in Iran today", recent_messages=[], attachments=[],
library_slug=None, router_model=None, chat_model="model", web_search_enabled=True,
)
self.assertEqual(result["workflow_slug"], "web-answer")
self.assertEqual(result["inputs"]["web_search_query"], "what happened in Iran today")
with patch(
"backend.agent.router.chat_typed",
new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({
"workflow_id": web.id,
"confidence": 0.95,
"reason": "current",
"inputs": {"web_search_queries": ["latest Iran news"]},
}))),
):
result = await select_workflow(
db, message="latest news today", recent_messages=[], attachments=[],
library_slug=None, router_model=None, chat_model="model", web_search_enabled=False,
)
self.assertEqual(result["workflow_slug"], "web-answer")
self.assertEqual(result["inputs"]["web_search_query"], "latest Iran news")
with patch(
"backend.agent.router.chat_typed",
new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({
"workflow_id": combined.id,
"confidence": 0.95,
"reason": "forced web with knowledge",
"inputs": {"web_search_query": "compare notes with latest docs", "web_search_queries": ["compare notes with latest docs"]},
}))),
):
result = await select_workflow(
db, message="hello", recent_messages=[], attachments=[],
library_slug="notes", router_model=None, chat_model="model", web_search_enabled=True,
)
self.assertEqual(result["workflow_slug"], "knowledge-web-answer")
self.assertEqual(result["inputs"]["web_search_query"], "compare notes with latest docs")
with patch("backend.agent.router.chat_typed", new=AsyncMock()) as mocked:
result = await select_workflow(
db, message="hello", recent_messages=[], attachments=[],
library_slug=None, router_model=None, chat_model="model", web_search_enabled=False,
)
self.assertEqual(result["workflow_slug"], "input-output")
mocked.assert_not_awaited()
result = await select_workflow(
db, message="remember this: I prefer short answers", recent_messages=[], attachments=[],
library_slug="memory", router_model=None, chat_model="model", web_search_enabled=False,
)
self.assertEqual(result["workflow_slug"], "remember-this")
db.close()
async def test_missing_router_model_falls_back_to_chat_model(self):
db = self.Session()
with patch("backend.agent.router.list_model_catalog", new=AsyncMock(return_value={"chat_models": ["chat"]})), patch("backend.agent.router.chat_typed", new=AsyncMock(side_effect=RuntimeError("bad response"))) as mocked:
result = await select_workflow(db, message="x", recent_messages=[], attachments=[], library_slug=None, router_model="missing", chat_model="chat")
self.assertEqual(result["workflow_slug"], "input-output")
self.assertEqual(mocked.await_args.args[0], "chat")
db.close()
async def test_router_prompt_formats_followup_context_and_filters_for_forced_web(self):
db = self.Session()
web = db.query(WorkflowDefinition).filter_by(slug="web-answer").one()
with patch(
"backend.agent.router.chat_typed",
new=AsyncMock(return_value=OllamaChatResult(content=json.dumps({
"workflow_id": web.id,
"confidence": 0.95,
"reason": "current follow-up",
"inputs": {"web_search_query": "Iran political news today", "web_search_queries": ["Iran political news today"]},
}))),
) as mocked:
result = await select_workflow(
db,
message="and what's the news politically?",
recent_messages=[
{"role": "user", "content": "what's the weather in Iran today"},
{"role": "assistant", "content": "The weather in Tehran, Iran today is clear."},
],
attachments=[],
library_slug=None,
router_model=None,
chat_model="model",
web_search_enabled=True,
)
prompt = mocked.await_args.args[1][0]["content"]
self.assertEqual(result["inputs"]["web_search_query"], "Iran political news today")
self.assertIn("Recent conversation:\n1. User: what's the weather in Iran today", prompt)
self.assertIn("search for 'Iran political news today'", prompt)
self.assertIn("Never choose Input -> Output for weather, news, politics today", prompt)
self.assertIn("Web Answer is for straightforward current lookups", prompt)
self.assertIn("Research is for explicit deeper investigation", prompt)
self.assertIn("Allowed workflows. This list is already filtered", prompt)
self.assertIn("web-answer", prompt)
self.assertNotIn("input-output", prompt)
self.assertNotIn("Cost:", prompt)
self.assertNotIn("Web search mode:", prompt)
self.assertNotIn("Selected knowledge database:", prompt)
db.close()
async def test_all_builtins_validate(self):
registry = NativeToolProvider(); register_native_tools(registry)
for item in BUILTIN_WORKFLOWS:
with self.subTest(item=item["slug"]):
self.assertTrue(validate_workflow_graph(item["graph"], registry).valid)