diff --git a/.env.example b/.env.example index f3575fe..dcb7517 100644 --- a/.env.example +++ b/.env.example @@ -21,6 +21,10 @@ OFFICEQA_DOC_MODE=parsed FRONTIER_AGENT_DATASETS_DIR= # Optional: web search and page fetching. +# Serper remains the default. Set WEB_SEARCH_PROVIDER=parallel to use the free, +# keyless Parallel Search MCP for web_search (up to 10 results per query); +# this does not change web_fetch. +WEB_SEARCH_PROVIDER=serper # Support any Serper.dev-compatible endpoint (like litescrape.com, serpbase.dev, # and others): set SERPER_BASE_URL and a provider-issued SERPER_API_KEY. SERPER_API_KEY= diff --git a/README.md b/README.md index f665ee2..f2ef38a 100644 --- a/README.md +++ b/README.md @@ -146,6 +146,7 @@ OPENAI_BASE_URL=https://your-openai-compatible-endpoint/v1 OPENAI_MODEL=your-model-name # Optional web research tools +WEB_SEARCH_PROVIDER=serper SERPER_API_KEY= SERPER_BASE_URL=https://google.serper.dev JINA_API_KEY= @@ -154,6 +155,14 @@ JINA_API_KEY= Support any Serper.dev-compatible endpoint (like litescrape.com, serpbase.dev, and others) by setting `SERPER_BASE_URL` and a provider-issued `SERPER_API_KEY`. +Set `WEB_SEARCH_PROVIDER=parallel` to use the free, keyless +[Parallel Search MCP](https://docs.parallel.ai/integrations/mcp/search-mcp) for +`web_search`. Serper remains the default, so existing keys and missing-key +errors keep their current behavior. Parallel MCP does not support the tool's +custom region, language, or time filters, or counts above 10 per query; choose +Serper when those are needed. +`web_fetch` continues to use its existing fetch provider. + Start the TUI: ```bash diff --git a/apodex/README.md b/apodex/README.md index 1297c94..66f0107 100644 --- a/apodex/README.md +++ b/apodex/README.md @@ -88,7 +88,9 @@ cp .env.example .env # OPENAI_API_KEY=... # OPENAI_BASE_URL=https://api.openai.com/v1 # OPENAI_MODEL=gpt-4o -# research mode additionally wants SERPER_API_KEY / JINA_API_KEY +# For keyless Parallel search set WEB_SEARCH_PROVIDER=parallel; it returns up +# to 10 results per query. Serper remains the default and uses SERPER_API_KEY. +# Page fetch uses JINA_API_KEY. # Interactive TUI, Stateful ReAct, against another repository frontier-agent --mode react --cwd /path/to/your/repo diff --git a/apodex/config.py b/apodex/config.py index f55e629..a77d759 100644 --- a/apodex/config.py +++ b/apodex/config.py @@ -222,7 +222,18 @@ def inspect_runtime_config( mode=active_mode, env=env, ) - if "web_search" in tool_names and not _configured(env.get("SERPER_API_KEY")): + search_provider = (env.get("WEB_SEARCH_PROVIDER") or "serper").strip().lower() + if "web_search" in tool_names and search_provider not in {"serper", "parallel"}: + issues.append(RuntimeConfigIssue( + code="invalid_web_search_provider", + message="WEB_SEARCH_PROVIDER must be set to 'serper' or 'parallel'.", + env_var="WEB_SEARCH_PROVIDER", + )) + if ( + "web_search" in tool_names + and search_provider == "serper" + and not _configured(env.get("SERPER_API_KEY")) + ): issues.append(RuntimeConfigIssue( code="missing_serper_api_key", message=( diff --git a/apodex/tests/test_config_preflight.py b/apodex/tests/test_config_preflight.py index d94c353..ba9452a 100644 --- a/apodex/tests/test_config_preflight.py +++ b/apodex/tests/test_config_preflight.py @@ -111,6 +111,22 @@ def test_search_credentials_are_checked_when_the_profile_binds_the_web_tools(): assert with_search.ok assert [issue.code for issue in with_search.warnings] == ["missing_jina_api_key"] + parallel = inspect_runtime_config( + cfg, + profile=_profile(tool_names=_WEB_TOOLS), + environ={"WEB_SEARCH_PROVIDER": "parallel"}, + ) + assert parallel.ok + assert [issue.code for issue in parallel.warnings] == ["missing_jina_api_key"] + + invalid_provider = inspect_runtime_config( + cfg, + profile=_profile(tool_names=_WEB_TOOLS), + environ={"WEB_SEARCH_PROVIDER": "unknown"}, + ) + assert not invalid_provider.ok + assert "invalid_web_search_provider" in {issue.code for issue in invalid_provider.errors} + no_web_tools = inspect_runtime_config( cfg, profile=_profile(tool_names=("bash", "read_file")), environ={}, ) diff --git a/apodex/tests/test_userenv.py b/apodex/tests/test_userenv.py index 3d04b87..fc0d4a8 100644 --- a/apodex/tests/test_userenv.py +++ b/apodex/tests/test_userenv.py @@ -322,7 +322,11 @@ def test_resolution_never_carries_a_value(launch, user_file) -> None: def test_forwarded_names_are_names_only_and_only_when_set(launch, user_file, monkeypatch) -> None: - _write(user_file, f"OPENAI_API_KEY={_SECRET}\nCUSTOM_PROVIDER_TOKEN=abc\n") + _write( + user_file, + f"OPENAI_API_KEY={_SECRET}\nCUSTOM_PROVIDER_TOKEN=abc\n" + "WEB_SEARCH_PROVIDER=parallel\n", + ) monkeypatch.setenv("SERPER_API_KEY", "serper-secret") monkeypatch.delenv("JINA_API_KEY", raising=False) @@ -331,6 +335,7 @@ def test_forwarded_names_are_names_only_and_only_when_set(launch, user_file, mon assert "OPENAI_API_KEY" in names # from the file assert "CUSTOM_PROVIDER_TOKEN" in names # file-defined, even if unlisted + assert "WEB_SEARCH_PROVIDER" in names # provider selected in the user file assert "SERPER_API_KEY" in names # well-known and exported assert "JINA_API_KEY" not in names # well-known but not set assert _SECRET not in " ".join(names) @@ -338,12 +343,14 @@ def test_forwarded_names_are_names_only_and_only_when_set(launch, user_file, mon def test_empty_resolution_forwards_only_exported_well_known_names(monkeypatch) -> None: monkeypatch.setenv("OPENAI_MODEL", "m") + monkeypatch.setenv("WEB_SEARCH_PROVIDER", "parallel") monkeypatch.delenv("OPENAI_API_KEY", raising=False) names = EnvResolution.empty().forwarded_names() assert "OPENAI_MODEL" in names assert "OPENAI_API_KEY" not in names + assert "WEB_SEARCH_PROVIDER" not in names # ── the CLI: --cwd, --model, and where .env is looked up ────────────────── diff --git a/docs/install/global-install.md b/docs/install/global-install.md index 7430466..3952aa6 100644 --- a/docs/install/global-install.md +++ b/docs/install/global-install.md @@ -70,6 +70,10 @@ chmod 600 "$config_dir/env" ``` Optional web tools take `SERPER_API_KEY` and `JINA_API_KEY` in the same file. +Add `WEB_SEARCH_PROVIDER=parallel` to this file to use the keyless Parallel +Search MCP for `web_search`; it supports up to 10 results per query. Serper +remains the default. Exporting the setting also works for native runs. +`web_fetch` keeps its existing provider. `APODEX_ENV_FILE=/path/to/file` points the CLI at a different file. Rules for the file: diff --git a/plugins/tools/_parallel_search.py b/plugins/tools/_parallel_search.py new file mode 100644 index 0000000..6cf6f8b --- /dev/null +++ b/plugins/tools/_parallel_search.py @@ -0,0 +1,348 @@ +"""Keyless Parallel Search MCP transport for the native web-search tools.""" + +from __future__ import annotations + +import asyncio +import json +import logging +import os +from contextlib import suppress +from importlib.metadata import PackageNotFoundError, version +from typing import Any +from uuid import uuid4 + +import httpx + +from frontier_agent.infra.session_context import get_task_session_id +from frontier_agent.infra.usage_meter import record_api_request + +logger = logging.getLogger(__name__) + +_MCP_URL = "https://search.parallel.ai/mcp" +_INITIAL_PROTOCOL_VERSION = "2025-03-26" +_ANONYMOUS_MAX_RESULTS = 10 +_PROVIDER_NAMES = frozenset({"serper", "parallel"}) + + +def selected_search_provider() -> str: + """Return the explicit search backend, keeping Serper as the default.""" + return (os.getenv("WEB_SEARCH_PROVIDER") or "serper").strip().lower() + + +def valid_search_provider(value: str | None = None) -> bool: + """Whether the selected backend is supported.""" + provider = selected_search_provider() if value is None else value.strip().lower() + return provider in _PROVIDER_NAMES + + +def _user_agent() -> str: + """Identify the calling project and its installed version truthfully.""" + try: + project_version = version("frontier-agent") + except PackageNotFoundError: + project_version = "dev" + return f"FrontierAgent/{project_version} (+https://github.com/ApodexAI/FrontierAgent)" + + +def _session_id() -> str: + """Use the current task id when available, or generate a request id.""" + task_id = get_task_session_id().strip() + if task_id and len(task_id) <= 100: + return task_id + return uuid4().hex + + +def _headers( + session_id: str | None = None, + protocol_version: str | None = None, +) -> dict[str, str]: + headers = { + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + "User-Agent": _user_agent(), + } + if session_id: + headers["Mcp-Session-Id"] = session_id + if protocol_version: + headers["MCP-Protocol-Version"] = protocol_version + return headers + + +def _event_response(body: str, request_id: int | None) -> dict[str, Any] | None: + """Read the JSON-RPC response from a Streamable HTTP SSE body.""" + for line in body.splitlines(): + if not line.startswith("data:"): + continue + payload = line[5:].strip() + if not payload or payload == "[DONE]": + continue + try: + value = json.loads(payload) + except json.JSONDecodeError: + continue + if not isinstance(value, dict): + continue + if request_id is None or value.get("id") == request_id: + return value + return None + + +async def _post( + client: httpx.AsyncClient, + payload: dict[str, Any], + *, + request_id: int | None, + session_id: str | None = None, + protocol_version: str | None = None, +) -> tuple[dict[str, Any] | None, httpx.Response]: + response = await client.post( + _MCP_URL, + json=payload, + headers=_headers(session_id, protocol_version), + ) + if not 200 <= response.status_code < 300: + raise httpx.HTTPStatusError( + f"Parallel Search MCP returned HTTP {response.status_code}", + request=response.request, + response=response, + ) + content_type = response.headers.get("content-type", "").lower() + if "text/event-stream" in content_type: + result = _event_response(response.text, request_id) + elif response.content: + value = response.json() + result = value if isinstance(value, dict) else None + else: + result = None + return result, response + + +def _records(value: Any, depth: int = 0) -> list[dict[str, Any]]: + """Find result rows in either MCP structured content or JSON text.""" + if depth > 16: + return [] + if isinstance(value, dict): + if any(key in value for key in ("url", "link")): + return [value] + for key in ("results", "organic", "search_results", "web_results", "data"): + rows = value.get(key) + if isinstance(rows, list): + records = [item for item in rows if isinstance(item, dict)] + if records: + return records + if isinstance(rows, dict): + found = _records(rows, depth + 1) + if found: + return found + for nested in value.values(): + found = _records(nested, depth + 1) + if found: + return found + elif isinstance(value, list): + records = [item for item in value if isinstance(item, dict)] + if records and any(any(key in item for key in ("url", "link")) for item in records): + return records + for item in value: + found = _records(item, depth + 1) + if found: + return found + return [] + + +def _normalise_result(result: dict[str, Any]) -> dict[str, Any]: + """Convert a Parallel MCP result into the native Serper-shaped result.""" + structured = result.get("structuredContent") + text_parts = [ + item.get("text", "") + for item in result.get("content", []) + if isinstance(item, dict) and item.get("type") == "text" + ] + text_value = "\n\n".join(item for item in text_parts if item) + parsed_text: Any = None + if text_value: + with suppress(json.JSONDecodeError): + parsed_text = json.loads(text_value) + source = structured if structured is not None else parsed_text + records = _records(source) + if not records and isinstance(source, dict): + nested = source.get("data") + records = _records(nested) + + organic: list[dict[str, Any]] = [] + for item in records: + link = item.get("link") or item.get("url") or "" + excerpts = item.get("excerpts") + if isinstance(excerpts, list): + snippet = "\n".join(str(part) for part in excerpts if part) + else: + snippet = item.get("snippet") or item.get("description") or item.get("content") or "" + organic.append({ + "title": str(item.get("title") or ""), + "link": str(link), + "snippet": str(snippet), + "date": str(item.get("publish_date") or item.get("published_date") or item.get("date") or ""), + }) + if organic: + return {"organic": organic} + + return {} + + +async def parallel_search_batch( + queries: list[str], + *, + num_results: int, + gl: str = "us", + hl: str = "en", + tbs: str = "", +) -> list[dict[str, Any]] | str: + """Discover and call Parallel's native MCP search tool for each query. + + Locale and time controls are rejected when requested because the anonymous + MCP tool does not expose them. Anonymous search uses the server's default + limit of ten results per query; larger counts are rejected because this + tool has no per-call result-count argument. The selected provider never + falls back to Serper on a transport, rate-limit, or tool error. + Return all candidates so callers can limit results after display filtering + and URL deduplication. + """ + if gl != "us" or hl != "en" or tbs: + return ( + "Parallel Search MCP does not support custom region, language, or " + "time filters. Select Serper to use those search options." + ) + if int(num_results) > _ANONYMOUS_MAX_RESULTS: + return ( + "Parallel Search MCP supports up to 10 results per query. " + "Select Serper for larger result counts." + ) + if not queries: + return [] + + task_id = _session_id() + protocol_version = _INITIAL_PROTOCOL_VERSION + server_session_id: str | None = None + try: + async with httpx.AsyncClient(timeout=httpx.Timeout(45.0, connect=10.0)) as client: + initialized, response = await _post( + client, + { + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": _INITIAL_PROTOCOL_VERSION, + "capabilities": {}, + "clientInfo": {"name": "FrontierAgent", "version": _user_agent().split("/")[1].split(" ")[0]}, + }, + }, + request_id=1, + ) + server_session_id = response.headers.get("Mcp-Session-Id") + if not initialized or initialized.get("error"): + raise ValueError("initialize failed") + protocol_version = str( + (initialized.get("result") or {}).get("protocolVersion") + or _INITIAL_PROTOCOL_VERSION + ) + + await _post( + client, + {"jsonrpc": "2.0", "method": "notifications/initialized"}, + request_id=None, + session_id=server_session_id, + protocol_version=protocol_version, + ) + + tools: list[dict[str, Any]] = [] + cursor: str | None = None + seen_cursors: set[str] = set() + request_id = 2 + for _ in range(20): + params = {"cursor": cursor} if cursor else {} + listed, _ = await _post( + client, + {"jsonrpc": "2.0", "id": request_id, "method": "tools/list", "params": params}, + request_id=request_id, + session_id=server_session_id, + protocol_version=protocol_version, + ) + request_id += 1 + if not listed or listed.get("error"): + raise ValueError("tool discovery failed") + result = listed.get("result") or {} + tools.extend(item for item in result.get("tools", []) if isinstance(item, dict)) + next_cursor = result.get("nextCursor") + if not next_cursor: + break + if next_cursor in seen_cursors: + raise ValueError("tool discovery did not advance") + seen_cursors.add(next_cursor) + cursor = str(next_cursor) + else: + raise ValueError("tool discovery page limit exceeded") + + if not any(item.get("name") == "web_search" for item in tools): + raise ValueError("web_search tool unavailable") + + async def call_search(query: str, call_id: int) -> dict[str, Any]: + called, _ = await _post( + client, + { + "jsonrpc": "2.0", + "id": call_id, + "method": "tools/call", + "params": { + "name": "web_search", + "arguments": { + "objective": query, + "search_queries": [query], + "session_id": task_id, + }, + }, + }, + request_id=call_id, + session_id=server_session_id, + protocol_version=protocol_version, + ) + if not called or called.get("error"): + raise ValueError("web_search call failed") + tool_result = called.get("result") or {} + if tool_result.get("isError"): + raise ValueError("web_search returned an error") + normalised = _normalise_result(tool_result) + record_api_request("parallel") + return normalised + + tasks = [ + asyncio.create_task(call_search(query, request_id + index)) + for index, query in enumerate(queries) + ] + try: + results = await asyncio.gather(*tasks) + except BaseException: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + raise + return results + except httpx.HTTPStatusError as error: + status = error.response.status_code + logger.warning("Parallel Search MCP request failed with HTTP %d", status) + return f"Parallel Search MCP returned HTTP {status}." + except (httpx.HTTPError, ValueError, TypeError, KeyError, json.JSONDecodeError): + logger.warning("Parallel Search MCP request failed during search") + return "Parallel Search MCP search failed. Check network access and try again." + finally: + if server_session_id: + try: + async with httpx.AsyncClient(timeout=httpx.Timeout(5.0, connect=2.0)) as client: + await client.delete( + _MCP_URL, + headers=_headers(server_session_id, protocol_version), + ) + except httpx.HTTPError: + logger.debug("Parallel Search MCP session cleanup failed") + + +__all__ = ["parallel_search_batch", "selected_search_provider", "valid_search_provider"] diff --git a/plugins/tools/web_search.py b/plugins/tools/web_search.py index 9ec1161..5698f43 100644 --- a/plugins/tools/web_search.py +++ b/plugins/tools/web_search.py @@ -14,6 +14,11 @@ from frontier_agent.infra.config import get_config from frontier_agent.infra.usage_meter import record_api_request from plugins.tools._coerce import coerce_json_list +from plugins.tools._parallel_search import ( + parallel_search_batch, + selected_search_provider, + valid_search_provider, +) from plugins.tools._single_flight import SingleFlightCoalescer logger = logging.getLogger(__name__) @@ -457,6 +462,36 @@ async def web_search( from plugins.tools._overflow import maybe_overflow + provider = selected_search_provider() + if not valid_search_provider(provider): + return "Error: WEB_SEARCH_PROVIDER must be 'serper' or 'parallel'." + if provider == "parallel": + datas = await parallel_search_batch( + queries, + num_results=num_results, + gl=gl, + hl=hl, + tbs=tbs, + ) + if isinstance(datas, str): + return f"Error: {datas}" + if not datas: + return f"No results found for: {queries[0]}" + if len(datas) == 1: + return maybe_overflow( + "web_search", + _format_results( + {**datas[0], "organic": _dedupe_display_organic( + datas[0].get("organic") or [], set(), num_results, + )}, + max_organic=num_results, + ), + ) + return maybe_overflow( + "web_search", + _format_parallel_query_results(queries, datas, num_results), + ) + if len(queries) == 1: data = await raw_web_search(queries[0], num_results, gl, hl, tbs) if not data: @@ -469,6 +504,26 @@ async def web_search( ) +def _format_parallel_query_results( + queries: list[str], datas: list[dict], max_results: int, +) -> str: + """Keep the canonical query-labelled output for Parallel MCP batches.""" + seen_urls: set[str] = set() + blocks: list[str] = [] + for query, data in zip(queries, datas, strict=False): + if not data: + blocks.append(f"## Query: {query}\n\nNo results (error or empty).") + continue + organic = _dedupe_display_organic( + data.get("organic") or [], seen_urls, max_results, + ) + blocks.append( + f"## Query: {query}\n\n" + f"{_format_results({**data, 'organic': organic}, max_organic=max_results)}" + ) + return "\n\n---\n\n".join(blocks) + + def _normalise_queries(query: str | list[str]) -> list[str]: """Collapse the LangChain payload into a clean list of non-empty strings. diff --git a/plugins/tools/web_search_aligned.py b/plugins/tools/web_search_aligned.py index dfbfc22..fbf1ea0 100644 --- a/plugins/tools/web_search_aligned.py +++ b/plugins/tools/web_search_aligned.py @@ -24,6 +24,11 @@ from frontier_agent.core.tool import tool from frontier_agent.infra.usage_meter import record_api_request +from plugins.tools._parallel_search import ( + parallel_search_batch, + selected_search_provider, + valid_search_provider, +) from plugins.tools.web_search import is_snippet_blocked_result logger = logging.getLogger(__name__) @@ -269,9 +274,6 @@ async def web_search_aligned( Returns: Numbered plain-text list of search results, each with Title, Snippet, and URL """ - if not _serper_api_key(): - return "[ERROR]: SERPER_API_KEY environment variable not set." - # The reference tool accepts ``Union[str, List[str]]`` and tolerates JSON-encoded # lists. We mirror both shapes so any prompt that worked there works here. q = _ensure_list(q) @@ -280,6 +282,48 @@ async def web_search_aligned( if not queries: return "[ERROR]: Search query 'q' is required and cannot be empty." + provider = selected_search_provider() + if not valid_search_provider(provider): + return "[ERROR]: WEB_SEARCH_PROVIDER must be 'serper' or 'parallel'." + if provider == "parallel": + if location is not None or page not in (None, 1) or autocorrect is not None: + return ( + "[ERROR]: Parallel Search MCP does not support location, page, or " + "autocorrect options. Select Serper to use those search options." + ) + result_limit = 10 if num is None else max(1, min(num, 100)) + datas = await parallel_search_batch( + queries, + num_results=result_limit, + gl=gl, + hl=hl, + tbs=tbs or "", + ) + if isinstance(datas, str): + return f"[ERROR]: {datas}" + merged: list[dict] = [] + seen_urls: set[str] = set() + for data in datas: + displayed = 0 + for item in data.get("organic", []): + link = item.get("link", "") + if _is_banned_url(link) or is_snippet_blocked_result(item): + continue + if link and link in seen_urls: + continue + if link: + seen_urls.add(link) + merged.append(item) + displayed += 1 + if displayed >= result_limit: + break + if not merged: + return "No search results found." + return _format_results_plaintext(merged) + + if not _serper_api_key(): + return "[ERROR]: SERPER_API_KEY environment variable not set." + try: all_organic: list[list[dict]] = await asyncio.gather( *[ diff --git a/tests/test_parallel_search.py b/tests/test_parallel_search.py new file mode 100644 index 0000000..9ee17cf --- /dev/null +++ b/tests/test_parallel_search.py @@ -0,0 +1,375 @@ +"""Native provider selection and anonymous Parallel Search MCP transport.""" + +from __future__ import annotations + +import asyncio +import importlib +import json +from types import SimpleNamespace +from typing import Any + +import pytest + +from plugins.tools import _parallel_search as parallel + + +class _Response: + def __init__( + self, + value: dict[str, Any] | None = None, + *, + session_id: str | None = None, + ) -> None: + self.status_code = 200 + self.request = object() + self.content = json.dumps(value).encode() if value is not None else b"" + self.text = self.content.decode() + self.headers = {"content-type": "application/json"} + if session_id: + self.headers["Mcp-Session-Id"] = session_id + + def json(self) -> dict[str, Any]: + return json.loads(self.text) + + +class _FakeClient: + calls: list[tuple[str, dict[str, Any] | None, dict[str, str]]] + + def __init__(self, **_: Any) -> None: + self.calls = _FakeClient.calls + + async def __aenter__(self) -> _FakeClient: + return self + + async def __aexit__(self, *_: Any) -> None: + return None + + async def post( + self, + url: str, + *, + json: dict[str, Any], + headers: dict[str, str], + ) -> _Response: + self.calls.append((url, json, headers)) + method = json.get("method") + if method == "initialize": + return _Response( + {"jsonrpc": "2.0", "id": 1, "result": {"protocolVersion": "2025-03-26"}}, + session_id="test-session", + ) + if method == "notifications/initialized": + return _Response() + if method == "tools/list": + return _Response({ + "jsonrpc": "2.0", + "id": json["id"], + "result": {"tools": [{"name": "web_search"}]}, + }) + if method == "tools/call": + return _Response({ + "jsonrpc": "2.0", + "id": json["id"], + "result": { + "content": [{"type": "text", "text": "Found a result."}], + "structuredContent": { + "results": [{ + "title": "Example", + "url": "https://example.com/page", + "excerpts": ["A useful excerpt."], + }], + }, + }, + }) + raise AssertionError(f"unexpected MCP method: {method}") + + async def delete(self, url: str, *, headers: dict[str, str]) -> _Response: + self.calls.append((url, None, headers)) + return _Response() + + +def test_serper_remains_the_implicit_provider(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("WEB_SEARCH_PROVIDER", raising=False) + assert parallel.selected_search_provider() == "serper" + assert parallel.valid_search_provider() + assert parallel.valid_search_provider("parallel") + assert not parallel.valid_search_provider("other") + + +@pytest.mark.asyncio +async def test_parallel_mcp_discovers_and_calls_search_with_project_user_agent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _FakeClient.calls = [] + monkeypatch.setattr(parallel.httpx, "AsyncClient", _FakeClient) + + result = await parallel.parallel_search_batch( + ["FrontierAgent MCP search"], num_results=5, + ) + + assert result == [{ + "organic": [{ + "title": "Example", + "link": "https://example.com/page", + "snippet": "A useful excerpt.", + "date": "", + }], + }] + assert [call[1]["method"] if call[1] else "delete" for call in _FakeClient.calls] == [ + "initialize", "notifications/initialized", "tools/list", "tools/call", "delete", + ] + for _, payload, headers in _FakeClient.calls: + assert headers["User-Agent"].startswith("FrontierAgent/") + assert "github.com/ApodexAI/FrontierAgent" in headers["User-Agent"] + assert "Authorization" not in headers + assert "X-API-KEY" not in headers + if payload and payload.get("method") in {"tools/list", "tools/call"}: + assert headers["MCP-Protocol-Version"] == "2025-03-26" + assert headers["Mcp-Session-Id"] == "test-session" + + +@pytest.mark.asyncio +async def test_parallel_mcp_rejects_filters_it_cannot_honor() -> None: + result = await parallel.parallel_search_batch( + ["current information"], num_results=5, gl="gb", + ) + assert isinstance(result, str) + assert "custom region" in result + + +@pytest.mark.asyncio +async def test_parallel_mcp_rejects_counts_above_anonymous_default( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _FakeClient.calls = [] + monkeypatch.setattr(parallel.httpx, "AsyncClient", _FakeClient) + + result = await parallel.parallel_search_batch( + ["current information"], num_results=11, + ) + + assert isinstance(result, str) + assert "up to 10 results per query" in result + assert "Select Serper" in result + assert _FakeClient.calls == [] + + +@pytest.mark.asyncio +async def test_parallel_mcp_dispatches_list_queries_concurrently( + monkeypatch: pytest.MonkeyPatch, +) -> None: + active = 0 + peak_active = 0 + + class ConcurrentClient(_FakeClient): + async def post( + self, + url: str, + *, + json: dict[str, Any], + headers: dict[str, str], + ) -> _Response: + nonlocal active, peak_active + if json.get("method") == "tools/call": + active += 1 + peak_active = max(peak_active, active) + await asyncio.sleep(0.02) + active -= 1 + return await super().post(url, json=json, headers=headers) + + _FakeClient.calls = [] + monkeypatch.setattr(parallel.httpx, "AsyncClient", ConcurrentClient) + + result = await parallel.parallel_search_batch( + ["first query", "second query"], num_results=3, + ) + + assert isinstance(result, list) + assert len(result) == 2 + assert peak_active == 2 + calls = [ + payload for _, payload, _ in _FakeClient.calls + if payload and payload.get("method") == "tools/call" + ] + assert len({payload["id"] for payload in calls}) == 2 + assert { + payload["params"]["arguments"]["objective"] for payload in calls + } == {"first query", "second query"} + + +@pytest.mark.asyncio +async def test_original_parallel_route_keeps_its_domain_exclusions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = importlib.import_module("plugins.tools.web_search") + monkeypatch.setenv("WEB_SEARCH_PROVIDER", "parallel") + async def search(*_args: Any, **_kwargs: Any) -> list[dict[str, Any]]: + return [{ + "organic": [ + {"title": "Video", "link": "https://youtube.com/watch/1", "snippet": ""}, + {"title": "Useful page", "link": "https://example.com/page", "snippet": "Text."}, + ], + }] + + monkeypatch.setattr(module, "parallel_search_batch", search) + output = await module.web_search.ainvoke({"q": "sample topic"}) + + assert "Useful page" in output + assert "youtube.com" not in output + + +@pytest.mark.asyncio +async def test_aligned_parallel_route_keeps_results_from_each_query( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = importlib.import_module("plugins.tools.web_search_aligned") + monkeypatch.setenv("WEB_SEARCH_PROVIDER", "parallel") + requested_limits: list[int] = [] + + async def search( + queries: list[str], *, num_results: int, **_: Any, + ) -> list[dict[str, Any]]: + requested_limits.append(num_results) + return [ + {"organic": [ + {"title": f"{query} result {index}", + "link": f"https://{query}{index}.example.com/page", + "snippet": "Useful result."} + for index in range(1, 3) + ]} + for query in queries + ] + + monkeypatch.setattr(module, "parallel_search_batch", search) + output = await module.web_search_aligned.ainvoke({ + "q": ["first", "second"], + "num": 2, + }) + + assert requested_limits == [2] + assert output.count("URL: https://") == 4 + assert "first result 1" in output + assert "second result 1" in output + + +@pytest.mark.asyncio +async def test_parallel_errors_do_not_fall_back_to_serper( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = importlib.import_module("plugins.tools.web_search") + monkeypatch.setenv("WEB_SEARCH_PROVIDER", "parallel") + monkeypatch.delenv("SERPER_API_KEY", raising=False) + + async def failed(*_args: Any, **_kwargs: Any) -> str: + return "Parallel Search MCP returned HTTP 429." + + async def no_serper(*_args: Any, **_kwargs: Any) -> dict[str, Any]: + raise AssertionError("explicit Parallel selection must not call Serper") + + monkeypatch.setattr(module, "parallel_search_batch", failed) + monkeypatch.setattr(module, "raw_web_search", no_serper) + output = await module.web_search.ainvoke({"q": "sample topic"}) + + assert "Parallel Search MCP returned HTTP 429" in output + + +def test_react_profile_loader_keeps_both_search_implementations_on_parallel( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from apodex.agent_tools import terminal_tool_registry + from plugins.tools.web_search import web_search + from plugins.tools.web_search_aligned import web_search_aligned + from workflows.stateful_react_agent.nodes.main_agent import ( + _replace_tool_impls, + _tools_for_stateful_react, + ) + from workflows.stateful_react_agent.profile import load_react_profile + + monkeypatch.setenv("WEB_SEARCH_PROVIDER", "parallel") + monkeypatch.delenv("REACT_NO_WEB", raising=False) + resource_mgr = SimpleNamespace( + all_tools=terminal_tool_registry(), + global_tool_policy=None, + ) + + for profile_name, expected in ( + ("simple", web_search), + ("tui", web_search_aligned), + ): + agent = load_react_profile(profile_name)["agent"] + tools = _replace_tool_impls( + _tools_for_stateful_react(resource_mgr, agent), agent, + ) + selected = next(tool for tool in tools if tool.name == "web_search") + assert selected is expected + assert parallel.selected_search_provider() == "parallel" + + +@pytest.mark.parametrize("implementation", ["original", "aligned"]) +@pytest.mark.parametrize("scenario", ["filtered", "duplicates", "multiple_queries"]) +@pytest.mark.asyncio +async def test_parallel_limits_displayed_results_after_filtering_and_deduplication( + monkeypatch: pytest.MonkeyPatch, implementation: str, scenario: str, +) -> None: + def row(title: str, url: str) -> dict[str, Any]: + return {"title": title, "url": url, "excerpts": ["Useful text."], + "publish_date": "2026-09-29"} + + first = row("First useful", "https://example.com/first") + second = row("Second useful", "https://example.com/second") + extra = row("Over limit", "https://example.com/extra") + blocked_url = ( + "https://youtube.com/watch/1" if implementation == "original" + else "https://huggingface.co/datasets/example" + ) + blocked = row("Blocked result", blocked_url) + if scenario == "filtered": + queries, limit = ["first"], 1 + responses = {"first": [blocked, first, extra]} + expected = ["First useful"] + elif scenario == "duplicates": + queries, limit = ["first"], 2 + responses = {"first": [first, first, second, extra]} + expected = ["First useful", "Second useful"] + else: + queries, limit = ["first", "second"], 1 + responses = {"first": [first, second], "second": [first, blocked, second, extra]} + expected = ["First useful", "Second useful"] + + class ResultsClient(_FakeClient): + async def post(self, url, *, json, headers): + if json.get("method") == "tools/call": + query = json["params"]["arguments"]["objective"] + return _Response({ + "jsonrpc": "2.0", "id": json["id"], + "result": {"structuredContent": {"results": responses[query]}}, + }) + return await super().post(url, json=json, headers=headers) + + _FakeClient.calls = [] + monkeypatch.setattr(parallel.httpx, "AsyncClient", ResultsClient) + monkeypatch.setenv("WEB_SEARCH_PROVIDER", "parallel") + if implementation == "original": + module = importlib.import_module("plugins.tools.web_search") + output = await module.web_search.ainvoke({"q": queries, "num_results": limit}) + else: + module = importlib.import_module("plugins.tools.web_search_aligned") + output = await module.web_search_aligned.ainvoke({"q": queries, "num": limit}) + assert output.count("Date: 2026-09-29") == len(expected) + + assert output.count("URL: https://") == len(expected) + for title in expected: + assert output.count(title) == 1 + assert "Blocked result" not in output + assert "Over limit" not in output + + +@pytest.mark.parametrize("field", ["publish_date", "published_date", "date"]) +@pytest.mark.parametrize("structured", [True, False]) +def test_parallel_normalises_publication_dates(field: str, structured: bool) -> None: + payload = {"results": [{"url": "https://example.com", field: "2026-09-29"}]} + result = ( + {"structuredContent": payload} if structured + else {"content": [{"type": "text", "text": json.dumps(payload)}]} + ) + assert parallel._normalise_result(result)["organic"][0]["date"] == "2026-09-29"