mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
208f76f8ad
commit
7b939b4558
8 changed files with 251 additions and 61 deletions
|
|
@ -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,
|
||||
|
|
|
|||
7
litellm/llms/exa_ai/search/__init__.py
Normal file
7
litellm/llms/exa_ai/search/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
Exa AI Search API module.
|
||||
"""
|
||||
from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig
|
||||
|
||||
__all__ = ["ExaAISearchConfig"]
|
||||
|
||||
183
litellm/llms/exa_ai/search/transformation.py
Normal file
183
litellm/llms/exa_ai/search/transformation.py
Normal 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",
|
||||
)
|
||||
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
18
tests/search_tests/test_exa_ai_search.py
Normal file
18
tests/search_tests/test_exa_ai_search.py
Normal 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"
|
||||
|
||||
Loading…
Add table
Reference in a new issue