diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index d0006f1a091..d5d1382e635 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -9,6 +9,8 @@ import asyncio +import hashlib +import hmac import json import re import threading @@ -29,6 +31,7 @@ from litellm.constants import ( PRESIDIO_ANALYZE_CHUNK_CONCURRENCY, PRESIDIO_ANALYZE_CHUNK_OVERLAP_CHARS, ) +from litellm.secret_managers.main import get_secret_str from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: @@ -169,6 +172,29 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener yield chunk +def _resolved_token_salt(value: str | None) -> str: + """Resolve the salt from an ``os.environ/`` reference, the way guardrail + api_key and api_base are already resolved. + + A literal is refused rather than used. The salt is the HMAC key, so it has to + survive as a secret for the tokens to mean anything, and a literal does not: + guardrail_registry logs the whole params mapping at debug before + initialization, so a literal reaches the proxy log in full and anyone with + the log can test candidate values against the tokens they can see.""" + if value is None: + return "" + if not value.startswith("os.environ/"): + raise ValueError( + "presidio_token_salt must be an os.environ/ reference, not a literal. " + "The salt is the HMAC key behind every stable token, and guardrail params " + "are logged in full at debug level." + ) + resolved: Final = get_secret_str(value) + if resolved is None or not resolved.strip(): + raise ValueError(f"presidio_token_salt: {value!r} resolves to an unset or blank environment variable") + return resolved + + class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): user_api_key_cache = None ad_hoc_recognizers: list[str] | None = None @@ -201,6 +227,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): presidio_entities_deny_list: list[PiiEntityType | str] | None = None, presidio_analyze_chunk_size_bytes: int | None = None, _callback_role: Literal["scan", "restore"] | None = None, + presidio_stable_tokens: bool | None = None, + presidio_token_salt: str | None = None, **kwargs, ): if logging_only is True: @@ -231,6 +259,13 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self.presidio_entities_deny_list: list[PiiEntityType | str] = presidio_entities_deny_list or [] self.presidio_language = presidio_language or "en" self.presidio_analyze_chunk_size_bytes: int = self._coerce_analyze_chunk_size(presidio_analyze_chunk_size_bytes) + self.presidio_stable_tokens: bool = bool(presidio_stable_tokens) + self.presidio_token_salt: str = _resolved_token_salt(presidio_token_salt) + if self.presidio_stable_tokens and not self.presidio_token_salt: + raise ValueError( + "presidio_stable_tokens requires presidio_token_salt. An unkeyed digest lets the " + "model provider recover a masked value by hashing candidates and comparing the prefix." + ) # Shared HTTP session to prevent memory leaks (issue #14540) self._http_session: aiohttp.ClientSession | None = None # Lock to prevent race conditions when creating session under concurrent load @@ -786,6 +821,16 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): masked_entity_count[entity_type] = masked_entity_count.get(entity_type, 0) + 1 return redacted_text["text"] + _STABLE_TOKEN_HEX_CHARS: Final[int] = 8 + + def _stable_token_suffix(self, entity_type: str, value: str) -> str: + # The guardrail name namespaces the digest, so two guardrails with + # different salts cannot produce the same token for the same value. + namespace: Final = self.guardrail_name or "" + message: Final = f"{namespace}\x00{entity_type}\x00{value}".encode() + digest: Final = hmac.new(self.presidio_token_salt.encode(), message, hashlib.sha256).hexdigest() + return digest[: self._STABLE_TOKEN_HEX_CHARS] + def _finalize_presidio_anonymize_numbered_tokens( self, text: str, @@ -814,7 +859,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): # Assign sequence numbers in forward (left-to-right) order so # that is the first entity in the text, etc. sorted_forward: Final = sorted(analyze_results, key=lambda x: x["start"]) - seq_map: Final = {} + seq_map: Final[dict[tuple[int, int], int]] = {} for idx, ar in enumerate(sorted_forward, start=1): seq_map[(ar["start"], ar["end"])] = idx @@ -824,14 +869,17 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): for ar in reversed(sorted_forward): start = ar["start"] end = ar["end"] - entity_type = ar["entity_type"] - replacement = f"<{entity_type}>" - seq = seq_map[(start, end)] - if replacement.endswith(">"): - replacement = f"{replacement[:-1]}_{seq}>" - else: - replacement = f"{replacement}_{seq}" - pii_tokens[replacement] = text[start:end] + # Annotated because the analyzer result is an untyped mapping: both now + # reach _stable_token_suffix as call arguments, where an Any counts. + entity_type: str = ar["entity_type"] + value: str = text[start:end] + suffix = ( + self._stable_token_suffix(entity_type, value) + if self.presidio_stable_tokens + else str(seq_map[(start, end)]) + ) + replacement = f"<{entity_type}_{suffix}>" + pii_tokens[replacement] = value new_text = new_text[:start] + replacement + new_text[end:] masked_entity_count[entity_type] = masked_entity_count.get(entity_type, 0) + 1 return new_text diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index b7f3726d017..413faa11961 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -137,6 +137,11 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> _OPTIONAL_PresidioPIIMasking, ) + # Read through getattr with a declared type: these two are guardrail-only + # params, so they are not fields on LitellmParams, and an unannotated + # getattr returns Any and counts against the unknown-argument budget. + stable_tokens: Final[bool | None] = getattr(litellm_params, "presidio_stable_tokens", None) + token_salt: Final[str | None] = getattr(litellm_params, "presidio_token_salt", None) explicit_filter_scope: Final = litellm_params.presidio_filter_scope filter_scope: Final = explicit_filter_scope or ("input" if _is_mcp_only_mode(litellm_params.mode) else "both") run_input: Final = filter_scope in ("input", "both") @@ -156,6 +161,8 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base, presidio_language=litellm_params.presidio_language, presidio_entities_deny_list=litellm_params.presidio_entities_deny_list, + presidio_stable_tokens=stable_tokens, + presidio_token_salt=token_salt, apply_to_output=False, timeout=litellm_params.timeout, _callback_role="scan", diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py index 08acec0d7ac..be52d8be77f 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -6,9 +6,10 @@ Tests PII detection and masking for different message formats import asyncio import copy import json +import os import re from contextlib import asynccontextmanager -from typing import Final, Literal +from typing import Dict, Final, List, Literal, Tuple from unittest.mock import MagicMock, patch from aiohttp import web @@ -4288,3 +4289,195 @@ async def test_standalone_restoration_preserves_post_call_selection(event_hook: data, UserAPIKeyAuth(request_route="/v1/chat/completions"), response ) assert response.choices[0].message.content == "Jane" + + +def _analyze_result(entity_type: str, text: str, value: str) -> Dict[str, object]: + start = text.index(value) + return {"entity_type": entity_type, "start": start, "end": start + len(value), "score": 0.9} + + +def _mask( + guardrail: _OPTIONAL_PresidioPIIMasking, text: str, results: List[Dict[str, object]] +) -> Tuple[str, Dict[str, str]]: + request_data: Dict[str, object] = {} + masked = guardrail._finalize_presidio_anonymize_numbered_tokens( + text=text, analyze_results=results, request_data=request_data, masked_entity_count={} + ) + return masked, request_data["metadata"]["pii_tokens"] + + +def _stable_guardrail(salt: str = "unit-test-salt", guardrail_name: str | None = None) -> _OPTIONAL_PresidioPIIMasking: + """Build a stable-token guardrail on a salt held where the real one has to be. + + The salt is only ever read through an os.environ reference, so the tests put + it there rather than passing a literal the constructor refuses.""" + var = f"PRESIDIO_TEST_SALT_{abs(hash(salt)) % 10**8}" + os.environ[var] = salt + return _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + presidio_stable_tokens=True, + presidio_token_salt=f"os.environ/{var}", + **({"guardrail_name": guardrail_name} if guardrail_name else {}), + ) + + +def test_stable_tokens_are_identical_across_requests(): + """The same value keeps its token as the conversation grows""" + guardrail = _stable_guardrail() + + first_text = "Call Alice Brenner" + second_text = "Earlier you asked. Then Bob Smith replied. Call Alice Brenner" + first, first_tokens = _mask(guardrail, first_text, [_analyze_result("PERSON", first_text, "Alice Brenner")]) + second, second_tokens = _mask( + guardrail, + second_text, + [ + _analyze_result("PERSON", second_text, "Bob Smith"), + _analyze_result("PERSON", second_text, "Alice Brenner"), + ], + ) + + alice_token = next(token for token, value in first_tokens.items() if value == "Alice Brenner") + assert alice_token in second + assert second_tokens[alice_token] == "Alice Brenner" + + +def test_counter_tokens_drift_across_requests(): + """The default numbering is what stable tokens exist to replace""" + guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) + + first_text = "Call Alice Brenner" + second_text = "Earlier you asked. Then Bob Smith replied. Call Alice Brenner" + _, first_tokens = _mask(guardrail, first_text, [_analyze_result("PERSON", first_text, "Alice Brenner")]) + _, second_tokens = _mask( + guardrail, + second_text, + [ + _analyze_result("PERSON", second_text, "Bob Smith"), + _analyze_result("PERSON", second_text, "Alice Brenner"), + ], + ) + + assert first_tokens[""] == "Alice Brenner" + assert second_tokens[""] == "Bob Smith" + assert second_tokens[""] == "Alice Brenner" + + +def test_stable_tokens_differ_per_value_and_per_entity_type(): + guardrail = _stable_guardrail() + + text = "Alice Brenner and Bob Smith" + _, tokens = _mask( + guardrail, + text, + [_analyze_result("PERSON", text, "Alice Brenner"), _analyze_result("PERSON", text, "Bob Smith")], + ) + assert len(set(tokens)) == 2 + + _, person = _mask(guardrail, "Toronto", [_analyze_result("PERSON", "Toronto", "Toronto")]) + _, location = _mask(guardrail, "Toronto", [_analyze_result("LOCATION", "Toronto", "Toronto")]) + assert set(person) != set(location) + + +def test_stable_tokens_depend_on_the_salt(): + """Without this, a token could be reversed by hashing candidate values""" + entity = [_analyze_result("PERSON", "Alice Brenner", "Alice Brenner")] + salted, salted_tokens = _mask(_stable_guardrail("salt-a"), "Alice Brenner", entity) + other, other_tokens = _mask(_stable_guardrail("salt-b"), "Alice Brenner", entity) + + assert salted != other + assert set(salted_tokens) != set(other_tokens) + + +def test_stable_tokens_round_trip_through_unmasking(): + guardrail = _stable_guardrail() + text = "Call Alice Brenner today" + masked, tokens = _mask(guardrail, text, [_analyze_result("PERSON", text, "Alice Brenner")]) + + assert guardrail._unmask_pii_text(masked, tokens) == "Call Alice Brenner today" + + +def test_stable_tokens_are_off_by_default(): + guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) + assert guardrail.presidio_stable_tokens is False + + text = "Call Alice Brenner" + masked, _ = _mask(guardrail, text, [_analyze_result("PERSON", text, "Alice Brenner")]) + assert masked == "Call " + + +def test_stable_token_config_reaches_the_guardrail(monkeypatch): + from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio + + monkeypatch.setenv("PRESIDIO_ANALYZER_API_BASE", "http://localhost:5002") + monkeypatch.setenv("PRESIDIO_ANONYMIZER_API_BASE", "http://localhost:5001") + params = LitellmParams( + guardrail="presidio", + mode="pre_call", + output_parse_pii=True, + presidio_stable_tokens=True, + presidio_token_salt="os.environ/PRESIDIO_CONFIG_SALT", + ) + monkeypatch.setenv("PRESIDIO_CONFIG_SALT", "from-config") + callbacks = initialize_presidio(params, {"guardrail_name": "presidio-unit", "litellm_params": params}) + + assert callbacks + for callback in callbacks: + assert callback.presidio_stable_tokens is True + assert callback.presidio_token_salt == "from-config" + + +def test_a_literal_salt_is_refused(monkeypatch): + """The salt is the HMAC key, and guardrail params are logged in full at debug. + + A literal would therefore reach the proxy log, where anyone holding it could + test candidate values against the tokens they can already see. + """ + monkeypatch.setenv("PRESIDIO_ANALYZER_API_BASE", "http://localhost:5002") + monkeypatch.setenv("PRESIDIO_ANONYMIZER_API_BASE", "http://localhost:5001") + + with pytest.raises(ValueError, match=r"os\.environ"): + _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + presidio_stable_tokens=True, + presidio_token_salt="a-literal-secret", + ) + + +def test_stable_token_salt_resolves_an_os_environ_reference(monkeypatch): + monkeypatch.setenv("MY_PII_SALT", "from-env") + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + presidio_stable_tokens=True, + presidio_token_salt="os.environ/MY_PII_SALT", + ) + assert guardrail.presidio_token_salt == "from-env" + + +def test_stable_token_salt_refuses_an_unset_os_environ_reference(monkeypatch): + monkeypatch.delenv("MY_PII_SALT", raising=False) + with pytest.raises(ValueError, match="unset or blank"): + _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + presidio_stable_tokens=True, + presidio_token_salt="os.environ/MY_PII_SALT", + ) + + +def test_stable_tokens_require_a_salt(): + """An unkeyed digest lets a provider recover the value by hashing candidates""" + with pytest.raises(ValueError, match="presidio_token_salt"): + _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True, presidio_stable_tokens=True) + + +def test_stable_tokens_are_namespaced_per_guardrail(): + """One salt shared by two guardrails must not produce one token""" + first = _stable_guardrail("shared", guardrail_name="tenant-a") + second = _stable_guardrail("shared", guardrail_name="tenant-b") + entity = [_analyze_result("PERSON", "Alice Brenner", "Alice Brenner")] + + assert _mask(first, "Alice Brenner", entity)[0] != _mask(second, "Alice Brenner", entity)[0]