From 19126a5208ca3cfc0b5dc38f0183df670c1725b4 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Fri, 22 May 2026 13:10:50 +0000 Subject: [PATCH] fix(purview): use threading.Lock and getattr for LitellmParams - Replace asyncio.Lock with threading.Lock in PurviewGuardrailBase. The cache lock is acquired both from the proxy's main event loop and from short-lived event loops created by the logging_hook thread fallback. In Python 3.10+ an asyncio.Lock is bound to the first event loop that acquires it, so the second loop would silently break audit logging with RuntimeError. All critical sections are in-memory dict ops with no awaits, so a synchronous lock is safe. - Use getattr() on LitellmParams in initialize_guardrail() instead of .get(), which does not exist on Pydantic BaseModel instances and would raise AttributeError at runtime. Tests updated to construct Mock objects with spec= so they reflect the real interface. Co-authored-by: Yassin Kortam --- .../microsoft_purview/__init__.py | 14 ++-- .../guardrail_hooks/microsoft_purview/base.py | 22 ++++-- .../guardrail_hooks/test_microsoft_purview.py | 69 +++++++++++++++---- 3 files changed, 78 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py index 273089d4b4a..42e0ea62dd7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py @@ -12,12 +12,14 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" import litellm from litellm.types.guardrails import GuardrailEventHooks - tenant_id = litellm_params.get("tenant_id") - client_id = litellm_params.get("client_id") + tenant_id = getattr(litellm_params, "tenant_id", None) + client_id = getattr(litellm_params, "client_id", None) # client_secret can be passed via the standard api_key field or as # a dedicated client_secret parameter. - client_secret = litellm_params.api_key or litellm_params.get("client_secret") + client_secret = litellm_params.api_key or getattr( + litellm_params, "client_secret", None + ) if not tenant_id: raise ValueError("Microsoft Purview: tenant_id is required") @@ -40,8 +42,10 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" tenant_id=str(tenant_id), client_id=str(client_id), client_secret=str(client_secret), - purview_app_name=str(litellm_params.get("purview_app_name") or "LiteLLM"), - user_id_field=str(litellm_params.get("user_id_field") or "user_id"), + purview_app_name=str( + getattr(litellm_params, "purview_app_name", None) or "LiteLLM" + ), + user_id_field=str(getattr(litellm_params, "user_id_field", None) or "user_id"), logging_only=logging_only, event_hook=litellm_params.mode, default_on=litellm_params.default_on, diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index 40a36a6e502..63b8ae2c963 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -1,4 +1,4 @@ -import asyncio +import threading import time import uuid from collections import OrderedDict @@ -68,7 +68,15 @@ class PurviewGuardrailBase: OrderedDict() ) self._scope_cache_maxsize = 1000 - self._cache_lock = asyncio.Lock() + # Use a threading.Lock (not asyncio.Lock) because this lock is acquired + # from both the proxy's main asyncio event loop and from short-lived + # event loops created by the logging_hook thread fallback. In Python + # 3.10+ an asyncio.Lock is bound to the first event loop that acquires + # it and raises RuntimeError from any other loop, which would silently + # break audit logging via the thread fallback. All critical sections + # below are pure in-memory dict ops with no awaits, so a synchronous + # lock is both correct and sufficient. + self._cache_lock = threading.Lock() @staticmethod def _encode_graph_user_id(user_id: str) -> str: @@ -82,7 +90,7 @@ class PurviewGuardrailBase: async def _get_access_token(self) -> str: """Acquire or return cached OAuth2 token via client_credentials grant.""" now = time.time() - async with self._cache_lock: + with self._cache_lock: if self._token_cache and self._token_cache[1] > now + 60: return self._token_cache[0] @@ -102,7 +110,7 @@ class PurviewGuardrailBase: token_data = response.json() access_token = token_data["access_token"] expires_in = int(token_data.get("expires_in", 3599)) - async with self._cache_lock: + with self._cache_lock: self._token_cache = (access_token, now + expires_in) verbose_proxy_logger.debug( "Purview: acquired new OAuth2 token (expires_in=%ds)", expires_in @@ -157,7 +165,7 @@ class PurviewGuardrailBase: encoded_user_id = self._encode_graph_user_id(user_id) now = time.time() - async with self._cache_lock: + with self._cache_lock: cached = self._scope_cache.get(user_id) if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS: self._scope_cache.move_to_end(user_id) @@ -180,7 +188,7 @@ class PurviewGuardrailBase: response_json, response_headers = await self._graph_post(url, body) etag = response_headers.get("etag", response_headers.get("ETag", "")) - async with self._cache_lock: + with self._cache_lock: self._scope_cache[user_id] = (etag, response_json, now) # Move refreshed entry to the end so it is treated as most-recently-used. # OrderedDict.__setitem__ preserves existing insertion order for known @@ -258,7 +266,7 @@ class PurviewGuardrailBase: # If policies changed, invalidate scope cache so next call re-fetches. if response_json.get("protectionScopeState") == "modified": - async with self._cache_lock: + with self._cache_lock: self._scope_cache.pop(user_id, None) return response_json diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py index 3bd91d85d2a..2c002b4b1c2 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py @@ -418,7 +418,9 @@ class TestPostCallHook: with pytest.raises(HTTPException) as exc_info: await guardrail.async_post_call_success_hook( data={}, - user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + user_api_key_dict=UserAPIKeyAuth( + api_key="test", user_id="user-123" + ), response=response, ) @@ -662,9 +664,10 @@ class TestResponsesAPIHooks: ) mock_check.assert_called_once() - assert "policy text in instructions only" in mock_check.call_args.kwargs[ - "text" - ] + assert ( + "policy text in instructions only" + in mock_check.call_args.kwargs["text"] + ) @pytest.mark.asyncio async def test_post_call_responses_api_output_text(self): @@ -1217,8 +1220,21 @@ class TestInitializerValidation: initialize_guardrail, ) - litellm_params = Mock() - litellm_params.get.return_value = None + litellm_params = Mock( + spec=[ + "tenant_id", + "client_id", + "client_secret", + "purview_app_name", + "user_id_field", + "api_key", + "mode", + "default_on", + ] + ) + litellm_params.tenant_id = None + litellm_params.client_id = None + litellm_params.client_secret = None litellm_params.api_key = "secret" litellm_params.mode = "pre_call" @@ -1230,10 +1246,21 @@ class TestInitializerValidation: initialize_guardrail, ) - litellm_params = Mock() - litellm_params.get.side_effect = lambda k, *a: ( - "test-tenant" if k == "tenant_id" else None + litellm_params = Mock( + spec=[ + "tenant_id", + "client_id", + "client_secret", + "purview_app_name", + "user_id_field", + "api_key", + "mode", + "default_on", + ] ) + litellm_params.tenant_id = "test-tenant" + litellm_params.client_id = None + litellm_params.client_secret = None litellm_params.api_key = "secret" litellm_params.mode = "pre_call" @@ -1245,11 +1272,21 @@ class TestInitializerValidation: initialize_guardrail, ) - litellm_params = Mock() - litellm_params.get.side_effect = lambda k, *a: { - "tenant_id": "test-tenant", - "client_id": "test-client", - }.get(k) + litellm_params = Mock( + spec=[ + "tenant_id", + "client_id", + "client_secret", + "purview_app_name", + "user_id_field", + "api_key", + "mode", + "default_on", + ] + ) + litellm_params.tenant_id = "test-tenant" + litellm_params.client_id = "test-client" + litellm_params.client_secret = None litellm_params.api_key = None litellm_params.mode = "pre_call" @@ -1745,7 +1782,9 @@ class TestGraphUserIdEncoding: guardrail.async_handler.post = AsyncMock(side_effect=_capture_post) - with patch.object(guardrail, "_get_access_token", new_callable=AsyncMock) as mock_token: + with patch.object( + guardrail, "_get_access_token", new_callable=AsyncMock + ) as mock_token: mock_token.return_value = "tok" await guardrail._compute_protection_scopes("user/with%special")