[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
This commit is contained in:
Ishaan Jaff 2025-10-21 17:06:23 -07:00 • committed by GitHub
parent 208f76f8ad
commit 7b939b4558
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 251 additions and 61 deletions

View file

@ -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,

View file

@ -0,0 +1,7 @@
"""
Exa AI Search API module.
"""
from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig
__all__ = ["ExaAISearchConfig"]

View file

@ -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",
)

View file

@ -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,

View file

@ -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,

View file

@ -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"

View file

@ -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:

View file

@ -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"