diff --git a/backend/agent/router.py b/backend/agent/router.py index 7afcdf1..0e66380 100644 --- a/backend/agent/router.py +++ b/backend/agent/router.py @@ -184,6 +184,14 @@ def _normalize_router_inputs(selected: Dict[str, Any], inputs: Any, message: str return normalized +def _forced_web_manifest(manifests: List[Dict[str, Any]], *, library_slug: Optional[str]) -> Optional[Dict[str, Any]]: + if library_slug: + combined = _manifest_by_slug(manifests, "knowledge-web-answer") + if combined: + return combined + return _manifest_by_slug(manifests, "web-answer") + + async def select_workflow( db: Session, *, @@ -262,7 +270,17 @@ async def select_workflow( workflow_id = str(parsed.get("workflow_id") or "") selected = next((item for item in manifests if item["workflow_id"] == workflow_id or item["slug"] == workflow_id), None) confidence = float(parsed.get("confidence") or 0) - if selected is None or confidence < confidence_threshold: + forced_web = _forced_web_manifest(manifests, library_slug=library_slug) if web_search_enabled else None + forced = False + if selected is None and forced_web: + selected = forced_web + confidence = max(confidence, 0.95) + forced = True + elif selected is not None and forced_web and "web" not in (selected.get("required_capabilities") or []): + selected = forced_web + confidence = max(confidence, 0.95) + forced = True + if selected is None or (confidence < confidence_threshold and not forced): return fallback_result if "rag" in selected["required_capabilities"] and not library_slug: return fallback_result @@ -271,7 +289,7 @@ async def select_workflow( "workflow_id": selected["workflow_id"], "workflow_slug": selected["slug"], "confidence": max(0.0, min(confidence, 1.0)), - "reason": str(parsed.get("reason") or ""), + "reason": (str(parsed.get("reason") or "") + (" Forced web search is enabled." if forced else "")).strip(), "inputs": inputs, "fallback": False, "model": model, diff --git a/backend/tests/test_workflow_router_builtins.py b/backend/tests/test_workflow_router_builtins.py index dc32916..73bb2ce 100644 --- a/backend/tests/test_workflow_router_builtins.py +++ b/backend/tests/test_workflow_router_builtins.py @@ -66,6 +66,23 @@ class RouterAndBuiltinTests(unittest.IsolatedAsyncioTestCase): 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({