style(websearch): apply ruff format to changed files

CI's lint gate checks ruff format, not black; black's output differs on
a few line splits. No logic changes.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Aidan Sinclair 2026-09-09 08:55:25 -04:00
parent 05908bbe57
commit 038a4bb384
4 changed files with 109 additions and 329 deletions

View file

@ -174,9 +174,7 @@ class _AcompletionNamedParams(TypedDict, total=False):
logprobs: ReadOnly[bool | None]
top_logprobs: ReadOnly[int | None]
deployment_id: ReadOnly[str | None]
reasoning_effort: ReadOnly[
Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None
]
reasoning_effort: ReadOnly[Literal["none", "minimal", "low", "medium", "high", "xhigh", "default"] | None]
verbosity: ReadOnly[Literal["low", "medium", "high"] | None]
safety_identifier: ReadOnly[str | None]
service_tier: ReadOnly[str | None]
@ -234,9 +232,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if enabled_providers is None:
self.enabled_providers = [LlmProviders.BEDROCK.value]
else:
self.enabled_providers = [
p.value if isinstance(p, LlmProviders) else p for p in enabled_providers
]
self.enabled_providers = [p.value if isinstance(p, LlmProviders) else p for p in enabled_providers]
self.search_tool_name = search_tool_name
self.max_agentic_loops = self._validated_max_agentic_loops(max_agentic_loops)
self._request_has_websearch = False # Track if current request has web search
@ -246,9 +242,7 @@ class WebSearchInterceptionLogger(CustomLogger):
"""
Reject loop ceilings the agentic loop cannot honor, at config load time.
"""
return validated_max_agentic_loops(
max_agentic_loops, field="websearch_interception_params.max_agentic_loops"
)
return validated_max_agentic_loops(max_agentic_loops, field="websearch_interception_params.max_agentic_loops")
async def try_short_circuit_search(
self,
@ -283,10 +277,7 @@ class WebSearchInterceptionLogger(CustomLogger):
# Check if provider is in enabled list
provider_str: Final = custom_llm_provider or ""
if (
self.enabled_providers is not None
and provider_str not in self.enabled_providers
):
if self.enabled_providers is not None and provider_str not in self.enabled_providers:
return None
# Only short-circuit for providers whose Anthropic Messages agentic loop
@ -302,15 +293,10 @@ class WebSearchInterceptionLogger(CustomLogger):
# web-search-only requests against it.
try:
provider_enum: Final = LlmProviders(provider_str)
anthropic_config: Final = (
ProviderConfigManager.get_provider_anthropic_messages_config(
model=model, provider=provider_enum
)
anthropic_config: Final = ProviderConfigManager.get_provider_anthropic_messages_config(
model=model, provider=provider_enum
)
if (
anthropic_config is not None
and anthropic_config.handles_web_search_natively()
):
if anthropic_config is not None and anthropic_config.handles_web_search_natively():
verbose_logger.debug(
"WebSearchInterception: Skipping short-circuit for %s (provider handles web search natively via the agentic loop)",
provider_str,
@ -355,13 +341,9 @@ class WebSearchInterceptionLogger(CustomLogger):
if kwargs is None:
search_result_text, structured = await self._execute_search(query)
else:
search_result_text, structured = await self._execute_search(
query, kwargs=kwargs
)
search_result_text, structured = await self._execute_search(query, kwargs=kwargs)
except Exception as e:
verbose_logger.error(
"WebSearchInterception: Short-circuit search failed: %s", e
)
verbose_logger.error("WebSearchInterception: Short-circuit search failed: %s", e)
search_result_text, structured = f"Search failed: {e}", None
content: Final[list[dict[str, object]]] = []
@ -421,14 +403,12 @@ class WebSearchInterceptionLogger(CustomLogger):
"litellm_params": kwargs.get("litellm_params", {}),
"model": kwargs.get("model", ""),
}
custom_llm_provider = call_kwargs_view[
"custom_llm_provider"
] or call_kwargs_view["litellm_params"].get("custom_llm_provider", "")
custom_llm_provider = call_kwargs_view["custom_llm_provider"] or call_kwargs_view["litellm_params"].get(
"custom_llm_provider", ""
)
if not custom_llm_provider:
try:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=call_kwargs_view["model"]
)
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=call_kwargs_view["model"])
except Exception:
custom_llm_provider = ""
if custom_llm_provider not in self.enabled_providers:
@ -447,9 +427,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if not has_websearch:
return None
verbose_logger.debug(
"WebSearchInterception: Converting native web_search tools to LiteLLM standard"
)
verbose_logger.debug("WebSearchInterception: Converting native web_search tools to LiteLLM standard")
# If the client sent an Anthropic-native web_search_* tool, mark the
# request so the agentic loop emits native web_search_tool_result
@ -479,9 +457,7 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs["tools"] = converted_tools
if kwargs.get("stream"):
verbose_logger.debug(
"WebSearchInterception: deployment hook converting stream=True to stream=False"
)
verbose_logger.debug("WebSearchInterception: deployment hook converting stream=True to stream=False")
kwargs["stream"] = False
kwargs["_websearch_interception_converted_stream"] = True
@ -494,34 +470,23 @@ class WebSearchInterceptionLogger(CustomLogger):
if not any(is_web_search_tool_responses(tool) for tool in tools):
return None
verbose_logger.debug(
"WebSearchInterception: Converting Responses web_search tools to LiteLLM standard"
)
verbose_logger.debug("WebSearchInterception: Converting Responses web_search tools to LiteLLM standard")
converted_tools: Final = [
(
get_litellm_web_search_tool_responses()
if is_web_search_tool_responses(tool)
else tool
)
for tool in tools
(get_litellm_web_search_tool_responses() if is_web_search_tool_responses(tool) else tool) for tool in tools
]
converted_kwargs: Final = {**kwargs, "tools": converted_tools}
if kwargs.get("stream"):
verbose_logger.debug(
"WebSearchInterception: deployment hook converting stream=True to stream=False"
)
verbose_logger.debug("WebSearchInterception: deployment hook converting stream=True to stream=False")
converted_kwargs["stream"] = False
converted_kwargs["_websearch_interception_converted_stream"] = True
return converted_kwargs
@classmethod
def from_config_yaml(
cls, config: WebSearchInterceptionConfig
) -> "WebSearchInterceptionLogger":
def from_config_yaml(cls, config: WebSearchInterceptionConfig) -> "WebSearchInterceptionLogger":
"""
Initialize WebSearchInterceptionLogger from proxy config.yaml parameters.
@ -576,9 +541,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return tool.get("name")
@classmethod
def _sync_forced_tool_choice(
cls, tool_choice: object, converted_tools: Sequence[Mapping[str, object]]
) -> object:
def _sync_forced_tool_choice(cls, tool_choice: object, converted_tools: Sequence[Mapping[str, object]]) -> object:
"""Repoint a forced ``tool_choice`` at ``litellm_web_search`` when it
names a web-search tool that was just converted away.
@ -595,9 +558,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return tool_choice
return {**tool_choice, "name": LITELLM_WEB_SEARCH_TOOL_NAME}
async def async_pre_request_hook(
self, model: str, messages: list[dict], kwargs: dict
) -> dict | None:
async def async_pre_request_hook(self, model: str, messages: list[dict], kwargs: dict) -> dict | None:
"""
Pre-request hook to convert native web search tools to LiteLLM standard.
@ -613,9 +574,7 @@ class WebSearchInterceptionLogger(CustomLogger):
Modified kwargs dict with converted tools, or None if no modifications needed
"""
# Check if this request is for an enabled provider
custom_llm_provider: Final = kwargs.get("litellm_params", {}).get(
"custom_llm_provider", ""
)
custom_llm_provider: Final = kwargs.get("litellm_params", {}).get("custom_llm_provider", "")
verbose_logger.debug(
"WebSearchInterception: Pre-request hook called - custom_llm_provider=%s - enabled_providers=%s",
@ -623,10 +582,7 @@ class WebSearchInterceptionLogger(CustomLogger):
self.enabled_providers or "ALL",
)
if (
self.enabled_providers is not None
and custom_llm_provider not in self.enabled_providers
):
if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers:
verbose_logger.debug(
"WebSearchInterception: Skipping - provider %s not in %s",
custom_llm_provider,
@ -651,9 +607,7 @@ class WebSearchInterceptionLogger(CustomLogger):
deployment_max_agentic_loops: Final = kwargs.get("max_agentic_loops")
if self.max_agentic_loops is not None and deployment_max_agentic_loops is None:
kwargs["max_agentic_loops"] = (
self.max_agentic_loops
) # rebind-ok: this hook returns the kwargs it edits
kwargs["max_agentic_loops"] = self.max_agentic_loops # rebind-ok: this hook returns the kwargs it edits
# If the client sent an Anthropic-native web_search_* tool, mark the
# request so the agentic loop emits native web_search_tool_result
@ -685,15 +639,11 @@ class WebSearchInterceptionLogger(CustomLogger):
)
if "tool_choice" in kwargs:
kwargs["tool_choice"] = self._sync_forced_tool_choice(
kwargs.get("tool_choice"), converted_tools
)
kwargs["tool_choice"] = self._sync_forced_tool_choice(kwargs.get("tool_choice"), converted_tools)
# Also convert here for direct callers that bypass the deployment hook.
if kwargs.get("stream"):
verbose_logger.debug(
"WebSearchInterception: Converting stream=True to stream=False"
)
verbose_logger.debug("WebSearchInterception: Converting stream=True to stream=False")
kwargs["stream"] = False
kwargs["_websearch_interception_converted_stream"] = True
@ -741,10 +691,7 @@ class WebSearchInterceptionLogger(CustomLogger):
# Check if provider should be intercepted
# Note: custom_llm_provider is already normalized by get_llm_provider()
# (e.g., "bedrock/invoke/..." -> "bedrock")
if (
self.enabled_providers is not None
and custom_llm_provider not in self.enabled_providers
):
if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers:
verbose_logger.debug(
"WebSearchInterception: Skipping provider %s (not in enabled list: %s)",
custom_llm_provider,
@ -766,9 +713,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
if not should_intercept:
verbose_logger.debug(
"WebSearchInterception: No WebSearch tool_use detected in response"
)
verbose_logger.debug("WebSearchInterception: No WebSearch tool_use detected in response")
return False, {}
verbose_logger.debug(
@ -801,9 +746,7 @@ class WebSearchInterceptionLogger(CustomLogger):
thinking_block_dict: dict = {"type": block_type}
if block_type == "thinking":
thinking_block_dict["thinking"] = getattr(block, "thinking", "")
thinking_block_dict["signature"] = getattr(
block, "signature", ""
)
thinking_block_dict["signature"] = getattr(block, "signature", "")
else: # redacted_thinking
thinking_block_dict["data"] = getattr(block, "data", "")
thinking_blocks.append(thinking_block_dict)
@ -848,10 +791,7 @@ class WebSearchInterceptionLogger(CustomLogger):
verbose_logger.debug("WebSearchInterception: Response type: %s", type(response))
# Check if provider should be intercepted
if (
self.enabled_providers is not None
and custom_llm_provider not in self.enabled_providers
):
if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers:
verbose_logger.debug(
"WebSearchInterception: Skipping provider %s (not in enabled list: %s)",
custom_llm_provider,
@ -860,13 +800,9 @@ class WebSearchInterceptionLogger(CustomLogger):
return False, {}
# Check if tools include any web search tool (strict check for chat completions)
has_websearch_tool: Final = any(
is_web_search_tool_chat_completion(t) for t in (tools or [])
)
has_websearch_tool: Final = any(is_web_search_tool_chat_completion(t) for t in (tools or []))
if not has_websearch_tool:
verbose_logger.debug(
"WebSearchInterception: No litellm_web_search tool in request"
)
verbose_logger.debug("WebSearchInterception: No litellm_web_search tool in request")
return False, {}
# Detect WebSearch tool_calls in response (OpenAI format)
@ -877,9 +813,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
if not should_intercept:
verbose_logger.debug(
"WebSearchInterception: No WebSearch tool_calls detected in response"
)
verbose_logger.debug("WebSearchInterception: No WebSearch tool_calls detected in response")
return False, {}
verbose_logger.debug(
@ -913,10 +847,7 @@ class WebSearchInterceptionLogger(CustomLogger):
stream,
)
if (
self.enabled_providers is not None
and custom_llm_provider not in self.enabled_providers
):
if self.enabled_providers is not None and custom_llm_provider not in self.enabled_providers:
verbose_logger.debug(
"WebSearchInterception: Skipping provider %s (not in enabled list: %s)",
custom_llm_provider,
@ -924,13 +855,9 @@ class WebSearchInterceptionLogger(CustomLogger):
)
return False, {}
has_websearch_tool: Final = any(
is_web_search_tool_responses(t) for t in (tools or [])
)
has_websearch_tool: Final = any(is_web_search_tool_responses(t) for t in (tools or []))
if not has_websearch_tool:
verbose_logger.debug(
"WebSearchInterception: No litellm_web_search tool in responses request"
)
verbose_logger.debug("WebSearchInterception: No litellm_web_search tool in responses request")
return False, {}
should_intercept, tool_calls = WebSearchTransformation.transform_request(
@ -940,9 +867,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
if not should_intercept:
verbose_logger.debug(
"WebSearchInterception: No WebSearch function_call detected in responses output"
)
verbose_logger.debug("WebSearchInterception: No WebSearch function_call detected in responses output")
return False, {}
verbose_logger.debug(
@ -1053,11 +978,9 @@ class WebSearchInterceptionLogger(CustomLogger):
# (while we still have the structured SearchResponse list) and stash
# them on plan metadata for the post-hook to inject.
if kwargs.get(WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY):
metadata[WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY] = (
self._build_native_result_blocks(
tool_calls=tool_calls,
structured_results=structured_results,
)
metadata[WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY] = self._build_native_result_blocks(
tool_calls=tool_calls,
structured_results=structured_results,
)
return AgenticLoopPlan(
@ -1083,9 +1006,7 @@ class WebSearchInterceptionLogger(CustomLogger):
render citations / sources alongside the model's textual reply.
"""
metadata_view: Final[_PlanMetadataView] = {
"websearch_native_blocks": plan.metadata.get(
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY
)
"websearch_native_blocks": plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)
}
native_blocks: Final = metadata_view["websearch_native_blocks"]
if not native_blocks:
@ -1110,9 +1031,7 @@ class WebSearchInterceptionLogger(CustomLogger):
for i, tool_call in enumerate(tool_calls)
for block in WebSearchInterceptionLogger._native_result_pair(
query=WebSearchInterceptionLogger._tool_call_query(tool_call),
search_response=(
structured_results[i] if i < len(structured_results) else None
),
search_response=(structured_results[i] if i < len(structured_results) else None),
)
)
@ -1131,9 +1050,7 @@ class WebSearchInterceptionLogger(CustomLogger):
) -> tuple[Mapping[str, object], Mapping[str, object]]:
tool_use_id: Final = f"srvtoolu_{uuid.uuid4().hex}"
return (
AnthropicServerToolUseBlock(
id=tool_use_id, input=AnthropicSearchQuery(query=query)
).model_dump(),
AnthropicServerToolUseBlock(id=tool_use_id, input=AnthropicSearchQuery(query=query)).model_dump(),
WebSearchTransformation.build_web_search_tool_result_block(
tool_use_id=tool_use_id,
search_response=search_response,
@ -1141,9 +1058,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
@staticmethod
def _inject_native_blocks(
response: _ResponseT, native_blocks: Sequence[Mapping[str, object]]
) -> _ResponseT:
def _inject_native_blocks(response: _ResponseT, native_blocks: Sequence[Mapping[str, object]]) -> _ResponseT:
"""Prepend native blocks to response content, dict or object form."""
if not native_blocks:
return response
@ -1153,9 +1068,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return response
existing = getattr(response, _RESPONSE_CONTENT_FIELD, None) or []
try:
setattr(
response, _RESPONSE_CONTENT_FIELD, list(native_blocks) + list(existing)
)
setattr(response, _RESPONSE_CONTENT_FIELD, list(native_blocks) + list(existing))
except (AttributeError, TypeError):
# Object refused write — fall through and leave the response
# untouched rather than crash the request.
@ -1269,8 +1182,7 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs=kwargs,
rich=self._rich_search_input(tool_call["input"]),
)
if isinstance(tool_call.get("input"), dict)
and tool_call["input"].get("query")
if isinstance(tool_call.get("input"), dict) and tool_call["input"].get("query")
else self._create_empty_search_result()
)
for tool_call in tool_calls
@ -1280,13 +1192,9 @@ class WebSearchInterceptionLogger(CustomLogger):
"WebSearchInterception: Executing %s responses search(es) in parallel",
len(search_tasks),
)
search_results: Final = await asyncio.gather(
*search_tasks, return_exceptions=True
)
search_results: Final = await asyncio.gather(*search_tasks, return_exceptions=True)
search_texts: Final = [
self._extract_search_text(result) for result in search_results
]
search_texts: Final = [self._extract_search_text(result) for result in search_results]
followup_items: Final = [
item
@ -1367,16 +1275,12 @@ class WebSearchInterceptionLogger(CustomLogger):
@staticmethod
def _extract_search_text(result: object) -> str:
if isinstance(result, Exception):
verbose_logger.error(
"WebSearchInterception: Responses search failed with error: %s", result
)
verbose_logger.error("WebSearchInterception: Responses search failed with error: %s", result)
return f"Search failed: {result}"
if isinstance(result, tuple) and len(result) == 2:
text_value, _ = result
return text_value if isinstance(text_value, str) else str(text_value)
verbose_logger.debug(
"WebSearchInterception: Unexpected search result type %s", type(result)
)
verbose_logger.debug("WebSearchInterception: Unexpected search result type %s", type(result))
return str(result)
@staticmethod
@ -1427,9 +1331,7 @@ class WebSearchInterceptionLogger(CustomLogger):
"""
_internal_keys: Final = {"litellm_logging_obj"}
return {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception") and k not in _internal_keys
k: v for k, v in kwargs.items() if not k.startswith("_websearch_interception") and k not in _internal_keys
}
async def _execute_agentic_loop(
@ -1449,9 +1351,7 @@ class WebSearchInterceptionLogger(CustomLogger):
messages=messages,
tool_calls=tool_calls,
thinking_blocks=thinking_blocks,
anthropic_messages_optional_request_params=dict[str, object](
anthropic_messages_optional_request_params
),
anthropic_messages_optional_request_params=dict[str, object](anthropic_messages_optional_request_params),
logging_obj=logging_obj,
kwargs=dict[str, object](kwargs),
)
@ -1469,15 +1369,13 @@ class WebSearchInterceptionLogger(CustomLogger):
max_tokens = cast(int, kwargs.get("max_tokens", 1024))
patch_kwargs: Final = dict[str, object](request_patch.kwargs)
response: AnthropicMessagesResponse | AsyncIterator[object] = (
await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=request_patch.messages,
model=request_patch.model or model,
**_NO_ACREATE_NAMED,
**optional_params,
**patch_kwargs,
)
response: AnthropicMessagesResponse | AsyncIterator[object] = await anthropic_messages.acreate(
max_tokens=max_tokens,
messages=request_patch.messages,
model=request_patch.model or model,
**_NO_ACREATE_NAMED,
**optional_params,
**patch_kwargs,
)
# Legacy path: the new path goes through the typed plan + core
@ -1517,9 +1415,7 @@ class WebSearchInterceptionLogger(CustomLogger):
for tool_call in tool_calls:
query = tool_call["input"].get("query")
if query:
verbose_logger.debug(
"WebSearchInterception: Queuing search for query='%s'", query
)
verbose_logger.debug("WebSearchInterception: Queuing search for query='%s'", query)
search_tasks.append(
self._execute_search(
query,
@ -1528,9 +1424,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
)
else:
verbose_logger.debug(
"WebSearchInterception: Tool call %s has no query", tool_call["id"]
)
verbose_logger.debug("WebSearchInterception: Tool call %s has no query", tool_call["id"])
# Add empty result for tools without query
search_tasks.append(self._create_empty_search_result())
@ -1539,9 +1433,7 @@ class WebSearchInterceptionLogger(CustomLogger):
"WebSearchInterception: Executing %s search(es) in parallel",
len(search_tasks),
)
search_results: Final = await asyncio.gather(
*search_tasks, return_exceptions=True
)
search_results: Final = await asyncio.gather(*search_tasks, return_exceptions=True)
# Split the gathered (text, structured) tuples into two parallel lists.
# The text list feeds the follow-up model call; the structured list
@ -1550,23 +1442,13 @@ class WebSearchInterceptionLogger(CustomLogger):
structured_results: Final[list[SearchResponse | None]] = []
for i, result in enumerate(search_results):
if isinstance(result, Exception):
verbose_logger.error(
"WebSearchInterception: Search %s failed with error: %s", i, result
)
verbose_logger.error("WebSearchInterception: Search %s failed with error: %s", i, result)
final_search_results.append(f"Search failed: {result}")
structured_results.append(None)
elif isinstance(result, tuple) and len(result) == 2:
text_value, structured_value = result
final_search_results.append(
cast(str, text_value)
if isinstance(text_value, str)
else str(text_value)
)
structured_results.append(
structured_value
if isinstance(structured_value, SearchResponse)
else None
)
final_search_results.append(cast(str, text_value) if isinstance(text_value, str) else str(text_value))
structured_results.append(structured_value if isinstance(structured_value, SearchResponse) else None)
else:
# Defensive: legacy callers / unexpected shape — preserve text,
# drop structure.
@ -1591,15 +1473,11 @@ class WebSearchInterceptionLogger(CustomLogger):
]
# Correlation context for structured logging
_call_id: Final = getattr(logging_obj, "litellm_call_id", None) or kwargs.get(
"litellm_call_id", "unknown"
)
_call_id: Final = getattr(logging_obj, "litellm_call_id", None) or kwargs.get("litellm_call_id", "unknown")
full_model_name = model # safe default before try block
max_tokens: Final = self._resolve_max_tokens(
anthropic_messages_optional_request_params, kwargs
)
max_tokens: Final = self._resolve_max_tokens(anthropic_messages_optional_request_params, kwargs)
verbose_logger.debug(
"WebSearchInterception: Using max_tokens=%s for follow-up request",
@ -1607,17 +1485,13 @@ class WebSearchInterceptionLogger(CustomLogger):
)
optional_params_without_max_tokens: Final = {
k: v
for k, v in anthropic_messages_optional_request_params.items()
if k != "max_tokens"
k: v for k, v in anthropic_messages_optional_request_params.items() if k != "max_tokens"
}
kwargs_for_followup: Final = self._prepare_followup_kwargs(kwargs)
if logging_obj is not None:
agentic_view: Final[_AgenticLoopParamsView] = {
"agentic_loop_params": logging_obj.model_call_details.get(
"agentic_loop_params", {}
)
"agentic_loop_params": logging_obj.model_call_details.get("agentic_loop_params", {})
}
full_model_name = agentic_view["agentic_loop_params"].get("model", model)
verbose_logger.debug(
@ -1647,9 +1521,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if not isinstance(tool_input, Mapping):
return None
objective = tool_input.get("objective")
valid_objective = (
objective if isinstance(objective, str) and objective.strip() else None
)
valid_objective = objective if isinstance(objective, str) and objective.strip() else None
raw_queries = tool_input.get("search_queries")
valid_queries: list[str] | None = None
if isinstance(raw_queries, Sequence) and not isinstance(raw_queries, str):
@ -1677,9 +1549,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return False
# SearchProviders is a str enum, so an unknown provider string simply
# misses the config map and returns None rather than raising.
config = ProviderConfigManager.get_provider_search_config(
search_provider
) # pyright: ignore[reportArgumentType] -- SearchProviders is a str enum, so the router's provider string hashes to the matching member; unknown strings miss the map and yield None
config = ProviderConfigManager.get_provider_search_config(search_provider) # pyright: ignore[reportArgumentType] -- SearchProviders is a str enum, so the router's provider string hashes to the matching member; unknown strings miss the map and yield None
return config is not None and config.supports_rich_search_input()
async def _execute_search(
@ -1709,21 +1579,13 @@ class WebSearchInterceptionLogger(CustomLogger):
)
llm_router = None
search_tool: Final = self._select_search_tool_from_router(
llm_router=llm_router
)
search_tool: Final = self._select_search_tool_from_router(llm_router=llm_router)
search_provider: str | None = None
search_litellm_params: Mapping[str, object] = {}
search_tool_name: Final = self._selected_search_tool_name(
search_tool=search_tool
)
search_tool_name: Final = self._selected_search_tool_name(search_tool=search_tool)
if search_tool is not None:
await self._authorize_search_tool(
search_tool=search_tool, kwargs=kwargs
)
tool_params: Final[_SearchToolLitellmParams] = (
search_tool.get("litellm_params", {}) or {}
)
await self._authorize_search_tool(search_tool=search_tool, kwargs=kwargs)
tool_params: Final[_SearchToolLitellmParams] = search_tool.get("litellm_params", {}) or {}
search_litellm_params = dict[str, object](tool_params)
search_provider = tool_params.get("search_provider")
@ -1783,9 +1645,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
# Format using transformation function
search_result_text: Final = WebSearchTransformation.format_search_response(
result
)
search_result_text: Final = WebSearchTransformation.format_search_response(result)
verbose_logger.debug(
"WebSearchInterception: Search completed for '%s', got %s chars",
@ -1794,9 +1654,7 @@ class WebSearchInterceptionLogger(CustomLogger):
)
return search_result_text, result
except Exception as e:
verbose_logger.error(
"WebSearchInterception: Search failed for '%s': %s", query, e
)
verbose_logger.error("WebSearchInterception: Search failed for '%s': %s", query, e)
raise
async def _authorize_search_tool(
@ -1856,9 +1714,7 @@ class WebSearchInterceptionLogger(CustomLogger):
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
user_api_key_metadata: Final[StandardLoggingUserAPIKeyMetadata] = (
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(
user_api_key_dict=user_api_key_auth
)
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_auth)
)
return { # mutable-ok: litellm's metadata channel is a plain dict its logging path reads and enriches
**user_api_key_metadata,
@ -1874,11 +1730,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if search_tool is None:
return None
search_tool_name: Final = search_tool.get("search_tool_name")
return (
search_tool_name
if isinstance(search_tool_name, str) and search_tool_name
else None
)
return search_tool_name if isinstance(search_tool_name, str) and search_tool_name else None
@staticmethod
def _get_user_api_key_auth_from_kwargs(
@ -1889,10 +1741,7 @@ class WebSearchInterceptionLogger(CustomLogger):
for metadata_key in ("metadata", "litellm_metadata"):
metadata = kwargs.get(metadata_key)
if (
isinstance(metadata, dict)
and metadata.get("user_api_key_auth") is not None
):
if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None:
return metadata["user_api_key_auth"]
litellm_params: Final = kwargs.get("litellm_params")
@ -1901,23 +1750,16 @@ class WebSearchInterceptionLogger(CustomLogger):
for metadata_key in ("metadata", "litellm_metadata"):
metadata = litellm_params.get(metadata_key)
if (
isinstance(metadata, dict)
and metadata.get("user_api_key_auth") is not None
):
if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None:
return metadata["user_api_key_auth"]
return None
def _select_search_tool_from_router(
self, llm_router: object
) -> "_SearchToolConfig | None":
def _select_search_tool_from_router(self, llm_router: object) -> "_SearchToolConfig | None":
if llm_router is None or not hasattr(llm_router, "search_tools"):
return None
search_tools: Final = tuple(getattr(llm_router, "search_tools", None) or ())
return self._select_search_tool_from_list(
search_tools=search_tools, source="router"
)
return self._select_search_tool_from_list(search_tools=search_tools, source="router")
def _select_search_tool_from_list(
self,
@ -1926,14 +1768,10 @@ class WebSearchInterceptionLogger(CustomLogger):
) -> "_SearchToolConfig | None":
if self.search_tool_name:
matching_tools: Final = tuple(
tool
for tool in search_tools
if tool.get("search_tool_name") == self.search_tool_name
tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name
)
if matching_tools:
search_provider = (
matching_tools[0].get("litellm_params", {}) or {}
).get("search_provider")
search_provider = (matching_tools[0].get("litellm_params", {}) or {}).get("search_provider")
verbose_logger.debug(
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
self.search_tool_name,
@ -1949,9 +1787,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if search_tools:
first_tool: Final = search_tools[0]
search_provider = (first_tool.get("litellm_params", {}) or {}).get(
"search_provider"
)
search_provider = (first_tool.get("litellm_params", {}) or {}).get("search_provider")
verbose_logger.debug(
"WebSearchInterception: Using first available search tool from %s with provider '%s'",
source,
@ -2024,14 +1860,8 @@ class WebSearchInterceptionLogger(CustomLogger):
query = args.get("query")
if query:
verbose_logger.debug(
"WebSearchInterception: Queuing search for query='%s'", query
)
search_tasks.append(
self._execute_search(
query, kwargs=kwargs, rich=self._rich_search_input(tool_args)
)
)
verbose_logger.debug("WebSearchInterception: Queuing search for query='%s'", query)
search_tasks.append(self._execute_search(query, kwargs=kwargs, rich=self._rich_search_input(tool_args)))
else:
verbose_logger.debug(
"WebSearchInterception: Tool call %s has no query",
@ -2045,26 +1875,18 @@ class WebSearchInterceptionLogger(CustomLogger):
"WebSearchInterception: Executing %s search(es) in parallel",
len(search_tasks),
)
search_results: Final = await asyncio.gather(
*search_tasks, return_exceptions=True
)
search_results: Final = await asyncio.gather(*search_tasks, return_exceptions=True)
# Chat-completion path only needs text — OpenAI tool_result format
# has no equivalent of Anthropic's web_search_tool_result block.
final_search_results: Final[list[str]] = []
for i, result in enumerate(search_results):
if isinstance(result, Exception):
verbose_logger.error(
"WebSearchInterception: Search %s failed with error: %s", i, result
)
verbose_logger.error("WebSearchInterception: Search %s failed with error: %s", i, result)
final_search_results.append(f"Search failed: {result}")
elif isinstance(result, tuple) and len(result) == 2:
text_value, _ = result
final_search_results.append(
cast(str, text_value)
if isinstance(text_value, str)
else str(text_value)
)
final_search_results.append(cast(str, text_value) if isinstance(text_value, str) else str(text_value))
else:
verbose_logger.debug(
"WebSearchInterception: Unexpected result type %s at index %s",
@ -2086,9 +1908,7 @@ class WebSearchInterceptionLogger(CustomLogger):
# Make follow-up request with search results
# For OpenAI format, tool_messages_or_user is a list of tool messages
if response_format == "openai":
follow_up_messages = (
messages + [assistant_message] + cast(list[dict], tool_messages_or_user)
)
follow_up_messages = messages + [assistant_message] + cast(list[dict], tool_messages_or_user)
else:
# For Anthropic format (shouldn't happen in this method, but handle it)
follow_up_messages = messages + [
@ -2096,9 +1916,7 @@ class WebSearchInterceptionLogger(CustomLogger):
cast(dict, tool_messages_or_user),
]
verbose_logger.debug(
"WebSearchInterception: Making follow-up chat completion request with search results"
)
verbose_logger.debug("WebSearchInterception: Making follow-up chat completion request with search results")
verbose_logger.debug(
"WebSearchInterception: Follow-up messages count: %s",
len(follow_up_messages),
@ -2115,9 +1933,7 @@ class WebSearchInterceptionLogger(CustomLogger):
"custom_prompt_dict",
}
kwargs_for_followup: Final = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception") and k not in internal_params
k: v for k, v in kwargs.items() if not k.startswith("_websearch_interception") and k not in internal_params
}
full_model_name = model
@ -2190,9 +2006,7 @@ class WebSearchInterceptionLogger(CustomLogger):
websearch_params: WebSearchInterceptionConfig = {}
if "websearch_interception_params" in litellm_settings:
settings_view: Final[_WebSearchSettingsView] = {
"websearch_interception_params": litellm_settings[
"websearch_interception_params"
]
"websearch_interception_params": litellm_settings["websearch_interception_params"]
}
websearch_params = settings_view["websearch_interception_params"]
elif "websearch_interception" in callback_specific_params and isinstance(

View file

@ -197,20 +197,12 @@ class BaseSearchConfig:
def sign_request(
self,
headers: dict[
str, str
], # mutable-ok: matches the request header dict every other hook on this base takes
optional_params: dict[
str, object
], # mutable-ok: matches every other hook on this base
request_data: (
dict[str, object] | list[dict[str, object]]
), # mutable-ok: transform_search_request's body
headers: dict[str, str], # mutable-ok: matches the request header dict every other hook on this base takes
optional_params: dict[str, object], # mutable-ok: matches every other hook on this base
request_data: (dict[str, object] | list[dict[str, object]]), # mutable-ok: transform_search_request's body
api_base: str,
api_key: str | None = None,
) -> tuple[
dict[str, str], bytes | None
]: # mutable-ok: the handler passes these headers straight to httpx
) -> tuple[dict[str, str], bytes | None]: # mutable-ok: the handler passes these headers straight to httpx
"""
OPTIONAL
@ -270,9 +262,7 @@ class BaseSearchConfig:
Returns:
Dict with request data
"""
raise NotImplementedError(
"transform_search_request must be implemented by provider"
)
raise NotImplementedError("transform_search_request must be implemented by provider")
def transform_search_response(
self,
@ -284,9 +274,7 @@ class BaseSearchConfig:
Transform provider-specific Search response to standard format.
Override in provider-specific implementations.
"""
raise NotImplementedError(
"transform_search_response must be implemented by provider"
)
raise NotImplementedError("transform_search_response must be implemented by provider")
def get_error_class(
self,

View file

@ -110,9 +110,7 @@ class ParallelAISearchConfig(BaseSearchConfig):
default_api_base=self.PARALLEL_AI_API_BASE,
)
if not resolved_api_key:
raise ValueError(
"PARALLEL_API_KEY is not set. Set `PARALLEL_API_KEY` environment variable."
)
raise ValueError("PARALLEL_API_KEY is not set. Set `PARALLEL_API_KEY` environment variable.")
headers["x-api-key"] = resolved_api_key
headers["Content-Type"] = "application/json"
return headers
@ -124,11 +122,7 @@ class ParallelAISearchConfig(BaseSearchConfig):
data: dict | list[dict] | None = None,
**kwargs,
) -> str:
resolved_api_base: Final = (
api_base
or get_secret_str("PARALLEL_AI_API_BASE")
or self.PARALLEL_AI_API_BASE
)
resolved_api_base: Final = api_base or get_secret_str("PARALLEL_AI_API_BASE") or self.PARALLEL_AI_API_BASE
trimmed: Final = resolved_api_base.rstrip("/")
if trimmed.endswith("/v1/search"):
@ -195,9 +189,7 @@ class ParallelAISearchConfig(BaseSearchConfig):
advanced_settings["location"] = params.pop("location")
if "max_chars_per_result" in params:
advanced_settings["excerpt_settings"] = {
"max_chars_per_result": params.pop("max_chars_per_result")
}
advanced_settings["excerpt_settings"] = {"max_chars_per_result": params.pop("max_chars_per_result")}
if "fetch_policy" in params:
advanced_settings["fetch_policy"] = params.pop("fetch_policy")
@ -290,6 +282,4 @@ class ParallelAISearchConfig(BaseSearchConfig):
}
)
return SearchResponse.model_validate(
MappingProxyType({"results": results, "object": "search", **extra_fields})
)
return SearchResponse.model_validate(MappingProxyType({"results": results, "object": "search", **extra_fields}))

View file

@ -71,9 +71,7 @@ class TestRichInputExtraction:
}
def test_returns_none_when_only_query_present(self):
assert (
WebSearchInterceptionLogger._rich_search_input({"query": "plain"}) is None
)
assert WebSearchInterceptionLogger._rich_search_input({"query": "plain"}) is None
def test_returns_none_for_non_mapping_input(self):
assert WebSearchInterceptionLogger._rich_search_input(None) is None
@ -90,12 +88,7 @@ class TestRichInputExtraction:
def test_ignores_string_valued_search_queries(self):
# A string is a Sequence; it must not be treated as a list of queries.
assert (
WebSearchInterceptionLogger._rich_search_input(
{"query": "q", "search_queries": "not a list"}
)
is None
)
assert WebSearchInterceptionLogger._rich_search_input({"query": "q", "search_queries": "not a list"}) is None
class TestProviderSupport:
@ -107,10 +100,7 @@ class TestProviderSupport:
def test_unknown_provider_is_unsupported(self):
assert WebSearchInterceptionLogger._provider_supports_rich_search(None) is False
assert (
WebSearchInterceptionLogger._provider_supports_rich_search("not_a_provider")
is False
)
assert WebSearchInterceptionLogger._provider_supports_rich_search("not_a_provider") is False
class TestExecuteSearchShape:
@ -202,9 +192,7 @@ class TestCallSiteWiring:
monkeypatch.setattr(proxy_server, "llm_router", _mock_router("parallel_ai"))
monkeypatch.setattr(litellm, "asearch", mock_asearch)
tool_calls = [
{"id": "toolu_1", "name": "litellm_web_search", "input": dict(RICH_INPUT)}
]
tool_calls = [{"id": "toolu_1", "name": "litellm_web_search", "input": dict(RICH_INPUT)}]
await logger._build_anthropic_request_patch(
model="claude",
messages=[{"role": "user", "content": "hi"}],