Merge pull request #42007 from BerriAI/litellm_/web-search-autoship-scope-01a237

feat(ui): configure web search interception from the Admin UI
This commit is contained in:
yuneng-jiang 2026-09-19 15:54:03 -07:00 committed by GitHub
commit 5fc5afd595
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1312 additions and 2 deletions

View file

@ -123,6 +123,7 @@ class SupportedDBObjectType(str, enum.Enum):
MODEL_COST_MAP = "model_cost_map"
TOOLS = "tools"
CONFIG_OVERRIDES = "config_overrides"
WEBSEARCH_INTERCEPTION_SETTINGS = "websearch_interception_settings"
def __str__(self):
return str(self.value)

View file

@ -4787,6 +4787,22 @@ def _swap_in_model_cost_map(new_model_cost_map: dict) -> int:
return adopt_model_cost_map(new_model_cost_map)
def _websearch_handler_params(stored: Mapping[str, object]) -> dict[str, object]:
"""
Translate stored web search interception settings into handler kwargs.
Drops ``enabled``, which gates the callback rather than configuring it, and
drops an ``enabled_providers`` that is not a non-empty list so the handler
applies its own default. An empty list otherwise matches no provider at all,
and a bare string is iterated one character at a time.
"""
params: Final = {key: value for key, value in stored.items() if key != "enabled"}
providers: Final = params.get("enabled_providers")
if not isinstance(providers, list) or not providers:
params.pop("enabled_providers", None)
return params
def should_load_db_object(object_type: str | SupportedDBObjectType) -> bool:
"""
Check if an object type should be loaded from the database based on general_settings.supported_db_objects.
@ -4888,6 +4904,7 @@ class ProxyConfig:
def __init__(self) -> None:
self.config: Mapping[str, object] = MappingProxyType({})
self._last_semantic_filter_config: dict[str, object] | None = None
self._last_websearch_interception_config: dict[str, object] | None = None
self._last_hashicorp_vault_config: dict[str, object] | None = None
self._last_cyberark_config: dict[str, object] | None = None # mutable-ok: change-detection cache
self._cyberark_boot_env: dict[str, str | None] | None = None # mutable-ok: deployment env snapshot, set once
@ -7709,6 +7726,9 @@ class ProxyConfig:
if self._should_load_db_object(object_type="semantic_filter_settings"):
await self._init_semantic_filter_settings_in_db(prisma_client=prisma_client)
if self._should_load_db_object(object_type=SupportedDBObjectType.WEBSEARCH_INTERCEPTION_SETTINGS):
await self.init_websearch_interception_settings_in_db(prisma_client=prisma_client)
if self._should_load_db_object(object_type="config_overrides"):
await self._init_hashicorp_vault_config_override(prisma_client=prisma_client)
await self._init_cyberark_config_override(prisma_client=prisma_client)
@ -7787,6 +7807,66 @@ class ProxyConfig:
except Exception as e:
verbose_proxy_logger.exception("Error initializing semantic filter settings from DB: %s", e)
async def init_websearch_interception_settings_in_db(self, prisma_client: PrismaClient):
"""
Initialize web search interception settings from database.
Called periodically (approximately every 10 seconds) by background task to hot-reload settings across all pods.
"""
import json
import litellm
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
try:
config_record: Final = await get_config_param(prisma_client, "litellm_settings")
if config_record is None or config_record.param_value is None:
return
litellm_settings = config_record.param_value
if isinstance(litellm_settings, str):
litellm_settings = json.loads(litellm_settings)
websearch_config: Final = litellm_settings.get("websearch_interception_params", None)
if not isinstance(websearch_config, Mapping):
return
if "enabled" not in websearch_config and self._last_websearch_interception_config is None:
verbose_proxy_logger.debug(
"Web search interception: stored settings carry no 'enabled' flag and none were applied "
"before, so litellm_settings.callbacks keeps ownership of the callback."
)
return
enabled: Final = bool(coerce_bool(websearch_config.get("enabled", True)))
registered: Final = bool(
litellm.logging_callback_manager.get_custom_loggers_for_type(WebSearchInterceptionLogger)
)
if self._last_websearch_interception_config == websearch_config and registered == enabled:
return
replacement: Final = (
WebSearchInterceptionLogger.from_config_yaml(_websearch_handler_params(websearch_config))
if enabled
else None
)
litellm.logging_callback_manager.remove_callbacks_by_type(litellm.callbacks, WebSearchInterceptionLogger)
if replacement is not None:
litellm.logging_callback_manager.add_litellm_callback(replacement)
verbose_proxy_logger.info("Web search interception reinitialized from DB")
else:
verbose_proxy_logger.info("Web search interception disabled")
self._last_websearch_interception_config = dict(websearch_config)
except Exception as e:
verbose_proxy_logger.exception("Error initializing web search interception settings from DB: %s", e)
async def _init_sso_settings_in_db(self, prisma_client: PrismaClient):
"""
Initialize SSO settings from database into the router on startup.

View file

@ -6,6 +6,7 @@ from collections import Counter
from collections.abc import Mapping, MutableMapping, Sequence
from types import MappingProxyType
from typing import (
Annotated,
Final,
NamedTuple,
Protocol,
@ -477,6 +478,72 @@ class MCPToolSearchSettingsResponse(SettingsResponse):
"""Response model for native MCP tool search settings"""
class WebSearchInterceptionSettings(BaseModel):
"""Configuration for server-side web search interception"""
enabled: bool = Field(
default=False,
description="Serve web search tool calls from a configured search tool instead of passing them upstream",
)
enabled_providers: list[str] = Field(
default_factory=list,
description="LLM providers to intercept for (e.g. 'bedrock', 'vertex_ai'). Empty intercepts Bedrock only.",
)
search_tool_name: str | None = Field(
default=None,
description="Name of the configured search tool to run searches through. Empty uses the first one available.",
)
max_agentic_loops: int | None = Field(
default=None,
ge=1,
description="How many follow-up model calls one intercepted request may chain. Empty applies the default of 3.",
)
class WebSearchInterceptionSettingsResponse(SettingsResponse):
"""Response model for web search interception settings"""
def _with_websearch_enabled_resolved(config: Mapping[str, object]) -> dict[str, object]:
"""
Answer with the stored flag when there is one, and only otherwise with what
this process is running.
A stored flag is the cluster's own answer, so it is the same on every pod and
is safe for the page to send back on save. Deriving the answer from this
process instead would report off on a pod that has not polled yet, and the
next save would persist that as a cluster-wide off. Without a stored flag the
only available answer is local: litellm_settings.callbacks activates
interception without storing one, and a write through the generic config
endpoint can drop the flag from a block that is still live. Reporting the
field default there would claim the feature is off while it serves.
"""
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
litellm_settings: Final[Mapping[str, object]] = _as_settings_section(config.get("litellm_settings"))
stored: Final[Mapping[str, object]] = _as_settings_section(litellm_settings.get("websearch_interception_params"))
if "enabled" in stored:
return dict(config)
resolved: Final = {
**stored,
"enabled": bool(litellm.logging_callback_manager.get_custom_loggers_for_type(WebSearchInterceptionLogger)),
}
return {
**config,
"litellm_settings": {**litellm_settings, "websearch_interception_params": resolved},
}
def _as_settings_section(value: object) -> Mapping[str, object]:
return cast("Mapping[str, object]", value) if isinstance(value, Mapping) else MappingProxyType({})
@router.get(
"/get/allowed_ips",
tags=["Budget & Spend Tracking"],
@ -875,7 +942,13 @@ async def update_default_team_member_budget(teams: list[NewUserRequestTeam], use
async def _update_litellm_setting(
settings: DefaultInternalUserParams | DefaultTeamSSOParams | MCPSemanticFilterSettings | MCPToolSearchSettings,
settings: (
DefaultInternalUserParams
| DefaultTeamSSOParams
| MCPSemanticFilterSettings
| MCPToolSearchSettings
| WebSearchInterceptionSettings
),
settings_key: str,
success_message: str,
user_api_key_dict: UserAPIKeyAuth,
@ -1399,6 +1472,77 @@ async def update_mcp_semantic_filter_settings(
return result
@router.get(
"/get/websearch_interception_settings",
tags=["Settings"],
dependencies=[Depends(user_api_key_auth)],
response_model=WebSearchInterceptionSettingsResponse,
)
async def get_websearch_interception_settings(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
"""
Get web search interception configuration.
Returns the current settings plus their schema, for the Admin UI to render.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_config
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": "Database not connected. Please connect a database."},
)
config: Final = await proxy_config.get_config()
return await _get_settings_with_schema(
settings_key="websearch_interception_params",
settings_class=WebSearchInterceptionSettings,
config=_with_websearch_enabled_resolved(config),
)
@router.patch(
"/update/websearch_interception_settings",
tags=["Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_websearch_interception_settings(
settings: WebSearchInterceptionSettings,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
"""
Update web search interception settings in database.
Settings will be picked up by all pods within approximately 10 seconds via background polling.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail="Only proxy admins can update web search interception settings.",
)
result: Final = await _update_litellm_setting(
settings=settings,
settings_key="websearch_interception_params",
success_message=(
"Web search interception settings updated successfully. "
"Changes will be applied across all pods within 10 seconds."
),
user_api_key_dict=user_api_key_dict,
)
try:
from litellm.proxy.proxy_server import prisma_client, proxy_config
if prisma_client is not None:
await proxy_config.init_websearch_interception_settings_in_db(prisma_client=prisma_client)
except Exception as e:
verbose_proxy_logger.warning("Failed to reinitialize web search interception settings immediately: %s", e)
return result
@router.get(
"/get/mcp_tool_search_settings",
tags=["Settings"], # mutable-ok: FastAPI's route decorator only accepts a list

View file

@ -4526,3 +4526,191 @@ async def test_add_deployment_syncs_ui_settings_even_when_the_model_reconcile_fa
await config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=MagicMock())
assert general_settings["allow_agents_for_team_admins"] is True
def _websearch_logger_cls():
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
return WebSearchInterceptionLogger
def _run_websearch_init(monkeypatch, stored_params, starting_callbacks):
pc = ProxyConfig()
monkeypatch.setattr(litellm, "callbacks", list(starting_callbacks))
monkeypatch.setattr(
"litellm.proxy.proxy_server.get_config_param",
AsyncMock(return_value=SimpleNamespace(param_value={"websearch_interception_params": stored_params}))
if stored_params is not None
else AsyncMock(return_value=SimpleNamespace(param_value={})),
)
asyncio.run(pc.init_websearch_interception_settings_in_db(prisma_client=MagicMock()))
return pc
def _poll_websearch_init(pc, monkeypatch, stored_params):
monkeypatch.setattr(
"litellm.proxy.proxy_server.get_config_param",
AsyncMock(return_value=SimpleNamespace(param_value={"websearch_interception_params": stored_params})),
)
asyncio.run(pc.init_websearch_interception_settings_in_db(prisma_client=MagicMock()))
def test_init_websearch_interception_resyncs_after_a_write_drops_the_enabled_flag(monkeypatch):
logger_cls = _websearch_logger_cls()
pc = ProxyConfig()
monkeypatch.setattr(litellm, "callbacks", [])
_poll_websearch_init(pc, monkeypatch, {"enabled": True, "search_tool_name": "old-tool"})
_poll_websearch_init(pc, monkeypatch, {"search_tool_name": "new-tool"})
registered = [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)]
assert len(registered) == 1
assert registered[0].search_tool_name == "new-tool"
def test_init_websearch_interception_ignores_a_non_list_providers_value(monkeypatch):
logger_cls = _websearch_logger_cls()
_run_websearch_init(
monkeypatch,
stored_params={"enabled": True, "enabled_providers": "bedrock", "search_tool_name": "stored-tool"},
starting_callbacks=[],
)
registered = [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)]
assert len(registered) == 1
assert registered[0].enabled_providers == ["bedrock"]
def test_init_websearch_interception_absent_key_leaves_callbacks_untouched(monkeypatch):
logger_cls = _websearch_logger_cls()
config_registered = logger_cls(search_tool_name="from-config-yaml")
_run_websearch_init(monkeypatch, stored_params=None, starting_callbacks=[config_registered])
assert litellm.callbacks == [config_registered]
def test_init_websearch_interception_without_enabled_key_leaves_callbacks_untouched(monkeypatch):
logger_cls = _websearch_logger_cls()
config_registered = logger_cls(search_tool_name="from-config-yaml")
_run_websearch_init(
monkeypatch,
stored_params={"search_tool_name": "stored-tool"},
starting_callbacks=[config_registered],
)
assert litellm.callbacks == [config_registered]
def test_init_websearch_interception_registers_when_explicitly_enabled(monkeypatch):
logger_cls = _websearch_logger_cls()
_run_websearch_init(
monkeypatch,
stored_params={"enabled": True, "search_tool_name": "stored-tool"},
starting_callbacks=[],
)
registered = [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)]
assert len(registered) == 1
assert registered[0].search_tool_name == "stored-tool"
def test_init_websearch_interception_treats_string_false_as_disabled(monkeypatch):
logger_cls = _websearch_logger_cls()
existing = logger_cls(search_tool_name="stored-tool")
_run_websearch_init(
monkeypatch,
stored_params={"enabled": "false", "search_tool_name": "stored-tool"},
starting_callbacks=[existing],
)
assert [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)] == []
def test_init_websearch_interception_empty_providers_falls_back_to_handler_default(monkeypatch):
logger_cls = _websearch_logger_cls()
_run_websearch_init(
monkeypatch,
stored_params={"enabled": True, "enabled_providers": [], "search_tool_name": "stored-tool"},
starting_callbacks=[],
)
registered = [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)]
assert len(registered) == 1
assert registered[0].enabled_providers == ["bedrock"]
def test_init_websearch_interception_keeps_working_callback_when_new_one_cannot_be_built(monkeypatch):
logger_cls = _websearch_logger_cls()
working = logger_cls(search_tool_name="stored-tool", max_agentic_loops=3)
_run_websearch_init(
monkeypatch,
stored_params={"enabled": True, "search_tool_name": "stored-tool", "max_agentic_loops": 0},
starting_callbacks=[working],
)
assert litellm.callbacks == [working]
def test_init_websearch_interception_disabled_removes_the_callback(monkeypatch):
logger_cls = _websearch_logger_cls()
existing = logger_cls(search_tool_name="stored-tool")
_run_websearch_init(
monkeypatch,
stored_params={"enabled": False, "search_tool_name": "stored-tool"},
starting_callbacks=[existing],
)
assert [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)] == []
def test_init_websearch_interception_replaces_stale_instance_on_param_change(monkeypatch):
logger_cls = _websearch_logger_cls()
stale = logger_cls(search_tool_name="old-tool", max_agentic_loops=2)
_run_websearch_init(
monkeypatch,
stored_params={"enabled": True, "search_tool_name": "new-tool", "max_agentic_loops": 7},
starting_callbacks=[stale],
)
registered = [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)]
assert len(registered) == 1
assert (registered[0].search_tool_name, registered[0].max_agentic_loops) == ("new-tool", 7)
def test_init_websearch_interception_honors_enabled_providers(monkeypatch):
logger_cls = _websearch_logger_cls()
_run_websearch_init(
monkeypatch,
stored_params={"enabled": True, "enabled_providers": ["bedrock", "vertex_ai"]},
starting_callbacks=[],
)
registered = [cb for cb in litellm.callbacks if isinstance(cb, logger_cls)]
assert len(registered) == 1
assert registered[0].enabled_providers == ["bedrock", "vertex_ai"]
def test_websearch_interception_settings_can_be_named_in_supported_db_objects(monkeypatch):
from litellm.proxy import proxy_server
from litellm.proxy._types import ConfigGeneralSettings
allowlist = ConfigGeneralSettings(supported_db_objects=["websearch_interception_settings"]).supported_db_objects
assert allowlist
monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": allowlist})
assert proxy_server.should_load_db_object(object_type="websearch_interception_settings") is True
monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["models"]})
assert proxy_server.should_load_db_object(object_type="websearch_interception_settings") is False

View file

@ -3188,6 +3188,182 @@ class TestMcpToolSearchSettingsEndpoints:
assert mock_proxy_config["save_call_count"]() == 0
class TestWebSearchInterceptionSettingsEndpoints:
@staticmethod
def _override_auth(role: LitellmUserRoles):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="u", api_key="hashed", user_role=role
)
def test_get_returns_stored_values_and_field_schema(self, mock_proxy_config, mock_auth, monkeypatch):
import litellm
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(search_tool_name="running")])
mock_proxy_config["config"]["litellm_settings"]["websearch_interception_params"] = {
"enabled": True,
"enabled_providers": ["bedrock", "vertex_ai"],
"search_tool_name": "my-perplexity-search",
}
resp = client.get("/get/websearch_interception_settings")
assert resp.status_code == 200, resp.text
assert resp.json()["values"] == {
"enabled": True,
"enabled_providers": ["bedrock", "vertex_ai"],
"search_tool_name": "my-perplexity-search",
"max_agentic_loops": None,
}
assert resp.json()["field_schema"]["properties"]["enabled_providers"]["type"] == "array"
def test_update_requires_proxy_admin(self, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
self._override_auth(LitellmUserRoles.INTERNAL_USER)
try:
resp = client.patch("/update/websearch_interception_settings", json={"enabled": True})
finally:
app.dependency_overrides.clear()
assert resp.status_code == 403
assert "proxy admin" in resp.json()["detail"].lower()
def test_update_persists_settings(self, mock_proxy_config, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
self._override_auth(LitellmUserRoles.PROXY_ADMIN)
payload = {
"enabled": True,
"enabled_providers": ["bedrock"],
"search_tool_name": "my-perplexity-search",
"max_agentic_loops": 5,
}
try:
resp = client.patch("/update/websearch_interception_settings", json=payload)
finally:
app.dependency_overrides.clear()
assert resp.status_code == 200, resp.text
assert mock_proxy_config["save_call_count"]() == 1
assert mock_proxy_config["config"]["litellm_settings"]["websearch_interception_params"] == payload
def test_get_reports_enabled_while_the_callback_is_running_without_a_stored_flag(
self, mock_proxy_config, mock_auth, monkeypatch
):
import litellm
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr(litellm, "callbacks", [WebSearchInterceptionLogger(search_tool_name="from-config")])
mock_proxy_config["config"]["litellm_settings"]["websearch_interception_params"] = {
"enabled_providers": ["bedrock"],
"search_tool_name": "my-perplexity-search",
}
resp = client.get("/get/websearch_interception_settings")
assert resp.status_code == 200, resp.text
assert resp.json()["values"]["enabled"] is True
def test_get_reports_disabled_when_nothing_is_stored_and_nothing_is_running(
self, mock_proxy_config, mock_auth, monkeypatch
):
import litellm
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr(litellm, "callbacks", [])
mock_proxy_config["config"]["litellm_settings"]["websearch_interception_params"] = {
"enabled_providers": ["bedrock"],
}
resp = client.get("/get/websearch_interception_settings")
assert resp.status_code == 200, resp.text
assert resp.json()["values"]["enabled"] is False
assert resp.json()["values"]["enabled_providers"] == ["bedrock"]
def test_update_reapplies_settings_to_the_running_proxy(self, mock_proxy_config, monkeypatch):
from unittest.mock import AsyncMock
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
reapply = AsyncMock()
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config.init_websearch_interception_settings_in_db",
reapply,
)
self._override_auth(LitellmUserRoles.PROXY_ADMIN)
try:
resp = client.patch("/update/websearch_interception_settings", json={"enabled": True})
finally:
app.dependency_overrides.clear()
assert resp.status_code == 200, resp.text
reapply.assert_awaited_once()
def test_get_keeps_the_stored_flag_when_this_pod_has_not_reinitialized(
self, mock_proxy_config, mock_auth, monkeypatch
):
import litellm
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr(litellm, "callbacks", [])
mock_proxy_config["config"]["litellm_settings"]["websearch_interception_params"] = {
"enabled": True,
"search_tool_name": "cluster-search",
}
resp = client.get("/get/websearch_interception_settings")
assert resp.status_code == 200, resp.text
assert resp.json()["values"]["enabled"] is True
def test_get_reports_no_database_instead_of_empty_settings(self, mock_proxy_config, mock_auth, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
resp = client.get("/get/websearch_interception_settings")
assert resp.status_code == 500, resp.text
assert "Database not connected" in resp.json()["detail"]["error"]
def test_update_still_saves_when_the_live_reinit_fails(self, mock_proxy_config, monkeypatch):
from unittest.mock import AsyncMock
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config.init_websearch_interception_settings_in_db",
AsyncMock(side_effect=RuntimeError("callback blew up")),
)
self._override_auth(LitellmUserRoles.PROXY_ADMIN)
try:
resp = client.patch("/update/websearch_interception_settings", json={"enabled": True})
finally:
app.dependency_overrides.clear()
assert resp.status_code == 200, resp.text
assert mock_proxy_config["save_call_count"]() == 1
def test_update_rejects_zero_max_agentic_loops(self, mock_proxy_config, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
self._override_auth(LitellmUserRoles.PROXY_ADMIN)
try:
resp = client.patch(
"/update/websearch_interception_settings",
json={"enabled": True, "max_agentic_loops": 0},
)
finally:
app.dependency_overrides.clear()
assert resp.status_code == 422
assert mock_proxy_config["save_call_count"]() == 0
def test_upload_logo_requires_proxy_admin(monkeypatch):
"""Any authenticated key could previously write a file to the server's disk here."""
from litellm.proxy._types import UserAPIKeyAuth

View file

@ -22,6 +22,7 @@ import UserBannerSettings from "@/components/Settings/AdminSettings/UserBannerSe
import CyberArk from "@/components/Settings/AdminSettings/CyberArk/CyberArk";
import HashicorpVault from "@/components/Settings/AdminSettings/HashicorpVault/HashicorpVault";
import PluginSettings from "@/components/Settings/AdminSettings/PluginSettings/PluginSettings";
import WebSearchInterceptionSettings from "@/components/Settings/AdminSettings/WebSearchInterceptionSettings/WebSearchInterceptionSettings";
import SSOModals from "@/components/SSOModals";
import {
emptySSOSettingsFormValues,
@ -408,6 +409,11 @@ const AdminPanel: React.FC<AdminPanelProps> = ({ proxySettings }) => {
label: "Plugins",
children: <PluginSettings />,
},
{
key: "web-search-interception",
label: "Web Search Interception",
children: <WebSearchInterceptionSettings />,
},
];
return (

View file

@ -0,0 +1,23 @@
import { updateWebSearchInterceptionSettings, type WebSearchInterceptionSettings } from "@/components/networking";
import { useMutation, useQueryClient } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
const webSearchInterceptionSettingsKeys = createQueryKeys("webSearchInterceptionSettings");
export const useUpdateWebSearchInterceptionSettings = (accessToken: string) => {
const queryClient = useQueryClient();
return useMutation({
mutationFn: async (settings: WebSearchInterceptionSettings) => {
if (!accessToken) {
throw new Error("Access token is required");
}
return updateWebSearchInterceptionSettings(accessToken, settings);
},
onSuccess: () => {
queryClient.invalidateQueries({
queryKey: webSearchInterceptionSettingsKeys.all,
});
},
});
};

View file

@ -0,0 +1,17 @@
import { getWebSearchInterceptionSettings, type WebSearchInterceptionSettingsResponse } from "@/components/networking";
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import useAuthorized from "../useAuthorized";
const webSearchInterceptionSettingsKeys = createQueryKeys("webSearchInterceptionSettings");
export const useWebSearchInterceptionSettings = () => {
const { accessToken } = useAuthorized();
return useQuery<WebSearchInterceptionSettingsResponse>({
queryKey: webSearchInterceptionSettingsKeys.list({}),
queryFn: async () => await getWebSearchInterceptionSettings(accessToken),
enabled: !!accessToken,
staleTime: 60 * 60 * 1000,
gcTime: 60 * 60 * 1000,
});
};

View file

@ -0,0 +1,192 @@
import React from "react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { render, screen, act } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import WebSearchInterceptionSettings from "./WebSearchInterceptionSettings";
import { useWebSearchInterceptionSettings } from "@/app/(dashboard)/hooks/webSearchInterceptionSettings/useWebSearchInterceptionSettings";
import { useUpdateWebSearchInterceptionSettings } from "@/app/(dashboard)/hooks/webSearchInterceptionSettings/useUpdateWebSearchInterceptionSettings";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
vi.mock("@/app/(dashboard)/hooks/webSearchInterceptionSettings/useWebSearchInterceptionSettings", () => ({
useWebSearchInterceptionSettings: vi.fn(),
}));
vi.mock("@/app/(dashboard)/hooks/webSearchInterceptionSettings/useUpdateWebSearchInterceptionSettings", () => ({
useUpdateWebSearchInterceptionSettings: vi.fn(),
}));
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: vi.fn(),
}));
vi.mock("@/components/networking", () => ({
fetchSearchTools: vi.fn().mockResolvedValue({
search_tools: [{ search_tool_name: "my-perplexity-search" }, { search_tool_name: "backup-search" }],
}),
}));
const mockMutate = vi.fn();
const ENABLED_PAYLOAD = {
enabled: true,
enabled_providers: ["bedrock"],
search_tool_name: "my-perplexity-search",
max_agentic_loops: null,
};
const storedSettings = {
field_schema: {
properties: {
enabled: { description: "Serve web search tool calls from a configured search tool" },
},
},
values: {
enabled: false,
enabled_providers: ["bedrock"],
search_tool_name: "my-perplexity-search",
max_agentic_loops: null,
},
};
async function renderSettings() {
const result = render(<WebSearchInterceptionSettings />);
await act(async () => {});
return result;
}
describe("WebSearchInterceptionSettings", () => {
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(useAuthorized).mockReturnValue({ accessToken: "test-token" } as any);
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: storedSettings,
isLoading: false,
isError: false,
error: null,
} as any);
vi.mocked(useUpdateWebSearchInterceptionSettings).mockReturnValue({
mutate: mockMutate,
isPending: false,
error: null,
} as any);
});
it("renders the settings section", async () => {
await renderSettings();
expect(screen.getByText("Web Search Interception")).toBeInTheDocument();
});
it("shows a login prompt when there is no access token", () => {
vi.mocked(useAuthorized).mockReturnValue({ accessToken: null } as any);
render(<WebSearchInterceptionSettings />);
expect(screen.getByText(/please log in/i)).toBeInTheDocument();
});
it("hides the settings while loading", async () => {
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: undefined,
isLoading: true,
isError: false,
error: null,
} as any);
await renderSettings();
expect(screen.queryByText("Enable Web Search Interception")).not.toBeInTheDocument();
});
it("surfaces a load failure", async () => {
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: undefined,
isLoading: false,
isError: true,
error: new Error("boom"),
} as any);
await renderSettings();
expect(screen.getByText("Could not load web search interception settings")).toBeInTheDocument();
expect(screen.getByText("boom")).toBeInTheDocument();
});
it("keeps save disabled until something changes", async () => {
const user = userEvent.setup();
await renderSettings();
const save = screen.getByRole("button", { name: /save settings/i });
expect(save).toBeDisabled();
await user.click(save);
expect(mockMutate).not.toHaveBeenCalled();
});
it("submits the stored values with the toggled enabled flag", async () => {
const user = userEvent.setup();
await renderSettings();
await user.click(screen.getByRole("switch"));
await user.click(screen.getByRole("button", { name: /save settings/i }));
expect(mockMutate).toHaveBeenCalledTimes(1);
expect(mockMutate.mock.calls[0][0]).toEqual(ENABLED_PAYLOAD);
});
it("ignores stored values whose types do not match the field", async () => {
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: {
...storedSettings,
values: {
enabled: "yes",
enabled_providers: "bedrock",
search_tool_name: 7,
max_agentic_loops: "3",
},
},
isLoading: false,
isError: false,
error: null,
} as any);
await renderSettings();
expect(screen.getByRole("switch")).not.toBeChecked();
expect(screen.getByLabelText(/max agentic loops/i)).toHaveValue(null);
expect(screen.queryByText("bedrock")).not.toBeInTheDocument();
});
it("reseeds the form when the stored settings change underneath it", async () => {
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: { ...storedSettings, values: { ...storedSettings.values, max_agentic_loops: 3 } },
isLoading: false,
isError: false,
error: null,
} as any);
const { rerender } = await renderSettings();
expect(screen.getByRole("spinbutton")).toHaveValue(3);
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: { ...storedSettings, values: { ...storedSettings.values, max_agentic_loops: 9 } },
isLoading: false,
isError: false,
error: null,
} as any);
await act(async () => {
rerender(<WebSearchInterceptionSettings />);
});
expect(screen.getByRole("spinbutton")).toHaveValue(9);
});
it("sends null rather than a number when the loop cap is cleared", async () => {
const user = userEvent.setup();
vi.mocked(useWebSearchInterceptionSettings).mockReturnValue({
data: { ...storedSettings, values: { ...storedSettings.values, max_agentic_loops: 5 } },
isLoading: false,
isError: false,
error: null,
} as any);
await renderSettings();
await user.clear(screen.getByRole("spinbutton"));
await user.click(screen.getByRole("button", { name: /save settings/i }));
expect(mockMutate).toHaveBeenCalledTimes(1);
expect(mockMutate.mock.calls[0][0].max_agentic_loops).toBeNull();
});
});

View file

@ -0,0 +1,317 @@
"use client";
import { useWebSearchInterceptionSettings } from "@/app/(dashboard)/hooks/webSearchInterceptionSettings/useWebSearchInterceptionSettings";
import { useUpdateWebSearchInterceptionSettings } from "@/app/(dashboard)/hooks/webSearchInterceptionSettings/useUpdateWebSearchInterceptionSettings";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { toast } from "@/lib/toast";
import { Skeleton } from "@/components/ui/skeleton";
import { CircleHelp, Info, Save } from "lucide-react";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { useEffect, useState } from "react";
import { useForm } from "react-hook-form";
import { FieldGroup } from "@/components/ui/field";
import { FormField } from "@/components/shared/form/FormField";
import { MultiSelect } from "@/components/shared/MultiSelect";
import { SearchSelect } from "@/components/shared/SearchSelect";
import { Button } from "@/components/ui/button";
import { Card, CardContent } from "@/components/ui/card";
import { Input } from "@/components/ui/input";
import { Switch } from "@/components/ui/switch";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { Providers, provider_map } from "@/components/provider_info_helpers";
import { fetchSearchTools } from "@/components/networking";
interface WebSearchInterceptionStoredValues {
enabled?: boolean;
enabled_providers?: string[];
search_tool_name?: string | null;
max_agentic_loops?: number | null;
}
interface WebSearchInterceptionFieldSchema {
properties?: {
enabled?: { description?: string };
enabled_providers?: { description?: string };
search_tool_name?: { description?: string };
max_agentic_loops?: { description?: string };
};
}
interface WebSearchInterceptionFormValues {
enabled: boolean;
enabled_providers: string[];
search_tool_name: string | null;
max_agentic_loops: number | null;
}
const NO_STORED_VALUES: Readonly<Record<string, unknown>> = {};
const MAX_AGENTIC_LOOPS_MIN = 1;
const PROVIDER_OPTIONS = Object.entries(provider_map)
.map(([enumKey, providerValue]) => ({
label: Providers[enumKey as keyof typeof Providers] ?? providerValue,
value: providerValue,
}))
.sort((a, b) => a.label.localeCompare(b.label));
const labelWithHint = (label: string, hint: string): React.ReactNode => (
<>
{label}
<Tooltip>
<TooltipTrigger render={<CircleHelp className="size-3.5 shrink-0 cursor-help text-muted-foreground" />} />
<TooltipContent>{hint}</TooltipContent>
</Tooltip>
</>
);
const parseLoops = (raw: string, rawAsNumber: number): number | null =>
raw === "" || Number.isNaN(rawAsNumber) ? null : rawAsNumber;
const isStringArray = (value: unknown): value is string[] =>
Array.isArray(value) && value.every((entry) => typeof entry === "string");
const toStoredValues = (raw: Readonly<Record<string, unknown>>): WebSearchInterceptionStoredValues => ({
enabled: typeof raw.enabled === "boolean" ? raw.enabled : undefined,
enabled_providers: isStringArray(raw.enabled_providers) ? raw.enabled_providers : undefined,
search_tool_name: typeof raw.search_tool_name === "string" ? raw.search_tool_name : null,
max_agentic_loops: typeof raw.max_agentic_loops === "number" ? raw.max_agentic_loops : null,
});
const toFormValues = (values: WebSearchInterceptionStoredValues): WebSearchInterceptionFormValues => ({
enabled: values.enabled ?? false,
enabled_providers: values.enabled_providers ?? [],
search_tool_name: values.search_tool_name ?? null,
max_agentic_loops: values.max_agentic_loops ?? null,
});
const readSearchToolNames = (response: unknown): string[] => {
const payload = response as { search_tools?: unknown; data?: unknown } | null;
const tools = Array.isArray(payload?.search_tools) ? payload.search_tools : payload?.data;
if (!Array.isArray(tools)) {
return [];
}
return tools
.map((tool: { search_tool_name?: string }) => tool?.search_tool_name)
.filter((name: unknown): name is string => typeof name === "string" && name.length > 0);
};
const useSearchToolNames = (accessToken: string) => {
const [searchTools, setSearchTools] = useState<string[]>([]);
const [loadingSearchTools, setLoadingSearchTools] = useState(true);
useEffect(() => {
const loadSearchTools = async () => {
if (!accessToken) return;
try {
setSearchTools(readSearchToolNames(await fetchSearchTools(accessToken)));
} catch (loadError) {
console.error("Error fetching search tools:", loadError);
} finally {
setLoadingSearchTools(false);
}
};
loadSearchTools();
}, [accessToken]);
return { searchTools, loadingSearchTools };
};
interface WebSearchInterceptionFormProps {
accessToken: string;
initial: WebSearchInterceptionFormValues;
schema: WebSearchInterceptionFieldSchema | undefined;
}
function WebSearchInterceptionForm({ accessToken, initial, schema }: WebSearchInterceptionFormProps) {
const {
mutate: updateSettings,
isPending: isUpdating,
error: updateError,
} = useUpdateWebSearchInterceptionSettings(accessToken);
const { searchTools, loadingSearchTools } = useSearchToolNames(accessToken);
const form = useForm<WebSearchInterceptionFormValues>({ defaultValues: initial });
const isDirty = form.formState.isDirty;
const handleSave = (formValues: WebSearchInterceptionFormValues) => {
updateSettings(formValues, {
onSuccess: () => {
form.reset(formValues);
toast.success("Settings updated successfully. Changes will be applied across all pods within 10 seconds.");
},
onError: (saveError) => {
toast.fromError(saveError);
},
});
};
return (
<>
{updateError && (
<Alert variant="error" className="mb-4">
<AlertTitle>Could not update settings</AlertTitle>
{updateError instanceof Error && <AlertDescription>{updateError.message}</AlertDescription>}
</Alert>
)}
<TooltipProvider>
<form onSubmit={(event) => event.preventDefault()} noValidate>
<Card className="mb-4">
<CardContent>
<FieldGroup>
<FormField
control={form.control}
name="enabled"
label={labelWithHint(
"Enable Web Search Interception",
"When enabled, web search tool calls are executed server-side through the selected search tool",
)}
description={schema?.properties?.enabled?.description}
>
{({ value, onChange, onBlur, id }) => (
<Switch id={id} checked={value} onCheckedChange={onChange} onBlur={onBlur} disabled={isUpdating} />
)}
</FormField>
<FormField
control={form.control}
name="enabled_providers"
label={labelWithHint(
"Providers",
"Which LLM providers to intercept for. Leave empty to intercept Bedrock only.",
)}
description={schema?.properties?.enabled_providers?.description}
>
{({ value, onChange, id }) => (
<MultiSelect
id={id}
options={PROVIDER_OPTIONS}
value={value}
onValueChange={onChange}
placeholder="Select providers (defaults to Bedrock)"
allowCustomValues
disabled={isUpdating}
/>
)}
</FormField>
<FormField
control={form.control}
name="search_tool_name"
label={labelWithHint(
"Search Tool",
"Which configured search tool runs the searches. Leave empty to use the first one available.",
)}
description={schema?.properties?.search_tool_name?.description}
>
{({ value, onChange, id }) => (
<SearchSelect
inputId={id}
options={searchTools.map((name) => ({ label: name, value: name }))}
value={value}
onValueChange={onChange}
placeholder="Select a search tool (defaults to the first available)"
disabled={isUpdating || loadingSearchTools}
/>
)}
</FormField>
<FormField
control={form.control}
name="max_agentic_loops"
label={labelWithHint(
"Max Agentic Loops",
"How many follow-up model calls one intercepted request may chain. Leave empty for the default of 3.",
)}
description={schema?.properties?.max_agentic_loops?.description}
>
{({ value, onChange, onBlur, id, ref }) => (
<Input
id={id}
ref={ref}
type="number"
min={MAX_AGENTIC_LOOPS_MIN}
value={value ?? ""}
onChange={(event) => onChange(parseLoops(event.target.value, event.target.valueAsNumber))}
onBlur={onBlur}
disabled={isUpdating}
/>
)}
</FormField>
</FieldGroup>
</CardContent>
</Card>
<div className="flex justify-end gap-2">
<Button
type="button"
onClick={() => void form.handleSubmit(handleSave)()}
disabled={!isDirty || isUpdating}
>
{isUpdating ? <UiLoadingSpinner className="size-4" /> : <Save />}
Save Settings
</Button>
</div>
</form>
</TooltipProvider>
</>
);
}
export default function WebSearchInterceptionSettings() {
const { accessToken } = useAuthorized();
const { data, isLoading, isError, error } = useWebSearchInterceptionSettings();
if (!accessToken) {
return (
<div className="p-6 text-center text-muted-foreground">
Please log in to configure web search interception settings.
</div>
);
}
if (isLoading) {
return (
<div className="flex flex-col gap-3">
<Skeleton className="h-4 w-2/5" />
<Skeleton className="h-4 w-full" />
<Skeleton className="h-4 w-full" />
<Skeleton className="h-4 w-3/5" />
</div>
);
}
if (isError) {
return (
<Alert variant="error" className="mb-6">
<AlertTitle>Could not load web search interception settings</AlertTitle>
{error instanceof Error && <AlertDescription>{error.message}</AlertDescription>}
</Alert>
);
}
const values: WebSearchInterceptionStoredValues = toStoredValues(data?.values ?? NO_STORED_VALUES);
return (
<div className="w-full">
<Alert variant="info" className="mb-6">
<Info />
<AlertTitle>Web Search Interception</AlertTitle>
<AlertDescription>
Serve web search tool calls from a configured search tool instead of passing them upstream, so models without
native web search can still answer with fresh results. Click &apos;Save Settings&apos; to apply changes across
all pods (takes effect within 10 seconds).
</AlertDescription>
</Alert>
<WebSearchInterceptionForm
key={JSON.stringify(values)}
accessToken={accessToken}
initial={toFormValues(values)}
schema={data?.field_schema}
/>
</div>
);
}

View file

@ -3667,6 +3667,34 @@ export const updateMCPSemanticFilterSettings = async (accessToken: string, setti
}
};
export type WebSearchInterceptionSettings = components["schemas"]["WebSearchInterceptionSettings"];
export type WebSearchInterceptionSettingsResponse = components["schemas"]["WebSearchInterceptionSettingsResponse"];
export const getWebSearchInterceptionSettings = async (
accessToken: string,
): Promise<WebSearchInterceptionSettingsResponse> => {
try {
return await apiClient.get<WebSearchInterceptionSettingsResponse>(`/get/websearch_interception_settings`, {
accessToken,
});
} catch (error) {
console.error("Failed to get web search interception settings:", error);
throw error;
}
};
export const updateWebSearchInterceptionSettings = async (
accessToken: string,
settings: WebSearchInterceptionSettings,
) => {
try {
return await apiClient.patch(`/update/websearch_interception_settings`, { accessToken, body: settings });
} catch (error) {
console.error("Failed to update web search interception settings:", error);
throw error;
}
};
export const testMCPSemanticFilter = async (accessToken: string, model: string, query: string) => {
/**
* Test MCP semantic filter by making a responses API call

View file

@ -5430,6 +5430,28 @@ export interface paths {
patch?: never;
trace?: never;
};
"/get/websearch_interception_settings": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* Get Websearch Interception Settings
* @description Get web search interception configuration.
*
* Returns the current settings plus their schema, for the Admin UI to render.
*/
get: operations["get_websearch_interception_settings_get_websearch_interception_settings_get"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/get_favicon": {
parameters: {
query?: never;
@ -16820,6 +16842,28 @@ export interface paths {
patch: operations["update_user_banner_update_user_banner_patch"];
trace?: never;
};
"/update/websearch_interception_settings": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
/**
* Update Websearch Interception Settings
* @description Update web search interception settings in database.
*
* Settings will be picked up by all pods within approximately 10 seconds via background polling.
*/
patch: operations["update_websearch_interception_settings_update_websearch_interception_settings_patch"];
trace?: never;
};
"/upload/logo": {
parameters: {
query?: never;
@ -38024,7 +38068,7 @@ export interface components {
* Use in general_settings.supported_db_objects to specify which objects to load from DB.
* @enum {string}
*/
SupportedDBObjectType: "models" | "mcp" | "guardrails" | "policies" | "vector_stores" | "pass_through_endpoints" | "prompts" | "model_cost_map" | "tools" | "config_overrides";
SupportedDBObjectType: "models" | "mcp" | "guardrails" | "policies" | "vector_stores" | "pass_through_endpoints" | "prompts" | "model_cost_map" | "tools" | "config_overrides" | "websearch_interception_settings";
/** SupportedEndpoint */
SupportedEndpoint: {
/** Endpoint */
@ -41156,6 +41200,47 @@ export interface components {
/** Vector Store Name */
vector_store_name?: string | null;
};
/**
* WebSearchInterceptionSettings
* @description Configuration for server-side web search interception
*/
WebSearchInterceptionSettings: {
/**
* Enabled
* @description Serve web search tool calls from a configured search tool instead of passing them upstream
* @default false
*/
enabled: boolean;
/**
* Enabled Providers
* @description LLM providers to intercept for (e.g. 'bedrock', 'vertex_ai'). Empty intercepts Bedrock only.
*/
enabled_providers?: string[];
/**
* Max Agentic Loops
* @description How many follow-up model calls one intercepted request may chain. Empty applies the default of 3.
*/
max_agentic_loops?: number | null;
/**
* Search Tool Name
* @description Name of the configured search tool to run searches through. Empty uses the first one available.
*/
search_tool_name?: string | null;
};
/**
* WebSearchInterceptionSettingsResponse
* @description Response model for web search interception settings
*/
WebSearchInterceptionSettingsResponse: {
/** Field Schema */
field_schema: {
[key: string]: unknown;
};
/** Values */
values: {
[key: string]: unknown;
};
};
/** WorkerRegistryEntry */
WorkerRegistryEntry: {
/** Name */
@ -49624,6 +49709,26 @@ export interface operations {
};
};
};
get_websearch_interception_settings_get_websearch_interception_settings_get: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["WebSearchInterceptionSettingsResponse"];
};
};
};
};
get_favicon_get_favicon_get: {
parameters: {
query?: never;
@ -62751,6 +62856,39 @@ export interface operations {
};
};
};
update_websearch_interception_settings_update_websearch_interception_settings_patch: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["WebSearchInterceptionSettings"];
};
};
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
upload_logo_upload_logo_post: {
parameters: {
query?: never;