diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 2394ae4a43b..470a0c03a0d 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -11,6 +11,7 @@ import math import uuid from collections.abc import AsyncIterator, Mapping, Sequence from dataclasses import dataclass +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, cast @@ -103,9 +104,16 @@ _RESPONSE_CONTENT_FIELD: Final = "content" _ResponseT: Final = TypeVar("_ResponseT") +def _web_search_domain_strings(tool: Mapping[str, object], key: str) -> tuple[str, ...]: + value = tool.get(key) + if not isinstance(value, list): + return () + return tuple(item for item in value if isinstance(item, str) and item) + + def _extract_web_search_domain_filters( tools: Sequence[dict[str, object]], -) -> dict[str, list[str]] | None: +) -> Mapping[str, tuple[str, ...]] | None: """Collect ``allowed_domains`` / ``blocked_domains`` from web search tools. Anthropic-native ``web_search_*`` tools carry optional domain limits. The @@ -116,23 +124,20 @@ def _extract_web_search_domain_filters( Returns None when no web search tool carries a domain limit. """ - allowed: list[str] = [] - blocked: list[str] = [] - for tool in tools: - if not is_web_search_tool(tool): - continue - for key, bucket in (("allowed_domains", allowed), ("blocked_domains", blocked)): - value = tool.get(key) - if isinstance(value, list): - bucket.extend(item for item in value if isinstance(item, str) and item) + web_tools: Final = tuple(tool for tool in tools if is_web_search_tool(tool)) + allowed: Final = tuple( + chain.from_iterable(_web_search_domain_strings(tool, "allowed_domains") for tool in web_tools) + ) + blocked: Final = tuple( + chain.from_iterable(_web_search_domain_strings(tool, "blocked_domains") for tool in web_tools) + ) if not allowed and not blocked: return None - domain_filters: dict[str, list[str]] = {} - if allowed: - domain_filters["allowed_domains"] = allowed - if blocked: - domain_filters["blocked_domains"] = blocked - return domain_filters + if not blocked: + return MappingProxyType({"allowed_domains": allowed}) + if not allowed: + return MappingProxyType({"blocked_domains": blocked}) + return MappingProxyType({"allowed_domains": allowed, "blocked_domains": blocked}) class _PlanMetadataView(TypedDict): @@ -1658,18 +1663,19 @@ class WebSearchInterceptionLogger(CustomLogger): # allowlist and ``blocked_domains`` as '-'-prefixed exclusions — # the convention litellm.asearch()'s providers (e.g. Perplexity) # use for search_domain_filter. - search_domain_filter: list[str] | None = None request_domain_filters: Final = kwargs.get(WEBSEARCH_DOMAIN_FILTER_KEY) if kwargs is not None else None - if isinstance(request_domain_filters, dict): - allowed: Final[list[str]] = [ - item for item in request_domain_filters.get("allowed_domains", []) if isinstance(item, str) and item - ] - blocked: Final[list[str]] = [ - item for item in request_domain_filters.get("blocked_domains", []) if isinstance(item, str) and item - ] - if allowed or blocked: - search_domain_filter = allowed + [f"-{item}" for item in blocked] - verbose_logger.debug("WebSearchInterception: Applying domain filter %s", search_domain_filter) + domain_view: Final = ( + request_domain_filters if isinstance(request_domain_filters, Mapping) else MappingProxyType({}) + ) + allowed: Final = tuple( + item for item in domain_view.get("allowed_domains", ()) if isinstance(item, str) and item + ) + blocked: Final = tuple( + f"-{item}" for item in domain_view.get("blocked_domains", ()) if isinstance(item, str) and item + ) + search_domain_filter: Final = [*allowed, *blocked] or None # mutable-ok: JSON request array, not mutated + if search_domain_filter is not None: + verbose_logger.debug("WebSearchInterception: Applying domain filter %s", search_domain_filter) search_kwargs: Final = MappingProxyType( {**configured_search_kwargs, **parent_correlation.as_search_kwargs()} ) diff --git a/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py b/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py index 70a082f83a3..2cbffab71a9 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_domain_filters.py @@ -43,8 +43,8 @@ class TestExtractWebSearchDomainFilters: } ] assert _extract_web_search_domain_filters(tools) == { - "allowed_domains": ["docs.litellm.ai"], - "blocked_domains": ["twitter.com", "x.com"], + "allowed_domains": ("docs.litellm.ai",), + "blocked_domains": ("twitter.com", "x.com"), } def test_returns_none_without_domain_limits(self): @@ -65,7 +65,7 @@ class TestExtractWebSearchDomainFilters: "allowed_domains": ["docs.litellm.ai", 42, None], } ] - assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ["docs.litellm.ai"]} + assert _extract_web_search_domain_filters(tools) == {"allowed_domains": ("docs.litellm.ai",)} class TestDeploymentHookStashesDomainFilters: @@ -86,8 +86,8 @@ class TestDeploymentHookStashesDomainFilters: out = await logger.async_pre_call_deployment_hook(kwargs, None) assert out is not None assert kwargs[WEBSEARCH_DOMAIN_FILTER_KEY] == { - "allowed_domains": ["docs.litellm.ai"], - "blocked_domains": ["twitter.com"], + "allowed_domains": ("docs.litellm.ai",), + "blocked_domains": ("twitter.com",), } @pytest.mark.asyncio