mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
5107f205a0
commit
b4bb2a77a2
3 changed files with 126 additions and 78 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue