test: inject the HIBP client and the MCP loop clock so two backend tests stop flaking (#44007)

* test(proxy): inject the HIBP client into the breached-password update test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): drive optional-discovery deadlines with an injected loop clock

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test: bound MCP deadline checks, support Python 3.10, pin HIBP URL

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-01 05:16:49 -07:00 • committed by GitHub
parent 5107f205a0
commit b4bb2a77a2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 126 additions and 78 deletions

View file

@ -181,7 +181,10 @@ def _team_membership_table(
async def _hash_password_in_dict(
data: dict, general_settings: Mapping[str, object], password_prevalidated: bool = False
data: dict,
general_settings: Mapping[str, object],
password_prevalidated: bool = False,
hibp_client: AsyncHTTPHandler | None = None,
) -> None:
"""Validate and hash password field in-place if present.
@ -193,7 +196,7 @@ async def _hash_password_in_dict(
if "password" in data and data["password"] is not None:
if not password_prevalidated:
validate_password_policy(data["password"], general_settings)
await validate_password_not_breached(data["password"], general_settings)
await validate_password_not_breached(data["password"], general_settings, hibp_client)
data["password"] = hash_password(data["password"])
data["password_reset_required"] = True
data["last_breach_check_at"] = None
@ -1459,6 +1462,7 @@ async def _update_single_user_helper(
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: str | None = None,
password_prevalidated: bool = False,
hibp_client: AsyncHTTPHandler | None = None,
) -> dict[str, Any]:
"""
Helper function to update a single user.
@ -1481,7 +1485,12 @@ async def _update_single_user_helper(
data_json: Final[dict] = user_request.model_dump(exclude_unset=True)
non_default_values = _update_internal_user_params(data_json=data_json, data=user_request)
await _hash_password_in_dict(non_default_values, general_settings, password_prevalidated=password_prevalidated)
await _hash_password_in_dict(
non_default_values,
general_settings,
password_prevalidated=password_prevalidated,
hibp_client=hibp_client,
)
existing_user_row: BaseModel | None = None
if user_request.user_id:

View file

@ -10,7 +10,6 @@ from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
import respx
from fastapi import HTTPException
from fastapi.testclient import TestClient
from pytest_mock import MockerFixture
@ -4547,7 +4546,6 @@ async def test_user_update_hashes_and_persists_strong_password(_admin_prisma, mo
@pytest.mark.asyncio
@respx.mock
async def test_user_update_rejects_breached_password(_admin_prisma):
"""A strength-passing password found in the HIBP corpus must be rejected
before it ever reaches the DB write."""
@ -4557,19 +4555,26 @@ async def test_user_update_rejects_breached_password(_admin_prisma):
password = "Str0ng!Passw0rd"
sha1 = hashlib.sha1(password.encode("utf-8"), usedforsecurity=False).hexdigest().upper()
respx.get(f"https://api.pwnedpasswords.com/range/{sha1[:5]}").mock(
return_value=httpx.Response(200, text=f"{sha1[5:]}:1387")
)
lookups: Final[list[tuple[str, str]]] = [] # mutable-ok: capture the injected handler request method and URL
def handler(request: httpx.Request) -> httpx.Response:
lookups.append((request.method, str(request.url)))
return httpx.Response(200, text=f"{sha1[5:]}:1387")
user_request = UpdateUserRequest(user_id="target-user", password=password)
admin_caller = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
with pytest.raises(ProxyException) as exc_info:
await _update_single_user_helper(user_request=user_request, user_api_key_dict=admin_caller)
await _update_single_user_helper(
user_request=user_request,
user_api_key_dict=admin_caller,
hibp_client=_hibp_client_with_handler(handler),
)
assert exc_info.value.code == "400"
assert "data breaches" in exc_info.value.message
_admin_prisma.db.litellm_usertable.find_first.assert_not_called()
assert lookups == [("GET", f"https://api.pwnedpasswords.com/range/{sha1[:5]}")]
@pytest.mark.asyncio

View file

@ -81,6 +81,20 @@ class _MockTransportClient(MCPClient):
return streamable_http_client(self.server_url, http_client=http_client), http_client
class _ManualClockLoop(asyncio.SelectorEventLoop):
"""An event loop whose clock moves only when the test advances it, so timeouts fire on test-controlled conditions"""
def __init__(self) -> None:
super().__init__()
self._now = 0.0
def time(self) -> float:
return self._now
def advance(self, seconds: float) -> None:
self._now += seconds
class _FakeExceptionGroup(Exception):
"""Duck-typed stand-in for an anyio/builtin ExceptionGroup.
@ -2309,16 +2323,20 @@ async def test_optional_discovery_collects_all_pages(method: str, session_id: st
assert sum(call.args[0].method == "DELETE" for call in responder.call_args_list) == (1 if session_id else 0)
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("prompts/list", "resources/list", "resources/templates/list"))
@pytest.mark.parametrize(
"failure", ("repeat", "cycle", "cap", "method_not_found", "internal_error", "unauthorized", "deadline")
)
@pytest.mark.parametrize("strict", (False, True))
async def test_optional_discovery_rejects_incomplete_walks(
def test_optional_discovery_rejects_incomplete_walks(
method: str, failure: str, strict: bool, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
monkeypatch.setattr(mcp_client_module, "MCP_TOOL_LISTING_MAX_PAGES", 3 if failure == "cycle" else 2, raising=False)
monkeypatch.setattr(
mcp_client_module,
"MCP_TOOL_LISTING_MAX_PAGES",
3 if failure in ("cycle", "repeat") else 2,
raising=False,
)
monkeypatch.setattr(mcp_client_module, "MCP_TOOL_LISTING_TIMEOUT", 0.05)
field: Final = {
"prompts/list": "prompts",
@ -2330,82 +2348,98 @@ async def test_optional_discovery_rejects_incomplete_walks(
"resources/list": {"name": "first", "uri": "test://first"},
"resources/templates/list": {"name": "first", "uriTemplate": "test://{name}"},
}[method]
cancelled: Final = asyncio.Event()
loop: Final = _ManualClockLoop()
async def respond(request: httpx2.Request) -> httpx2.Response:
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
if payload.method == "initialize":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {"prompts": {}, "resources": {}},
"serverInfo": {"name": "interrupted", "version": "1"},
},
},
)
assert payload.method == method
cursor: Final = (payload.params or {}).get("cursor")
if cursor is not None:
if failure == "deadline":
try:
await asyncio.Event().wait()
finally:
cancelled.set()
if failure == "unauthorized":
return httpx2.Response(401)
if failure in ("method_not_found", "internal_error"):
async def run() -> None:
cancelled: Final = asyncio.Event()
async def respond(request: httpx2.Request) -> httpx2.Response:
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
if payload.method == "initialize":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"error": {
"code": -32601 if failure == "method_not_found" else -32603,
"message": "Later page unavailable",
"result": {
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {"prompts": {}, "resources": {}},
"serverInfo": {"name": "interrupted", "version": "1"},
},
},
)
next_cursor: Final = (
"private-cursor-2" if cursor == "private-cursor-1" and failure != "repeat" else "private-cursor-1"
)
return httpx2.Response(
200, json={"jsonrpc": "2.0", "id": payload.id, "result": {field: [entry], "nextCursor": next_cursor}}
)
assert payload.method == method
cursor: Final = (payload.params or {}).get("cursor")
if cursor is not None:
if failure == "deadline":
loop.advance(0.15)
try:
for _ in range(1_000):
await asyncio.sleep(0)
except asyncio.CancelledError:
cancelled.set()
raise
return httpx2.Response(500)
if failure == "unauthorized":
return httpx2.Response(401)
if failure in ("method_not_found", "internal_error"):
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"error": {
"code": -32601 if failure == "method_not_found" else -32603,
"message": "Later page unavailable",
},
},
)
if failure == "deadline" and cursor is None:
loop.advance(0.1)
next_cursor: Final = (
"private-cursor-2" if cursor == "private-cursor-1" and failure != "repeat" else "private-cursor-1"
)
return httpx2.Response(
200, json={"jsonrpc": "2.0", "id": payload.id, "result": {field: [entry], "nextCursor": next_cursor}}
)
responder: Final = AsyncMock(side_effect=respond)
client: Final = _MockTransportClient(responder, server_url="https://example.com/mcp", timeout=0.2)
operation: Final = {
"prompts/list": client.list_prompts,
"resources/list": client.list_resources,
"resources/templates/list": client.list_resource_templates,
}[method]
if strict:
error_type: Final = {
"internal_error": MCPError,
"unauthorized": httpx2.HTTPStatusError,
"deadline": TimeoutError,
}.get(failure, RuntimeError)
with pytest.raises(error_type):
await operation(raise_on_error=True)
else:
assert await operation() == []
assert len(
tuple(
payload
for call in responder.call_args_list
if isinstance(payload := _JSONRPC_MESSAGE_ADAPTER.validate_json(call.args[0].content), JSONRPCRequest)
and payload.method == method
)
) == (3 if failure == "cycle" else 2)
assert "private-cursor" not in caplog.text
if failure == "deadline":
assert cancelled.is_set()
responder: Final = AsyncMock(side_effect=respond)
client: Final = _MockTransportClient(responder, server_url="https://example.com/mcp", timeout=0.2)
operation: Final = {
"prompts/list": client.list_prompts,
"resources/list": client.list_resources,
"resources/templates/list": client.list_resource_templates,
}[method]
if strict:
error_type: Final = {
"internal_error": MCPError,
"unauthorized": httpx2.HTTPStatusError,
"deadline": TimeoutError,
}.get(failure, RuntimeError)
with pytest.raises(error_type):
await operation(raise_on_error=True)
else:
assert await operation() == []
assert len(
tuple(
payload
for call in responder.call_args_list
if isinstance(payload := _JSONRPC_MESSAGE_ADAPTER.validate_json(call.args[0].content), JSONRPCRequest)
and payload.method == method
)
) == (3 if failure == "cycle" else 2)
assert "private-cursor" not in caplog.text
if failure == "deadline":
assert cancelled.is_set()
try:
loop.run_until_complete(run())
finally:
loop.run_until_complete(loop.shutdown_asyncgens())
loop.run_until_complete(loop.shutdown_default_executor())
loop.close()
@pytest.mark.asyncio