This commit is contained in:
Vigilans 2026-06-27 18:40:41 +00:00 • committed by GitHub
commit 0c16542e73
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 305 additions and 2 deletions

View file

@ -36,6 +36,7 @@ from litellm.constants import (
STREAM_SSE_DATA_PREFIX,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.websearch_interception.tools import is_web_search_tool
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.llm_response_utils.get_headers import (
@ -918,7 +919,7 @@ class ProxyBaseLLMRequestProcessing:
return custom_headers
async def common_processing_pre_call_logic(
async def common_processing_pre_call_logic( # noqa: PLR0915
self,
request: Request,
general_settings: dict,
@ -1100,6 +1101,32 @@ class ProxyBaseLLMRequestProcessing:
):
self.data["model"] = user_api_key_dict.aliases[self.data["model"]]
### WEB SEARCH REDIRECT ###
# if request only contains web_search tools, check for redirect:
# 1. per-deployment force: always redirect to `force_websearch_model`
# 2. global fallback: redirect to `websearch_fallback_model` only when
# the current model group does not support web search
if (
isinstance(self.data.get("model"), str)
and llm_router is not None
and self.data.get("tools")
and all(
isinstance(t, dict) and is_web_search_tool(t)
for t in self.data["tools"]
)
):
for dep in llm_router.get_model_list(model_name=self.data["model"]) or []:
if model := dep.get("litellm_params", {}).get("force_websearch_model"):
self.data["model"] = model
break
if (
(dep_model := dep.get("litellm_params", {}).get("model", ""))
and not litellm.supports_web_search(dep_model)
and (model := getattr(litellm, "websearch_fallback_model", None))
):
self.data["model"] = model
break
self.data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)

View file

@ -236,6 +236,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
)
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
merge_reasoning_content_in_choices: Optional[bool] = False
force_websearch_model: Optional[str] = None
model_info: Optional[Dict] = None
mock_response: Optional[Union[str, ModelResponse, Exception, Any]] = None
@ -365,6 +366,8 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
drop_params: Optional[bool]
## RESPONSES API → CHAT COMPLETIONS BRIDGE ##
use_chat_completions_api: Optional[bool]
## WEB SEARCH REDIRECT ##
force_websearch_model: Optional[str]
## UNIFIED PROJECT/REGION ##
region_name: Optional[str]
## VERTEX AI ##

View file

@ -5715,7 +5715,7 @@ def _check_provider_match(model_info: dict, custom_llm_provider: Optional[str])
# as a last attempt if the model is not on Azure AI, Azure then fallback to OpenAI cost
# tracking the cost is better than attributing 0 cost to it.
return True
elif custom_llm_provider == "github":
elif custom_llm_provider in ("github", "github_copilot"):
# Allow github/<model> aliases to reuse existing provider metadata.
return True
else:

View file

@ -0,0 +1,273 @@
"""
Tests for web search model redirect in common_processing_pre_call_logic.
Covers per-deployment force_websearch_model and global websearch_fallback_model.
The redirect logic lives inline in common_processing_pre_call_logic, so tests
construct a ProxyBaseLLMRequestProcessing instance and call the method with
mocked dependencies, then verify self.data["model"] after the redirect block.
"""
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
WEB_SEARCH_TOOL = {"type": "web_search_20250305", "name": "web_search"}
REGULAR_TOOL = {
"type": "custom",
"name": "get_weather",
"description": "Get weather",
"input_schema": {"type": "object", "properties": {"location": {"type": "string"}}},
}
def _make_mock_router(deployments):
"""Build a mock router that returns the given deployments for get_model_list."""
router = MagicMock()
router.get_model_list.return_value = deployments
router.get_model_group_info.return_value = None
return router
def _make_mock_request():
request = MagicMock()
request.headers = {}
request.url = MagicMock()
request.url.path = "/v1/messages"
return request
def _make_user_api_key_dict():
user_api_key_dict = MagicMock()
user_api_key_dict.aliases = {}
user_api_key_dict.models = []
user_api_key_dict.api_key = "sk-test"
user_api_key_dict.user_id = "test"
user_api_key_dict.team_id = None
user_api_key_dict.metadata = {}
return user_api_key_dict
async def _run_pre_call_until_redirect(data, llm_router):
"""Run common_processing_pre_call_logic far enough to trigger the redirect,
then let it fail on subsequent steps — we only care about data['model']."""
proc = ProxyBaseLLMRequestProcessing(data=data)
async def _passthrough_add_litellm_data(data, **kwargs):
return data
with patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
side_effect=_passthrough_add_litellm_data,
):
try:
await proc.common_processing_pre_call_logic(
request=_make_mock_request(),
user_api_key_dict=_make_user_api_key_dict(),
llm_router=llm_router,
proxy_config=MagicMock(),
general_settings={},
proxy_logging_obj=MagicMock(),
route_type="anthropic_messages",
version=None,
)
except Exception:
pass
return proc.data.get("model")
class TestForceWebsearchModel:
@pytest.mark.asyncio
async def test_pure_websearch_redirected(self):
router = _make_mock_router(
[
{
"litellm_params": {
"model": "openai/gpt-5.4",
"force_websearch_model": "gpt-5.5",
}
}
]
)
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
result = await _run_pre_call_until_redirect(data, router)
assert result == "gpt-5.5"
@pytest.mark.asyncio
async def test_mixed_tools_not_redirected(self):
router = _make_mock_router(
[
{
"litellm_params": {
"model": "openai/gpt-5.4",
"force_websearch_model": "gpt-5.5",
}
}
]
)
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL, REGULAR_TOOL],
}
result = await _run_pre_call_until_redirect(data, router)
assert result == "my-model"
@pytest.mark.asyncio
async def test_multiple_websearch_tools_redirected(self):
router = _make_mock_router(
[
{
"litellm_params": {
"model": "openai/gpt-5.4",
"force_websearch_model": "gpt-5.5",
}
}
]
)
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "search"}],
"tools": [
WEB_SEARCH_TOOL,
{"name": "WebSearch", "description": "search"},
],
}
result = await _run_pre_call_until_redirect(data, router)
assert result == "gpt-5.5"
@pytest.mark.asyncio
async def test_no_force_no_redirect(self):
router = _make_mock_router([{"litellm_params": {"model": "openai/gpt-5.5"}}])
data = {
"model": "gpt-5.5",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
with patch("litellm.supports_web_search", return_value=True):
result = await _run_pre_call_until_redirect(data, router)
assert result == "gpt-5.5"
class TestWebsearchFallbackModel:
@pytest.mark.asyncio
async def test_fallback_when_model_lacks_search(self):
router = _make_mock_router(
[{"litellm_params": {"model": "openai/my-local-llm"}}]
)
data = {
"model": "my-local-llm",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
with (
patch("litellm.supports_web_search", return_value=False),
patch.object(
litellm,
"websearch_fallback_model",
"gpt-5.4-mini",
create=True,
),
):
result = await _run_pre_call_until_redirect(data, router)
assert result == "gpt-5.4-mini"
@pytest.mark.asyncio
async def test_no_fallback_when_model_supports_search(self):
router = _make_mock_router([{"litellm_params": {"model": "openai/gpt-5.5"}}])
data = {
"model": "gpt-5.5",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
with patch("litellm.supports_web_search", return_value=True):
result = await _run_pre_call_until_redirect(data, router)
assert result == "gpt-5.5"
@pytest.mark.asyncio
async def test_no_fallback_when_setting_not_configured(self):
router = _make_mock_router(
[{"litellm_params": {"model": "openai/my-local-llm"}}]
)
data = {
"model": "my-local-llm",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
with (
patch("litellm.supports_web_search", return_value=False),
patch.object(litellm, "websearch_fallback_model", None, create=True),
):
result = await _run_pre_call_until_redirect(data, router)
assert result == "my-local-llm"
class TestWebsearchForceOverFallbackPriority:
@pytest.mark.asyncio
async def test_force_takes_priority_over_fallback(self):
router = _make_mock_router(
[
{
"litellm_params": {
"model": "openai/gpt-5.4",
"force_websearch_model": "gpt-5.5",
}
}
]
)
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
with (
patch("litellm.supports_web_search", return_value=False),
patch.object(
litellm,
"websearch_fallback_model",
"gpt-5.4-mini",
create=True,
),
):
result = await _run_pre_call_until_redirect(data, router)
assert result == "gpt-5.5"
class TestWebsearchRedirectGuardConditions:
@pytest.mark.asyncio
async def test_no_tools(self):
router = _make_mock_router([])
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "hi"}],
}
result = await _run_pre_call_until_redirect(data, router)
assert result == "my-model"
@pytest.mark.asyncio
async def test_no_router(self):
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "search"}],
"tools": [WEB_SEARCH_TOOL],
}
result = await _run_pre_call_until_redirect(data, None)
assert result == "my-model"
@pytest.mark.asyncio
async def test_empty_tools(self):
router = _make_mock_router([])
data = {
"model": "my-model",
"messages": [{"role": "user", "content": "hi"}],
"tools": [],
}
result = await _run_pre_call_until_redirect(data, router)
assert result == "my-model"