mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
c92bb00a23
commit
19126a5208
3 changed files with 78 additions and 27 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue