diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index ee1ebcb5ce9..4176f57d9de 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index c663e63414c..adcdfea4711 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -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 diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 1a56227b008..504219a64e1 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -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