From 7b939b4558115100cd4e9595a6c230089b2617e7 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 21 Oct 2025 17:06:23 -0700 Subject: [PATCH] [Feat] Add EXA AI Search API to LiteLLM (#15774) * add BaseSearchConfig * add BaseSearchConfig * validate_environment * fix handlers * add PerplexitySearchConfig * add PerplexitySearchConfig * add LiteLLM Search API module. * add BaseSearchConfig * add _build_search_optional_params * add search_testing * add BaseSearchTest * add TestPerplexitySearch * fix BASE * fix handler * add search API * add to init * fix: working perplexity search API * add _hidden_params to search * add TAVILY to LlmProviders * add TavilySearchConfig * add TavilySearchConfig * TestTavilySearch * add tavily transform * TestParallelAISearch * add LlmProviders * add ParallelAISearchConfig * add ParallelAISearchConfig * ParallelAISearchConfig * add EXA AI Search API * add ExaAISearchConfig * TestExaAISearch * add get_supported_perplexity_optional_params * add Exa AI Search API * add transform_search_request * add ExaAISearchConfig * fix linting errors * transform_search_request --- .../llms/base_llm/search/transformation.py | 16 ++ litellm/llms/exa_ai/search/__init__.py | 7 + litellm/llms/exa_ai/search/transformation.py | 183 ++++++++++++++++++ .../llms/parallel_ai/search/transformation.py | 24 ++- litellm/llms/tavily/search/transformation.py | 59 ++---- litellm/types/utils.py | 1 + litellm/utils.py | 4 + tests/search_tests/test_exa_ai_search.py | 18 ++ 8 files changed, 251 insertions(+), 61 deletions(-) create mode 100644 litellm/llms/exa_ai/search/__init__.py create mode 100644 litellm/llms/exa_ai/search/transformation.py create mode 100644 tests/search_tests/test_exa_ai_search.py diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py index a1a04bc0f7b..20bfd5d7104 100644 --- a/litellm/llms/base_llm/search/transformation.py +++ b/litellm/llms/base_llm/search/transformation.py @@ -49,6 +49,22 @@ class BaseSearchConfig: def __init__(self) -> None: pass + @staticmethod + def get_supported_perplexity_optional_params() -> set: + """ + Get the set of Perplexity unified search parameters. + These are the standard parameters that providers should transform from. + + Returns: + Set of parameter names that are part of the unified spec + """ + return { + "max_results", + "search_domain_filter", + "country", + "max_tokens_per_page", + } + def validate_environment( self, headers: Dict, diff --git a/litellm/llms/exa_ai/search/__init__.py b/litellm/llms/exa_ai/search/__init__.py new file mode 100644 index 00000000000..b647d2cd80f --- /dev/null +++ b/litellm/llms/exa_ai/search/__init__.py @@ -0,0 +1,7 @@ +""" +Exa AI Search API module. +""" +from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig + +__all__ = ["ExaAISearchConfig"] + diff --git a/litellm/llms/exa_ai/search/transformation.py b/litellm/llms/exa_ai/search/transformation.py new file mode 100644 index 00000000000..3678d33bed4 --- /dev/null +++ b/litellm/llms/exa_ai/search/transformation.py @@ -0,0 +1,183 @@ +""" +Calls Exa AI's /search endpoint to search the web. + +Exa AI API Reference: https://docs.exa.ai/reference/search +""" +from typing import Dict, List, Optional, TypedDict, Union + +import httpx + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + SearchResponse, + SearchResult, +) +from litellm.secret_managers.main import get_secret_str + + +class _ExaAISearchRequestRequired(TypedDict): + """Required fields for Exa AI Search API request.""" + query: str # Required - search query + + +class ExaAISearchRequest(_ExaAISearchRequestRequired, total=False): + """ + Exa AI Search API request format. + Based on: https://docs.exa.ai/reference/search + """ + type: str # Optional - search type ('keyword', 'neural', 'fast', 'auto'), default 'auto' + category: str # Optional - data category ('company', 'research paper', 'news', 'pdf', 'github', 'tweet', 'personal site', 'linkedin profile', 'financial report') + userLocation: str # Optional - two-letter ISO country code + numResults: int # Optional - number of results (max 100), default 10 + includeDomains: List[str] # Optional - list of domains to include + excludeDomains: List[str] # Optional - list of domains to exclude + startCrawlDate: str # Optional - crawl date filter (ISO 8601 format) + endCrawlDate: str # Optional - crawl date filter (ISO 8601 format) + startPublishedDate: str # Optional - published date filter (ISO 8601 format) + endPublishedDate: str # Optional - published date filter (ISO 8601 format) + includeText: List[str] # Optional - strings that must be present in webpage text + excludeText: List[str] # Optional - strings that must not be present in webpage text + context: Union[bool, dict] # Optional - format results for LLMs + moderation: bool # Optional - enable content moderation, default false + contents: dict # Optional - content retrieval options + + +class ExaAISearchConfig(BaseSearchConfig): + EXA_AI_API_BASE = "https://api.exa.ai" + + def validate_environment( + self, + headers: Dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + **kwargs, + ) -> Dict: + """ + Validate environment and return headers. + """ + api_key = api_key or get_secret_str("EXA_API_KEY") + if not api_key: + raise ValueError("EXA_API_KEY is not set. Set `EXA_API_KEY` environment variable.") + headers["x-api-key"] = api_key + headers["Content-Type"] = "application/json" + return headers + + def get_complete_url( + self, + api_base: Optional[str], + optional_params: dict, + **kwargs, + ) -> str: + """ + Get complete URL for Search endpoint. + """ + api_base = api_base or get_secret_str("EXA_API_BASE") or self.EXA_AI_API_BASE + + # Append "/search" to the api base if it's not already there + if not api_base.endswith("/search"): + api_base = f"{api_base}/search" + + return api_base + + + def transform_search_request( + self, + query: Union[str, List[str]], + optional_params: dict, + **kwargs, + ) -> Dict: + """ + Transform Search request to Exa AI API format. + + Transforms Perplexity unified spec parameters: + - query → query (same) + - max_results → numResults + - search_domain_filter → includeDomains + - country → userLocation + - max_tokens_per_page → (not applicable, ignored) + + All other Exa-specific parameters are passed through as-is. + + Args: + query: Search query (string or list of strings). Exa AI only supports single string queries. + optional_params: Optional parameters for the request + + Returns: + Dict with typed request data following ExaAISearchRequest spec + """ + if isinstance(query, list): + # Exa AI only supports single string queries, join with spaces + query = " ".join(query) + + request_data: ExaAISearchRequest = { + "query": query, + } + + # Transform Perplexity unified spec parameters to Exa format + if "max_results" in optional_params: + request_data["numResults"] = optional_params["max_results"] + + if "search_domain_filter" in optional_params: + request_data["includeDomains"] = optional_params["search_domain_filter"] + + if "country" in optional_params: + request_data["userLocation"] = optional_params["country"] + + # Convert to dict before dynamic key assignments + result_data = dict(request_data) + + # 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 + + # By default, request text content if not explicitly specified + # Exa AI doesn't return content/text unless explicitly requested + if "contents" not in result_data: + result_data["contents"] = {"text": True} + + return result_data + + def transform_search_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs, + ) -> SearchResponse: + """ + Transform Exa AI API response to LiteLLM unified SearchResponse format. + + Exa AI → LiteLLM mappings: + - results[].title → SearchResult.title + - results[].url → SearchResult.url + - results[].text → SearchResult.snippet + - results[].publishedDate → SearchResult.date + - No last_updated field in Exa AI response (set to None) + + Args: + raw_response: Raw httpx response from Exa AI API + logging_obj: Logging object for tracking + + Returns: + SearchResponse with standardized format + """ + response_json = raw_response.json() + + # Transform results to SearchResult objects + results = [] + for result in response_json.get("results", []): + search_result = SearchResult( + title=result.get("title", ""), + url=result.get("url", ""), + snippet=result.get("text", ""), # Exa AI uses "text" for content + date=result.get("publishedDate"), # ISO 8601 datetime string + last_updated=None, # Exa AI doesn't provide last_updated in response + ) + results.append(search_result) + + return SearchResponse( + results=results, + object="search", + ) + diff --git a/litellm/llms/parallel_ai/search/transformation.py b/litellm/llms/parallel_ai/search/transformation.py index 2f465c7bcb6..fa8ac9262c4 100644 --- a/litellm/llms/parallel_ai/search/transformation.py +++ b/litellm/llms/parallel_ai/search/transformation.py @@ -117,26 +117,16 @@ class ParallelAISearchConfig(BaseSearchConfig): """ request_data: ParallelAISearchRequest = {} - # Map query to objective (string) or search_queries (list) + # Map query to objective (string or list both become objective) if isinstance(query, list): - # List of queries -> search_queries request_data["objective"] = self._transform_query_to_objective(query) else: - # Single string -> objective (natural language description) request_data["objective"] = query - # Map max_results (same field name) + # Transform Perplexity unified spec parameters to Parallel AI format if "max_results" in optional_params: request_data["max_results"] = optional_params["max_results"] - # Map processor (same field name) - if "processor" in optional_params: - request_data["processor"] = optional_params["processor"] - - # Map max_chars_per_result (same field name) - if "max_chars_per_result" in optional_params: - request_data["max_chars_per_result"] = optional_params["max_chars_per_result"] - # Map domain filters to source_policy source_policy: _ParallelAISourcePolicy = {} @@ -149,7 +139,15 @@ class ParallelAISearchConfig(BaseSearchConfig): if source_policy: request_data["source_policy"] = source_policy - return dict(request_data) + # Convert to dict before dynamic key assignments + result_data = dict(request_data) + + # 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 + + return result_data def transform_search_response( self, diff --git a/litellm/llms/tavily/search/transformation.py b/litellm/llms/tavily/search/transformation.py index a622eb0753f..dda10679242 100644 --- a/litellm/llms/tavily/search/transformation.py +++ b/litellm/llms/tavily/search/transformation.py @@ -118,63 +118,26 @@ class TavilySearchConfig(BaseSearchConfig): "query": query, } - # Map max_results (same field name) + # Transform Perplexity unified spec parameters to Tavily format if "max_results" in optional_params: request_data["max_results"] = optional_params["max_results"] - # Map search_domain_filter → include_domains (different field name in Tavily) if "search_domain_filter" in optional_params: request_data["include_domains"] = optional_params["search_domain_filter"] - # Map exclude_domains (same field name) - if "exclude_domains" in optional_params: - request_data["exclude_domains"] = optional_params["exclude_domains"] - - # Map topic (same field name) - if "topic" in optional_params: - request_data["topic"] = optional_params["topic"] - - # Map search_depth (same field name) - if "search_depth" in optional_params: - request_data["search_depth"] = optional_params["search_depth"] - - # Map include_answer (same field name) - if "include_answer" in optional_params: - request_data["include_answer"] = optional_params["include_answer"] - - # Map include_raw_content (same field name) - if "include_raw_content" in optional_params: - request_data["include_raw_content"] = optional_params["include_raw_content"] - - # Map include_images (same field name) - if "include_images" in optional_params: - request_data["include_images"] = optional_params["include_images"] - - # Map include_image_descriptions (same field name) - if "include_image_descriptions" in optional_params: - request_data["include_image_descriptions"] = optional_params["include_image_descriptions"] - - # Map include_favicon (same field name) - if "include_favicon" in optional_params: - request_data["include_favicon"] = optional_params["include_favicon"] - - # Map time_range (same field name) - if "time_range" in optional_params: - request_data["time_range"] = optional_params["time_range"] - - # Map start_date (same field name) - if "start_date" in optional_params: - request_data["start_date"] = optional_params["start_date"] - - # Map end_date (same field name) - if "end_date" in optional_params: - request_data["end_date"] = optional_params["end_date"] - - # Map country (same field name, but lowercase for Tavily) if "country" in optional_params: + # Tavily expects lowercase country names request_data["country"] = optional_params["country"].lower() - return dict(request_data) + # Convert to dict before dynamic key assignments + result_data = dict(request_data) + + # 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 + + return result_data def transform_search_response( self, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 34b008bd5a8..45267f8c011 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2515,6 +2515,7 @@ class LlmProviders(str, Enum): TOPAZ = "topaz" TAVILY = "tavily" PARALLEL_AI = "parallel_ai" + EXA_AI = "exa_ai" ASSEMBLYAI = "assemblyai" GITHUB_COPILOT = "github_copilot" SNOWFLAKE = "snowflake" diff --git a/litellm/utils.py b/litellm/utils.py index c720a713c57..9afddd6bb6d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7627,6 +7627,9 @@ class ProviderConfigManager: """ Get Search configuration for a given provider. """ + from litellm.llms.exa_ai.search.transformation import ( + ExaAISearchConfig, + ) from litellm.llms.parallel_ai.search.transformation import ( ParallelAISearchConfig, ) @@ -7641,6 +7644,7 @@ class ProviderConfigManager: litellm.LlmProviders.PERPLEXITY: PerplexitySearchConfig, litellm.LlmProviders.TAVILY: TavilySearchConfig, litellm.LlmProviders.PARALLEL_AI: ParallelAISearchConfig, + litellm.LlmProviders.EXA_AI: ExaAISearchConfig, } config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None) if config_class is None: diff --git a/tests/search_tests/test_exa_ai_search.py b/tests/search_tests/test_exa_ai_search.py new file mode 100644 index 00000000000..f2455e5d4b5 --- /dev/null +++ b/tests/search_tests/test_exa_ai_search.py @@ -0,0 +1,18 @@ +import pytest +import litellm +from typing import List, Union + +from tests.search_tests.base_search_unit_tests import BaseSearchTest + + +class TestExaAISearch(BaseSearchTest): + """ + Tests for Exa AI Search functionality. + """ + + def get_custom_llm_provider(self) -> str: + """ + Return custom_llm_provider for Exa AI Search. + """ + return "exa_ai" +