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 <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-22 13:10:50 +00:00
parent c92bb00a23
commit 19126a5208
No known key found for this signature in database
3 changed files with 78 additions and 27 deletions

View file

@ -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,

View file

@ -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

View file

@ -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")