mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(ui): save the MCP private ranges and client allowlist sequentially
The proxy stores both fields with a whole-row read-modify-write of general_settings, so two concurrent writes from one save can drop one of them. Also drops docstrings and suppressions the diff did not need Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
21cafc8780
commit
f29be6e1ee
5 changed files with 45 additions and 13 deletions
|
|
@ -71,7 +71,6 @@ def parse_allowed_mcp_clients(raw_setting: object) -> frozenset[str] | None:
|
|||
|
||||
|
||||
def allowed_mcp_clients_from_general_settings(general_settings: object) -> frozenset[str] | None:
|
||||
"""Reads the allowlist out of the proxy's untyped general_settings mapping."""
|
||||
return parse_allowed_mcp_clients(
|
||||
_GENERAL_SETTINGS_ADAPTER.validate_python(general_settings).get(MCP_ALLOWED_CLIENTS_SETTING)
|
||||
)
|
||||
|
|
@ -88,7 +87,6 @@ def extract_mcp_client_name(body: bytes) -> str | None:
|
|||
|
||||
|
||||
def check_mcp_client_allowed(body: bytes, allowed_clients: frozenset[str] | None) -> MCPClientRejection | None:
|
||||
"""None when the initialize is admitted, otherwise the rejection to send back as a 403."""
|
||||
if allowed_clients is None:
|
||||
return None
|
||||
client_name: Final = extract_mcp_client_name(body)
|
||||
|
|
|
|||
|
|
@ -3833,7 +3833,6 @@ if MCP_AVAILABLE:
|
|||
body: bytes,
|
||||
client_ip: str | None,
|
||||
) -> bool:
|
||||
"""Send a 403 and return True when the initialize body names a client the gateway does not admit."""
|
||||
rejection: Final = check_mcp_client_allowed(body, _load_allowed_mcp_clients())
|
||||
if rejection is None:
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -1999,7 +1999,7 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
|
|||
stateless_handle: Final = AsyncMock(side_effect=handle_request)
|
||||
stateful_handle: Final = AsyncMock(side_effect=handle_request)
|
||||
with (
|
||||
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(
|
||||
|
|
@ -2012,11 +2012,11 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
|
|||
),
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True),
|
||||
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.session_manager_stateless",
|
||||
SimpleNamespace(handle_request=stateless_handle),
|
||||
),
|
||||
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.session_manager_stateful",
|
||||
SimpleNamespace(handle_request=stateful_handle),
|
||||
),
|
||||
|
|
@ -2339,7 +2339,7 @@ async def test_mcp_routing_chunked_initialize_to_stateful():
|
|||
stateful_called.append(1)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
|
|
@ -2451,7 +2451,7 @@ async def test_mcp_routing_caps_body_peek_for_oversized_chunked_body():
|
|||
raise AssertionError("non-initialize POST should not reach stateful manager")
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
|
|
@ -2600,7 +2600,7 @@ async def test_mcp_routing_initialize_rejected_when_owner_at_session_cap():
|
|||
stateful_called.append(1)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
|
|
@ -2692,7 +2692,7 @@ async def test_stateful_mcp_requests_refresh_session_auth_context():
|
|||
captured_context = callback_context.run(get_auth_context)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(
|
||||
|
|
@ -3228,7 +3228,7 @@ async def test_stateful_mcp_session_owner_mismatch_returns_403():
|
|||
handle_request_mock = AsyncMock()
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(intruder_auth, None, None, None, None, None),
|
||||
|
|
|
|||
|
|
@ -217,4 +217,35 @@ describe("MCPNetworkSettings", () => {
|
|||
await waitFor(() => expect(toast.fromError).toHaveBeenCalledWith(rangeFailure));
|
||||
expect(toast.success).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("writes the private ranges and the allowed clients one after the other, never concurrently", async () => {
|
||||
vi.mocked(fetchMCPClientIp).mockResolvedValue("203.0.113.45");
|
||||
let finishRangeWrite: (() => void) | undefined;
|
||||
vi.mocked(updateConfigFieldSetting).mockImplementation(
|
||||
(_token, fieldName) =>
|
||||
new Promise<void>((resolve) => {
|
||||
if (fieldName === "mcp_internal_ip_ranges") {
|
||||
finishRangeWrite = resolve;
|
||||
} else {
|
||||
resolve();
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
renderSettings();
|
||||
await userEvent.click(await screen.findByText("203.0.113.0/24"));
|
||||
await userEvent.type(screen.getByRole("textbox", { name: "Allowed client names" }), "codex-mcp-client{Enter}");
|
||||
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
|
||||
|
||||
await waitFor(() =>
|
||||
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges", ["203.0.113.0/24"]),
|
||||
);
|
||||
expect(updateConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_allowed_clients", expect.anything());
|
||||
|
||||
finishRangeWrite?.();
|
||||
await waitFor(() =>
|
||||
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", ["codex-mcp-client"]),
|
||||
);
|
||||
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("MCP network settings saved"));
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -92,12 +92,14 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
|
|||
const handleSave = async () => {
|
||||
if (!accessToken) return;
|
||||
setSaving(true);
|
||||
const results = await Promise.allSettled([
|
||||
const [rangeResult] = await Promise.allSettled([
|
||||
persistList(accessToken, "mcp_internal_ip_ranges", {
|
||||
value: privateRanges,
|
||||
stored: storedRanges,
|
||||
setStored: setStoredRanges,
|
||||
}),
|
||||
]);
|
||||
const [clientResult] = await Promise.allSettled([
|
||||
persistList(accessToken, "mcp_allowed_clients", {
|
||||
value: allowedClients,
|
||||
stored: storedClients,
|
||||
|
|
@ -105,7 +107,9 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
|
|||
}),
|
||||
]);
|
||||
setSaving(false);
|
||||
const failures = results.filter((result): result is PromiseRejectedResult => result.status === "rejected");
|
||||
const failures = [rangeResult, clientResult].filter(
|
||||
(result): result is PromiseRejectedResult => result.status === "rejected",
|
||||
);
|
||||
if (failures.length === 0) {
|
||||
toast.success("MCP network settings saved");
|
||||
return;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue