mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(websearch_interception): match domain path rules on whole path segments
A rule like example.com/docs/ lost its trailing slash and was then matched as a plain prefix, so it also admitted example.com/docs-private. Match the path exactly or at a segment boundary instead
This commit is contained in:
parent
758950a432
commit
b7a62f4a71
2 changed files with 55 additions and 2 deletions
|
|
@ -34,12 +34,13 @@ def domain_host(domain: str) -> str:
|
|||
|
||||
|
||||
def _url_matches_domain(url: str, domain: str) -> bool:
|
||||
"""Anthropic web_search domain semantics: subdomains are included and an optional path is a prefix."""
|
||||
"""Anthropic web_search domain semantics: subdomains are included and an optional path matches whole path segments."""
|
||||
target: Final = urlsplit(url)
|
||||
host: Final = (target.hostname or "").lower()
|
||||
rule_host, rule_path = _split_domain(domain)
|
||||
host_matches: Final = host == rule_host or host.endswith(f".{rule_host}")
|
||||
return host_matches and target.path.startswith(rule_path)
|
||||
path_matches: Final = not rule_path or target.path == rule_path or target.path.startswith(f"{rule_path}/")
|
||||
return host_matches and path_matches
|
||||
|
||||
|
||||
def _url_passes_domain_filter(url: str, domains: WebSearchDomainFilter) -> bool:
|
||||
|
|
|
|||
|
|
@ -948,6 +948,58 @@ async def test_messages_web_search_honors_the_native_tool_domain_filter(
|
|||
assert mock_asearch.await_args.kwargs.get("search_domain_filter") == expected_provider_filter
|
||||
|
||||
|
||||
_PATH_RULE_URLS = (
|
||||
"https://example.com/docs/intro",
|
||||
"https://example.com/docs-private/x",
|
||||
"https://example.com/other/y",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("domain_field", "rule", "expected_urls"),
|
||||
[
|
||||
("allowed_domains", "example.com/docs/", ["https://example.com/docs/intro"]),
|
||||
("allowed_domains", "example.com/docs", ["https://example.com/docs/intro"]),
|
||||
("blocked_domains", "example.com/docs/", ["https://example.com/docs-private/x", "https://example.com/other/y"]),
|
||||
],
|
||||
)
|
||||
async def test_messages_web_search_domain_path_rules_match_whole_path_segments(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
domain_field: str,
|
||||
rule: str,
|
||||
expected_urls: list[str],
|
||||
):
|
||||
import litellm
|
||||
from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages
|
||||
from litellm.llms.base_llm.search.transformation import SearchResult
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _perplexity_router())
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"asearch",
|
||||
AsyncMock(
|
||||
return_value=SearchResponse(
|
||||
object="search",
|
||||
results=[SearchResult(title=url, url=url, snippet="snippet") for url in _PATH_RULE_URLS],
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(enabled_providers=["bedrock"])])
|
||||
|
||||
response = await anthropic_messages(
|
||||
max_tokens=512,
|
||||
messages=[{"role": "user", "content": "litellm"}],
|
||||
model="bedrock/converse/test-model",
|
||||
custom_llm_provider="bedrock",
|
||||
tools=[{"type": "web_search_20250305", "name": "web_search", domain_field: [rule]}],
|
||||
)
|
||||
|
||||
text = next(block["text"] for block in response["content"] if block["type"] == "text")
|
||||
assert [url for url in _PATH_RULE_URLS if url in text] == expected_urls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("domain_field", "expected_urls", "expected_provider_filter"),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue