diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py index 12d570f1733..85602bf1d86 100644 --- a/litellm/llms/parallel_ai/search/transformation.py +++ b/litellm/llms/parallel_ai/search/transformation.py @@ -1,7 +1,7 @@ """ -Calls Parallel AI's /search endpoint to search the web. +Calls Parallel AI's /v1/search endpoint to search the web. -Parallel AI API Reference: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search +Parallel AI API Reference: https://docs.parallel.ai/api-reference/search/search """ from typing import Dict, List, Optional, TypedDict, Union @@ -18,36 +18,43 @@ from litellm.secret_managers.main import get_secret_str class _ParallelAISourcePolicy(TypedDict, total=False): - """Source policy for Parallel AI search results.""" - - allowed_domains: List[str] # Optional - list of allowed domains - disallowed_domains: List[str] # Optional - list of disallowed domains + include_domains: List[str] + exclude_domains: List[str] + after_date: str -class _ParallelAISearchRequestRequired(TypedDict): - """Required fields for Parallel AI Search API request.""" - - # Note: At least one of objective or search_queries must be provided - pass +class _ParallelAIExcerptSettings(TypedDict, total=False): + max_chars_per_result: int -class ParallelAISearchRequest(_ParallelAISearchRequestRequired, total=False): +class _ParallelAIAdvancedSettings(TypedDict, total=False): + source_policy: _ParallelAISourcePolicy + excerpt_settings: _ParallelAIExcerptSettings + fetch_policy: Dict + location: str + max_results: int + + +class ParallelAISearchRequest(TypedDict, total=False): """ - Parallel AI Search API request format. - Based on: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search + Parallel AI v1 Search API request format. + Based on: https://docs.parallel.ai/api-reference/search/search """ + search_queries: List[str] # Required - at least one keyword search query objective: str # Optional - natural-language description of search goal - search_queries: List[str] # Optional - list of keyword search queries - processor: str # Optional - search processor ('base', 'pro'), default 'base' - max_results: int # Optional - maximum number of results, default 10 - max_chars_per_result: int # Optional - max characters per result excerpt - source_policy: _ParallelAISourcePolicy # Optional - source policy for allowed/disallowed domains + mode: str # Optional - 'turbo', 'basic', or 'advanced' (default 'advanced') + max_chars_total: int # Optional - upper bound on total excerpt characters + session_id: str # Optional - tracks calls across search/extract requests + client_model: str # Optional - model consuming the results + advanced_settings: _ParallelAIAdvancedSettings + + +LEGACY_PROCESSOR_TO_MODE = {"base": "basic", "pro": "advanced"} class ParallelAISearchConfig(BaseSearchConfig): PARALLEL_AI_API_BASE = "https://api.parallel.ai" - PARALLEL_HEADER_SEARCH_EXTRACT_VALUE = "search-extract-2025-10-10" @staticmethod def ui_friendly_name() -> str: @@ -60,9 +67,6 @@ class ParallelAISearchConfig(BaseSearchConfig): api_base: Optional[str] = None, **kwargs, ) -> Dict: - """ - Validate environment and return headers. - """ api_key = ( api_key or get_secret_str("PARALLEL_AI_API_KEY") @@ -74,7 +78,6 @@ class ParallelAISearchConfig(BaseSearchConfig): ) headers["x-api-key"] = api_key headers["Content-Type"] = "application/json" - headers["parallel-beta"] = self.PARALLEL_HEADER_SEARCH_EXTRACT_VALUE return headers def get_complete_url( @@ -84,32 +87,18 @@ class ParallelAISearchConfig(BaseSearchConfig): data: Optional[Union[Dict, List[Dict]]] = None, **kwargs, ) -> str: - """ - Get complete URL for Search endpoint. - """ api_base = ( api_base or get_secret_str("PARALLEL_AI_API_BASE") or self.PARALLEL_AI_API_BASE ) - # Parallel AI search endpoint is at /v1beta/search - if not api_base.endswith("/v1beta/search"): - if api_base.endswith("/"): - api_base = f"{api_base}v1beta/search" - else: - api_base = f"{api_base}/v1beta/search" + api_base = api_base.rstrip("/") + if not api_base.endswith("/v1/search"): + api_base = f"{api_base.removesuffix('/v1')}/v1/search" return api_base - def _transform_query_to_objective(self, query: Union[str, List[str]]) -> str: - """ - Transform query to objective. - """ - if isinstance(query, list): - return " ".join(query) - return query - def transform_search_request( self, query: Union[str, List[str]], @@ -117,57 +106,78 @@ class ParallelAISearchConfig(BaseSearchConfig): **kwargs, ) -> Dict: """ - Transform Search request to Parallel AI API format. + Transform Search request to Parallel AI v1 API format. Args: query: Search query (string or list of strings) - - If string: maps to `objective` (natural language) + - If string: maps to `search_queries` (single item) and `objective` - If list: maps to `search_queries` (keyword queries) optional_params: Optional parameters for the request - - max_results: Maximum number of search results (default 10) - - search_domain_filter: List of domains to include -> maps to `source_policy.allowed_domains` - - exclude_domains: List of domains to exclude -> maps to `source_policy.disallowed_domains` - - processor: Search processor ('base', 'pro') - - max_chars_per_result: Max characters per result excerpt + - mode: Search mode ('turbo', 'basic', 'advanced'); defaults to 'basic' + - processor: Legacy v1beta param; 'base' maps to mode 'basic', 'pro' to 'advanced' + - max_results: Maximum number of search results -> `advanced_settings.max_results` + - search_domain_filter: Domains to include -> `advanced_settings.source_policy.include_domains` + - exclude_domains: Domains to exclude -> `advanced_settings.source_policy.exclude_domains` + - country: ISO 3166-1 alpha-2 code -> `advanced_settings.location` + - max_chars_per_result: -> `advanced_settings.excerpt_settings.max_chars_per_result` + - Any other params are passed through to the request body as-is Returns: - Dict with typed request data following ParallelAISearchRequest spec + Dict with request data following the v1 search request spec """ + params = dict(optional_params) + request_data: ParallelAISearchRequest = {} - # Map query to objective (string or list both become objective) if isinstance(query, list): - request_data["objective"] = self._transform_query_to_objective(query) + request_data["search_queries"] = query else: + request_data["search_queries"] = [query] request_data["objective"] = query - # Transform Perplexity unified spec parameters to Parallel AI format - if "max_results" in optional_params: - request_data["max_results"] = optional_params["max_results"] + mode = params.pop("mode", None) + processor = params.pop("processor", None) + if mode is None and processor is not None: + mode = LEGACY_PROCESSOR_TO_MODE.get(processor, processor) + # the v1 API defaults to 'advanced' when mode is omitted; default to 'basic' + # instead to keep v1beta's default tier (processor 'base') and litellm's + # $0.004/query cost map entry for `parallel_ai/search` accurate + request_data["mode"] = mode or "basic" + + advanced_settings: _ParallelAIAdvancedSettings = {} + + if "max_results" in params: + advanced_settings["max_results"] = params.pop("max_results") + + if "country" in params: + advanced_settings["location"] = params.pop("country") + + if "max_chars_per_result" in params: + advanced_settings["excerpt_settings"] = { + "max_chars_per_result": params.pop("max_chars_per_result") + } - # Map domain filters to source_policy source_policy: _ParallelAISourcePolicy = {} - if "search_domain_filter" in optional_params: - source_policy["allowed_domains"] = optional_params["search_domain_filter"] + if "search_domain_filter" in params: + source_policy["include_domains"] = params.pop("search_domain_filter") - if "exclude_domains" in optional_params: - source_policy["disallowed_domains"] = optional_params["exclude_domains"] + if "exclude_domains" in params: + source_policy["exclude_domains"] = params.pop("exclude_domains") if source_policy: - request_data["source_policy"] = source_policy + advanced_settings["source_policy"] = source_policy - # Convert to dict before dynamic key assignments - result_data = dict(request_data) + advanced_settings.update(params.pop("advanced_settings", {})) - # pass through all other parameters as-is - for param, value in optional_params.items(): - if ( - param not in self.get_supported_perplexity_optional_params() - and param not in result_data - ): - result_data[param] = value + if advanced_settings: + request_data["advanced_settings"] = advanced_settings + # unified-spec param with no v1 equivalent + params.pop("max_tokens_per_page", None) + + result_data: Dict = dict(request_data) + result_data.update(params) return result_data def transform_search_response( @@ -177,36 +187,27 @@ class ParallelAISearchConfig(BaseSearchConfig): **kwargs, ) -> SearchResponse: """ - Transform Parallel AI API response to LiteLLM unified SearchResponse format. + Transform Parallel AI v1 API response to LiteLLM unified SearchResponse format. - Parallel AI → LiteLLM mappings: - - results[].title → SearchResult.title - - results[].url → SearchResult.url - - results[].excerpts (array) → SearchResult.snippet (joined string) - - No date/last_updated fields in Parallel AI response (set to None) - - Args: - raw_response: Raw httpx response from Parallel AI API - logging_obj: Logging object for tracking - - Returns: - SearchResponse with standardized format + Parallel AI -> LiteLLM mappings: + - results[].title -> SearchResult.title + - results[].url -> SearchResult.url + - results[].excerpts (array) -> SearchResult.snippet (joined string) + - results[].publish_date -> SearchResult.date """ response_json = raw_response.json() - # Transform results to SearchResult objects results = [] for result in response_json.get("results", []): - # Join excerpts array into a single snippet string - excerpts = result.get("excerpts", []) + excerpts = result.get("excerpts") or [] snippet = " ... ".join(excerpts) if excerpts else "" search_result = SearchResult( - title=result.get("title", ""), - url=result.get("url", ""), + title=result.get("title") or "", + url=result.get("url") or "", snippet=snippet, - date=None, # Parallel AI doesn't provide date in response - last_updated=None, # Parallel AI doesn't provide last_updated in response + date=result.get("publish_date"), + last_updated=None, ) results.append(search_result) diff --git a/tests/test_litellm/llms/parallel_ai/__init__.py b/tests/test_litellm/llms/parallel_ai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py new file mode 100644 index 00000000000..b5c1a86205b --- /dev/null +++ b/tests/test_litellm/llms/parallel_ai/test_parallel_ai_search.py @@ -0,0 +1,324 @@ +""" +Tests for Parallel AI Search API integration (v1 endpoint). +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm + +MOCK_V1_RESPONSE = { + "search_id": "search_abc123", + "session_id": "session_xyz", + "results": [ + { + "url": "https://example.com/1", + "title": "Test Result 1", + "publish_date": "2026-01-15", + "excerpts": ["First excerpt.", "Second excerpt."], + }, + { + "url": "https://example.com/2", + "title": None, + "publish_date": None, + "excerpts": ["Only excerpt."], + }, + ], + "usage": [{"name": "search_advanced", "count": 1}], +} + + +def _mock_response(): + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = MOCK_V1_RESPONSE + return mock_response + + +class TestParallelAISearch: + @pytest.fixture(autouse=True) + def _set_api_key(self, monkeypatch): + monkeypatch.setenv("PARALLEL_API_KEY", "test-api-key") + monkeypatch.delenv("PARALLEL_AI_API_BASE", raising=False) + + @pytest.mark.asyncio + async def test_v1_endpoint_and_headers(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="latest developments in AI", + search_provider="parallel_ai", + ) + + call_args = mock_post.call_args + assert call_args.kwargs["url"] == "https://api.parallel.ai/v1/search" + + headers = call_args.kwargs.get("headers", {}) + assert headers["x-api-key"] == "test-api-key" + assert headers["Content-Type"] == "application/json" + assert "parallel-beta" not in headers + + @pytest.mark.asyncio + async def test_string_query_maps_to_search_queries_and_objective(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="latest developments in AI", + search_provider="parallel_ai", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["search_queries"] == ["latest developments in AI"] + assert json_data["objective"] == "latest developments in AI" + + @pytest.mark.asyncio + async def test_list_query_maps_to_search_queries(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query=["AI developments", "machine learning trends"], + search_provider="parallel_ai", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["search_queries"] == [ + "AI developments", + "machine learning trends", + ] + assert "objective" not in json_data + + @pytest.mark.asyncio + async def test_mode_param_passthrough(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + mode="turbo", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == "turbo" + + @pytest.mark.asyncio + async def test_default_mode_is_basic(self): + """v1 defaults to 'advanced' server-side; litellm must send 'basic' to keep v1beta's default tier and cost tracking accurate.""" + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == "basic" + + @pytest.mark.parametrize( + "processor,expected_mode", [("base", "basic"), ("pro", "advanced")] + ) + @pytest.mark.asyncio + async def test_legacy_processor_maps_to_mode(self, processor, expected_mode): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + processor=processor, + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == expected_mode + assert "processor" not in json_data + + @pytest.mark.asyncio + async def test_explicit_mode_wins_over_processor(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + mode="turbo", + processor="pro", + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["mode"] == "turbo" + assert "processor" not in json_data + + @pytest.mark.asyncio + async def test_top_level_v1_params_pass_through(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + session_id="session_123", + max_chars_total=4000, + max_tokens_per_page=1024, + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["session_id"] == "session_123" + assert json_data["max_chars_total"] == 4000 + assert "max_tokens_per_page" not in json_data + + @pytest.mark.asyncio + async def test_optional_params_nest_under_advanced_settings(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + max_results=5, + country="US", + search_domain_filter=["arxiv.org", "nature.com"], + exclude_domains=["reddit.com"], + max_chars_per_result=1500, + ) + + json_data = mock_post.call_args.kwargs.get("json") + advanced_settings = json_data["advanced_settings"] + assert advanced_settings["max_results"] == 5 + assert advanced_settings["location"] == "US" + assert advanced_settings["source_policy"]["include_domains"] == [ + "arxiv.org", + "nature.com", + ] + assert advanced_settings["source_policy"]["exclude_domains"] == [ + "reddit.com" + ] + assert advanced_settings["excerpt_settings"]["max_chars_per_result"] == 1500 + + assert "max_results" not in json_data + assert "source_policy" not in json_data + assert "search_domain_filter" not in json_data + assert "exclude_domains" not in json_data + assert "max_chars_per_result" not in json_data + assert "country" not in json_data + + @pytest.mark.asyncio + async def test_explicit_advanced_settings_take_precedence(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + max_results=5, + advanced_settings={"max_results": 7}, + ) + + json_data = mock_post.call_args.kwargs.get("json") + assert json_data["advanced_settings"]["max_results"] == 7 + + @pytest.mark.asyncio + async def test_response_transformation(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + response = await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + ) + + assert response.object == "search" + assert len(response.results) == 2 + + first = response.results[0] + assert first.title == "Test Result 1" + assert first.url == "https://example.com/1" + assert first.snippet == "First excerpt. ... Second excerpt." + assert first.date == "2026-01-15" + + second = response.results[1] + assert second.title == "" + assert second.snippet == "Only excerpt." + assert second.date is None + + @pytest.mark.parametrize( + "api_base", + [ + "https://proxy.internal.example.com", + "https://proxy.internal.example.com/", + "https://proxy.internal.example.com/v1", + "https://proxy.internal.example.com/v1/", + "https://proxy.internal.example.com/v1/search", + ], + ) + @pytest.mark.asyncio + async def test_custom_api_base_appends_v1_search(self, api_base): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _mock_response() + + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + api_base=api_base, + ) + + call_args = mock_post.call_args + assert ( + call_args.kwargs["url"] + == "https://proxy.internal.example.com/v1/search" + ) + + @pytest.mark.asyncio + async def test_missing_api_key_raises(self, monkeypatch): + monkeypatch.delenv("PARALLEL_API_KEY", raising=False) + monkeypatch.delenv("PARALLEL_AI_API_KEY", raising=False) + + with pytest.raises(Exception, match="PARALLEL_API_KEY"): + await litellm.asearch( + query="AI developments", + search_provider="parallel_ai", + )