mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(mcp): keep a created MCP server when post-write registry refresh fails
add_mcp_server wrote the new row and then reloaded the entire registry from the database inside one try block. A single pre-existing malformed row made the reload raise, so the endpoint returned 500 even though the new server was already persisted; callers assumed failure and retried, creating duplicate servers. Split the flow so the database write is the commit point and still 500s on failure, while the in-memory registry refresh is best-effort and only logged on error. Add regression tests for both the refresh-fails-after-commit path and the db-write-fails path
This commit is contained in:
parent
0dcf316f59
commit
ba20078c05
2 changed files with 119 additions and 5 deletions
|
|
@ -1471,23 +1471,34 @@ if MCP_AVAILABLE:
|
|||
payload.submitted_by = None
|
||||
payload.submitted_at = None
|
||||
|
||||
# Attempt to create the mcp server
|
||||
# The database write is the commit point: if it fails nothing was
|
||||
# persisted and the request is a genuine failure.
|
||||
try:
|
||||
new_mcp_server = await create_mcp_server(
|
||||
prisma_client,
|
||||
payload,
|
||||
touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||
)
|
||||
await global_mcp_server_manager.add_server(new_mcp_server)
|
||||
|
||||
# Ensure registry is up to date by reloading from database
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating mcp server: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": f"Error creating mcp server: {str(e)}"},
|
||||
)
|
||||
|
||||
# Registry refresh is best-effort: the row is already committed, so a
|
||||
# failure here (e.g. an unrelated malformed row in the table) must not
|
||||
# surface as a 500 and orphan the created server, which would push the
|
||||
# caller to retry and create duplicates.
|
||||
try:
|
||||
await global_mcp_server_manager.add_server(new_mcp_server)
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
f"MCP server {new_mcp_server.server_id} created but in-memory "
|
||||
f"registry refresh failed: {str(e)}"
|
||||
)
|
||||
|
||||
return _redact_mcp_credentials(new_mcp_server)
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -2407,6 +2407,109 @@ class TestUpdateMCPServer:
|
|||
assert result.alias == "Updated Test Server"
|
||||
|
||||
|
||||
class TestAddMCPServerAtomicity:
|
||||
"""A committed MCP server must survive a post-write registry refresh failure.
|
||||
|
||||
Regression: add_mcp_server inserted the row and then reloaded the whole
|
||||
registry from the database inside the same try block. One unrelated malformed
|
||||
row made the reload raise, so the endpoint returned 500 even though the new
|
||||
row was already persisted. Callers assumed failure and retried, creating
|
||||
duplicate servers.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_succeeds_when_registry_refresh_fails(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
add_mcp_server,
|
||||
)
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
alias="echo",
|
||||
url="https://echo.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
admin = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"
|
||||
)
|
||||
created_server = generate_mock_mcp_server_db_record(
|
||||
server_id="created-1", alias="echo"
|
||||
)
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.add_server = AsyncMock()
|
||||
mock_manager.reload_servers_from_database = AsyncMock(
|
||||
side_effect=Exception("malformed pre-existing row")
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
|
||||
MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
|
||||
AsyncMock(return_value=created_server),
|
||||
) as create_mock,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
result = await add_mcp_server(payload=payload, user_api_key_dict=admin)
|
||||
|
||||
create_mock.assert_awaited_once()
|
||||
mock_manager.reload_servers_from_database.assert_awaited_once()
|
||||
assert result.server_id == "created-1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_500s_and_skips_registry_when_db_write_fails(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
add_mcp_server,
|
||||
)
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
alias="echo",
|
||||
url="https://echo.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
admin = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"
|
||||
)
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.add_server = AsyncMock()
|
||||
mock_manager.reload_servers_from_database = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
|
||||
MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
|
||||
AsyncMock(side_effect=Exception("db down")),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await add_mcp_server(payload=payload, user_api_key_dict=admin)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
mock_manager.add_server.assert_not_awaited()
|
||||
mock_manager.reload_servers_from_database.assert_not_awaited()
|
||||
|
||||
|
||||
class TestHealthCheckServers:
|
||||
"""Test suite for health check servers endpoint"""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue