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:
yassin 2026-09-17 20:05:35 +00:00
parent 21cafc8780
commit f29be6e1ee
5 changed files with 45 additions and 13 deletions

View file

@ -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)

View file

@ -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

View file

@ -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),

View file

@ -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"));
});
});

View file

@ -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;