mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat: add support for configurable confidence score thresholds and scope in Presidio PII masking (#17817)
* feat: add support for configurable confidence score thresholds in Presidio PII masking * feat: enhance Presidio PII masking with configurable score thresholds and behavior documentation * feat: add configurable output masking and filter scope for Presidio PII guardrail
This commit is contained in:
parent
2a864e25f8
commit
756c60540e
7 changed files with 575 additions and 52 deletions
|
|
@ -220,11 +220,28 @@ When connecting Litellm to Langfuse, you can see the guardrail information on th
|
|||
style={{width: '60%', display: 'block', margin: '0'}}
|
||||
/>
|
||||
|
||||
## Entity Type Configuration
|
||||
## Entity Types, Detection Confidence Score Threshold, and Scope Configuration
|
||||
|
||||
You can configure specific entity types for PII detection and decide how to handle each entity type (mask or block).
|
||||
- **Entity Types**
|
||||
- You can configure specific entity types for PII detection and decide how to handle each entity type (mask or block).
|
||||
- **Detection Confidence Score Threshold**
|
||||
- You can also provide an optional confidence score threshold at which detections will be passed to the anonymizer. Entities without an entry in `presidio_score_thresholds` keep all detections (no minimum score).
|
||||
- **Scope**
|
||||
- Use the optional `presidio_filter_scope` to choose where checks run:
|
||||
|
||||
### Configure Entity Types in config.yaml
|
||||
- `input`: only user → model content is scanned
|
||||
- `output`: only model → user content is scanned
|
||||
- `both` (default): scan both directions
|
||||
|
||||
**What about `output_parse_pii`?**
|
||||
This flag only un-masks tokens back to the originals after the model call; it does not run Presidio detection on outputs. Use `presidio_filter_scope: output` (or `both`) when you want Presidio to actively scan and mask the model’s response before it reaches the user.
|
||||
|
||||
**When to pick input vs output:**
|
||||
- `input`: Protect upstream providers; strip PII before it leaves your boundary.
|
||||
- `output`: Catch PII the model might generate or leak back to users.
|
||||
- `both`: End-to-end protection in both directions.
|
||||
|
||||
### Configure Entity Types, Detection Confidence Score Threshold, and Scope in `config.yaml`
|
||||
|
||||
Define your guardrails with specific entity type configuration:
|
||||
|
||||
|
|
@ -240,6 +257,11 @@ guardrails:
|
|||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: "pre_mcp_call" # Use this mode for MCP requests
|
||||
presidio_filter_scope: both # input | output | both, optional
|
||||
presidio_score_thresholds: # Optional
|
||||
ALL: 0.7 # Default confidence threshold applied to all entities
|
||||
CREDIT_CARD: 0.8 # Override for credit cards
|
||||
EMAIL_ADDRESS: 0.6 # Override for emails
|
||||
pii_entities_config:
|
||||
CREDIT_CARD: "MASK" # Will mask credit card numbers
|
||||
EMAIL_ADDRESS: "MASK" # Will mask email addresses
|
||||
|
|
@ -248,10 +270,19 @@ guardrails:
|
|||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: "pre_call" # Use this mode for regular LLM requests
|
||||
presidio_filter_scope: both # input | output | both, optional
|
||||
presidio_score_thresholds: # Optional
|
||||
CREDIT_CARD: 0.8 # Only keep credit card detections scoring 0.8+
|
||||
pii_entities_config:
|
||||
CREDIT_CARD: "BLOCK" # Will block requests containing credit card numbers
|
||||
```
|
||||
|
||||
#### Confidence threshold behavior:
|
||||
- No `presidio_score_thresholds`: keep all detections (no thresholds applied)
|
||||
- `presidio_score_thresholds.ALL`: apply this confidence threshold to every detection
|
||||
- `presidio_score_thresholds.<ENTITY>`: apply only to that entity
|
||||
- If both `ALL` and an entity override exist, `ALL` applies globally and the entity override takes precedence for that entity
|
||||
|
||||
### Supported Entity Types
|
||||
|
||||
LiteLLM Supports all Presidio entity types. See the complete list of presidio entity types [here](https://microsoft.github.io/presidio/supported_entities/).
|
||||
|
|
@ -357,6 +388,10 @@ guardrails:
|
|||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: "pre_mcp_call"
|
||||
presidio_filter_scope: both # input | output | both
|
||||
presidio_score_thresholds:
|
||||
CREDIT_CARD: 0.8 # Only keep credit card detections scoring 0.8+
|
||||
EMAIL_ADDRESS: 0.6 # Only keep email detections scoring 0.6+
|
||||
pii_entities_config:
|
||||
CREDIT_CARD: "MASK" # Will mask credit card numbers
|
||||
EMAIL_ADDRESS: "BLOCK" # Will block email addresses
|
||||
|
|
@ -674,5 +709,3 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
|||
```text title="Logged Response with Masked PII" showLineNumbers
|
||||
Hi, my name is <PERSON>!
|
||||
```
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -45,6 +45,20 @@ guardrails:
|
|||
description: "Score between 0-1 indicating content toxicity level"
|
||||
- name: "pii_detection"
|
||||
type: "boolean"
|
||||
|
||||
# Example Presidio guardrail config with entity actions + confidence score thresholds
|
||||
- guardrail_name: "presidio-pii"
|
||||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: "pre_call"
|
||||
presidio_language: "en"
|
||||
pii_entities_config:
|
||||
CREDIT_CARD: "MASK"
|
||||
EMAIL_ADDRESS: "MASK"
|
||||
US_SSN: "MASK"
|
||||
presidio_score_thresholds: # minimum confidence scores for keeping detections
|
||||
CREDIT_CARD: 0.8
|
||||
EMAIL_ADDRESS: 0.6
|
||||
```
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -123,6 +123,9 @@ guardrails:
|
|||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: "pre_call" # Run before LLM call
|
||||
presidio_score_thresholds: # optional confidence score thresholds for detections
|
||||
CREDIT_CARD: 0.8
|
||||
EMAIL_ADDRESS: 0.6
|
||||
pii_entities_config:
|
||||
CREDIT_CARD: "MASK"
|
||||
EMAIL_ADDRESS: "MASK"
|
||||
|
|
|
|||
|
|
@ -72,12 +72,16 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
presidio_analyzer_api_base: Optional[str] = None,
|
||||
presidio_anonymizer_api_base: Optional[str] = None,
|
||||
output_parse_pii: Optional[bool] = False,
|
||||
apply_to_output: bool = False,
|
||||
presidio_ad_hoc_recognizers: Optional[str] = None,
|
||||
logging_only: Optional[bool] = None,
|
||||
pii_entities_config: Optional[
|
||||
Dict[Union[PiiEntityType, str], PiiAction]
|
||||
] = None,
|
||||
presidio_language: Optional[str] = None,
|
||||
presidio_score_thresholds: Optional[
|
||||
Dict[Union[PiiEntityType, str], float]
|
||||
] = None,
|
||||
**kwargs,
|
||||
):
|
||||
if logging_only is True:
|
||||
|
|
@ -90,9 +94,13 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
) # mapping of PII token to original text - only used with Presidio `replace` operation
|
||||
self.mock_redacted_text = mock_redacted_text
|
||||
self.output_parse_pii = output_parse_pii or False
|
||||
self.apply_to_output = apply_to_output
|
||||
self.pii_entities_config: Dict[Union[PiiEntityType, str], PiiAction] = (
|
||||
pii_entities_config or {}
|
||||
)
|
||||
self.presidio_score_thresholds: Dict[Union[PiiEntityType, str], float] = (
|
||||
presidio_score_thresholds or {}
|
||||
)
|
||||
self.presidio_language = presidio_language or "en"
|
||||
if mock_testing is True: # for testing purposes only
|
||||
return
|
||||
|
|
@ -239,7 +247,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
async with session.post(analyze_url, json=analyze_payload) as response:
|
||||
analyze_results = await response.json()
|
||||
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
|
||||
|
||||
|
||||
# Handle error responses from Presidio (e.g., {'error': 'No text provided'})
|
||||
# Presidio may return a dict instead of a list when errors occur
|
||||
if isinstance(analyze_results, dict):
|
||||
|
|
@ -261,7 +269,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
e
|
||||
)
|
||||
return []
|
||||
|
||||
|
||||
# Normal case: list of results
|
||||
final_results = []
|
||||
for item in analyze_results:
|
||||
|
|
@ -272,7 +280,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
verbose_proxy_logger.warning(
|
||||
"Skipping invalid Presidio result item: %s (error: %s)",
|
||||
item,
|
||||
te
|
||||
te,
|
||||
)
|
||||
continue
|
||||
return final_results
|
||||
|
|
@ -290,6 +298,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
Send analysis results to the Presidio anonymizer endpoint to get redacted text
|
||||
"""
|
||||
try:
|
||||
# If there are no detections after filtering, return the original text
|
||||
if isinstance(analyze_results, list) and len(analyze_results) == 0:
|
||||
return text
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Make the request to /anonymize
|
||||
anonymize_url = f"{self.presidio_anonymizer_api_base}anonymize"
|
||||
|
|
@ -333,6 +345,37 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
def filter_analyze_results_by_score(
|
||||
self, analyze_results: Union[List[PresidioAnalyzeResponseItem], Dict]
|
||||
) -> Union[List[PresidioAnalyzeResponseItem], Dict]:
|
||||
"""
|
||||
Drop detections that fall below configured per-entity score thresholds.
|
||||
"""
|
||||
if not self.presidio_score_thresholds:
|
||||
return analyze_results
|
||||
|
||||
if not isinstance(analyze_results, list):
|
||||
return analyze_results
|
||||
|
||||
filtered_results: List[PresidioAnalyzeResponseItem] = []
|
||||
for item in analyze_results:
|
||||
entity_type = item.get("entity_type")
|
||||
score = item.get("score")
|
||||
|
||||
threshold = None
|
||||
if entity_type is not None:
|
||||
threshold = self.presidio_score_thresholds.get(entity_type)
|
||||
if threshold is None:
|
||||
threshold = self.presidio_score_thresholds.get("ALL")
|
||||
|
||||
if threshold is not None:
|
||||
if score is None or score < threshold:
|
||||
continue
|
||||
|
||||
filtered_results.append(item)
|
||||
|
||||
return filtered_results
|
||||
|
||||
def raise_exception_if_blocked_entities_detected(
|
||||
self, analyze_results: Union[List[PresidioAnalyzeResponseItem], Dict]
|
||||
):
|
||||
|
|
@ -389,6 +432,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
|
||||
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
|
||||
|
||||
# Apply score threshold filtering if configured
|
||||
analyze_results = self.filter_analyze_results_by_score(
|
||||
analyze_results=analyze_results
|
||||
)
|
||||
|
||||
####################################################
|
||||
# Blocked Entities check
|
||||
####################################################
|
||||
|
|
@ -455,9 +503,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
if messages is None:
|
||||
return data
|
||||
tasks = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = (
|
||||
[]
|
||||
) # Track (message_index, content_index) for each task
|
||||
task_mappings: List[
|
||||
Tuple[int, Optional[int]]
|
||||
] = [] # Track (message_index, content_index) for each task
|
||||
|
||||
for msg_idx, m in enumerate(messages):
|
||||
content = m.get("content", None)
|
||||
|
|
@ -558,9 +606,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
): # /chat/completions requests
|
||||
messages: Optional[List] = kwargs.get("messages", None)
|
||||
tasks = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = (
|
||||
[]
|
||||
) # Track (message_index, content_index) for each task
|
||||
task_mappings: List[
|
||||
Tuple[int, Optional[int]]
|
||||
] = [] # Track (message_index, content_index) for each task
|
||||
|
||||
if messages is None:
|
||||
return kwargs, result
|
||||
|
|
@ -635,6 +683,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
f"PII Masking Args: self.output_parse_pii={self.output_parse_pii}; type of response={type(response)}"
|
||||
)
|
||||
|
||||
if self.apply_to_output is True:
|
||||
return await self._mask_output_response(
|
||||
response=response, request_data=data
|
||||
)
|
||||
|
||||
if self.output_parse_pii is False and litellm.output_parse_pii is False:
|
||||
return response
|
||||
|
||||
|
|
@ -651,6 +704,52 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
].message.content.replace(key, value)
|
||||
return response
|
||||
|
||||
async def _mask_output_response(
|
||||
self,
|
||||
response: Union[ModelResponse, EmbeddingResponse, ImageResponse],
|
||||
request_data: dict,
|
||||
):
|
||||
"""
|
||||
Apply Presidio masking on model responses (non-streaming).
|
||||
"""
|
||||
if not isinstance(response, ModelResponse):
|
||||
return response
|
||||
|
||||
# skip streaming here; handled in async_post_call_streaming_iterator_hook
|
||||
if response.choices and isinstance(response.choices[0], StreamingChoices):
|
||||
return response
|
||||
|
||||
presidio_config = self.get_presidio_settings_from_request_data(
|
||||
request_data or {}
|
||||
)
|
||||
|
||||
for choice in response.choices:
|
||||
content = getattr(choice.message, "content", None)
|
||||
if content is None:
|
||||
continue
|
||||
if isinstance(content, str):
|
||||
choice.message.content = await self.check_pii(
|
||||
text=content,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=request_data,
|
||||
)
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
text_value = item.get("text")
|
||||
if text_value is None:
|
||||
continue
|
||||
item["text"] = await self.check_pii(
|
||||
text=text_value,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -663,6 +762,74 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
If PII processing is enabled, this collects all chunks, applies PII unmasking,
|
||||
and returns a reconstructed stream. Otherwise, it passes through the original stream.
|
||||
"""
|
||||
# If we need to mask model output, collect the full stream, apply masking, and replay it.
|
||||
if self.apply_to_output:
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
try:
|
||||
collected_content = ""
|
||||
last_chunk = None
|
||||
|
||||
async for chunk in response:
|
||||
last_chunk = chunk
|
||||
|
||||
if (
|
||||
hasattr(chunk, "choices")
|
||||
and chunk.choices
|
||||
and hasattr(chunk.choices[0], "delta")
|
||||
and hasattr(chunk.choices[0].delta, "content")
|
||||
and isinstance(chunk.choices[0].delta.content, str)
|
||||
):
|
||||
collected_content += chunk.choices[0].delta.content
|
||||
|
||||
if not last_chunk:
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
presidio_config = self.get_presidio_settings_from_request_data(
|
||||
request_data or {}
|
||||
)
|
||||
masked_content = await self.check_pii(
|
||||
text=collected_content,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
mock_response = MockResponseIterator(
|
||||
model_response=ModelResponse(
|
||||
id=last_chunk.id,
|
||||
object=last_chunk.object,
|
||||
created=last_chunk.created,
|
||||
model=last_chunk.model,
|
||||
choices=[
|
||||
Choices(
|
||||
message=Message(
|
||||
role="assistant",
|
||||
content=masked_content,
|
||||
),
|
||||
index=0,
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
),
|
||||
json_mode=False,
|
||||
)
|
||||
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"Error masking streaming PII output: {str(e)}"
|
||||
)
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
# If PII unmasking not needed, just pass through the original stream
|
||||
if not (self.output_parse_pii and self.pii_tokens):
|
||||
async for chunk in response:
|
||||
|
|
@ -787,3 +954,5 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
super().update_in_memory_litellm_params(litellm_params)
|
||||
if litellm_params.pii_entities_config:
|
||||
self.pii_entities_config = litellm_params.pii_entities_config
|
||||
if litellm_params.presidio_score_thresholds:
|
||||
self.presidio_score_thresholds = litellm_params.presidio_score_thresholds
|
||||
|
|
|
|||
|
|
@ -75,34 +75,51 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
_OPTIONAL_PresidioPIIMasking,
|
||||
)
|
||||
|
||||
_presidio_callback = _OPTIONAL_PresidioPIIMasking(
|
||||
guardrail_name=guardrail.get("guardrail_name", ""),
|
||||
event_hook=litellm_params.mode,
|
||||
output_parse_pii=litellm_params.output_parse_pii,
|
||||
presidio_ad_hoc_recognizers=litellm_params.presidio_ad_hoc_recognizers,
|
||||
mock_redacted_text=litellm_params.mock_redacted_text,
|
||||
default_on=litellm_params.default_on,
|
||||
pii_entities_config=litellm_params.pii_entities_config,
|
||||
presidio_analyzer_api_base=litellm_params.presidio_analyzer_api_base,
|
||||
presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base,
|
||||
presidio_language=litellm_params.presidio_language,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_presidio_callback)
|
||||
filter_scope = getattr(litellm_params, "presidio_filter_scope", None) or "both"
|
||||
run_input = filter_scope in ("input", "both")
|
||||
run_output = filter_scope in ("output", "both")
|
||||
|
||||
if litellm_params.output_parse_pii:
|
||||
_success_callback = _OPTIONAL_PresidioPIIMasking(
|
||||
output_parse_pii=True,
|
||||
def _make_presidio_callback(**overrides):
|
||||
params = dict(
|
||||
guardrail_name=guardrail.get("guardrail_name", ""),
|
||||
event_hook=GuardrailEventHooks.post_call.value,
|
||||
event_hook=litellm_params.mode,
|
||||
output_parse_pii=litellm_params.output_parse_pii,
|
||||
presidio_ad_hoc_recognizers=litellm_params.presidio_ad_hoc_recognizers,
|
||||
mock_redacted_text=litellm_params.mock_redacted_text,
|
||||
default_on=litellm_params.default_on,
|
||||
pii_entities_config=litellm_params.pii_entities_config,
|
||||
presidio_score_thresholds=litellm_params.presidio_score_thresholds,
|
||||
presidio_analyzer_api_base=litellm_params.presidio_analyzer_api_base,
|
||||
presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base,
|
||||
presidio_language=litellm_params.presidio_language,
|
||||
apply_to_output=False,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_success_callback)
|
||||
params.update(overrides)
|
||||
callback = _OPTIONAL_PresidioPIIMasking(**params)
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback)
|
||||
return callback
|
||||
|
||||
return _presidio_callback
|
||||
primary_callback = None
|
||||
|
||||
if run_input:
|
||||
primary_callback = _make_presidio_callback()
|
||||
|
||||
if litellm_params.output_parse_pii:
|
||||
_make_presidio_callback(
|
||||
output_parse_pii=True,
|
||||
event_hook=GuardrailEventHooks.post_call.value,
|
||||
)
|
||||
|
||||
if run_output:
|
||||
output_callback = _make_presidio_callback(
|
||||
apply_to_output=True,
|
||||
event_hook=GuardrailEventHooks.post_call.value,
|
||||
output_parse_pii=False,
|
||||
)
|
||||
if primary_callback is None:
|
||||
primary_callback = output_callback
|
||||
|
||||
return primary_callback
|
||||
|
||||
|
||||
def initialize_hide_secrets(litellm_params: LitellmParams, guardrail: Guardrail):
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from typing import Any, Dict, List, Literal, Optional, Union
|
|||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk, ChatCompletionToolParam
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionToolCallChunk,
|
||||
|
|
@ -269,6 +269,13 @@ class PresidioPresidioConfigModelUserInterface(BaseModel):
|
|||
default=None,
|
||||
description="Base URL for the Presidio anonymizer API",
|
||||
)
|
||||
presidio_filter_scope: Optional[Literal["input", "output", "both"]] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Where to apply Presidio checks: 'input' (user -> model), "
|
||||
"'output' (model -> user), or 'both' (default)."
|
||||
),
|
||||
)
|
||||
output_parse_pii: Optional[bool] = Field(
|
||||
default=None,
|
||||
description="When True, LiteLLM will replace the masked text with the original text in the response",
|
||||
|
|
@ -279,6 +286,10 @@ class PresidioPresidioConfigModelUserInterface(BaseModel):
|
|||
default="en",
|
||||
description="Language code for Presidio PII analysis (e.g., 'en', 'de', 'es', 'fr')",
|
||||
)
|
||||
presidio_run_on: Optional[Literal["input", "output", "both"]] = Field(
|
||||
default=None,
|
||||
description="Where to apply Presidio checks: input, output, or both (default).",
|
||||
)
|
||||
|
||||
|
||||
class PresidioConfigModel(PresidioPresidioConfigModelUserInterface):
|
||||
|
|
@ -287,6 +298,22 @@ class PresidioConfigModel(PresidioPresidioConfigModelUserInterface):
|
|||
pii_entities_config: Optional[Dict[Union[PiiEntityType, str], PiiAction]] = Field(
|
||||
default=None, description="Configuration for PII entity types and actions"
|
||||
)
|
||||
presidio_filter_scope: Literal["input", "output", "both"] = Field(
|
||||
default="both",
|
||||
description=(
|
||||
"Where to apply Presidio checks: 'input' runs on user → model traffic, "
|
||||
"'output' runs on model → user traffic, and 'both' applies to both."
|
||||
),
|
||||
)
|
||||
presidio_score_thresholds: Optional[
|
||||
Dict[Union[PiiEntityType, str], float]
|
||||
] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Optional per-entity minimum confidence scores for Presidio detections. "
|
||||
"Entities below the threshold are ignored."
|
||||
),
|
||||
)
|
||||
presidio_ad_hoc_recognizers: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Path to a JSON file containing ad-hoc recognizers for Presidio",
|
||||
|
|
|
|||
|
|
@ -18,7 +18,9 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
|
||||
_OPTIONAL_PresidioPIIMasking,
|
||||
)
|
||||
from litellm.types.guardrails import PiiAction, PiiEntityType
|
||||
from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
import litellm
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -604,6 +606,7 @@ async def test_request_data_flows_to_apply_guardrail():
|
|||
presidio = _OPTIONAL_PresidioPIIMasking(
|
||||
guardrail_name="test_presidio",
|
||||
output_parse_pii=True,
|
||||
mock_testing=True,
|
||||
)
|
||||
|
||||
request_data = {
|
||||
|
|
@ -634,6 +637,109 @@ async def test_request_data_flows_to_apply_guardrail():
|
|||
print("✓ request_data correctly passed to apply_guardrail")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_output_masking_apply_to_output_only(mock_user_api_key):
|
||||
"""
|
||||
Ensure output masking runs when apply_to_output is enabled.
|
||||
"""
|
||||
|
||||
presidio = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.MASK},
|
||||
)
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
return text.replace("4111-1111-1111-1111", "[CREDIT_CARD]")
|
||||
|
||||
presidio.check_pii = mock_check_pii
|
||||
|
||||
response = ModelResponse(
|
||||
id="1",
|
||||
object="chat.completion",
|
||||
created=0,
|
||||
model="gpt-test",
|
||||
choices=[
|
||||
Choices(
|
||||
message=Message(
|
||||
role="assistant",
|
||||
content="Card is 4111-1111-1111-1111",
|
||||
),
|
||||
index=0,
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
result = await presidio.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
response=response,
|
||||
)
|
||||
|
||||
assert "[CREDIT_CARD]" in result.choices[0].message.content
|
||||
assert "4111-1111-1111-1111" not in result.choices[0].message.content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_presidio_filter_scope_initializer(monkeypatch):
|
||||
"""
|
||||
Ensure initializer respects presidio_filter_scope for input/output/both.
|
||||
"""
|
||||
|
||||
created = []
|
||||
|
||||
class DummyGuardrail:
|
||||
def __init__(self, apply_to_output: bool = False, event_hook=None, **kwargs):
|
||||
self.apply_to_output = apply_to_output
|
||||
self.event_hook = event_hook
|
||||
created.append(self)
|
||||
|
||||
def update_in_memory_litellm_params(self, litellm_params):
|
||||
pass
|
||||
|
||||
class DummyManager:
|
||||
def __init__(self):
|
||||
self.added = []
|
||||
|
||||
def add_litellm_callback(self, cb):
|
||||
self.added.append(cb)
|
||||
|
||||
mgr = DummyManager()
|
||||
monkeypatch.setattr(litellm, "logging_callback_manager", mgr, raising=False)
|
||||
import litellm.proxy.guardrails.guardrail_initializers as gi
|
||||
import litellm.proxy.guardrails.guardrail_hooks.presidio as presidio_mod
|
||||
monkeypatch.setattr(
|
||||
presidio_mod, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False
|
||||
)
|
||||
monkeypatch.setattr(gi, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False)
|
||||
|
||||
# input-only
|
||||
created.clear()
|
||||
from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio
|
||||
|
||||
params_input = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="input")
|
||||
guardrail_dict = {"guardrail_name": "g1"}
|
||||
cb = initialize_presidio(params_input, guardrail_dict)
|
||||
assert cb is created[0]
|
||||
assert created[0].apply_to_output is False
|
||||
|
||||
# output-only
|
||||
created.clear()
|
||||
params_output = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="output")
|
||||
cb = initialize_presidio(params_output, guardrail_dict)
|
||||
assert len(created) == 1
|
||||
assert created[0].apply_to_output is True
|
||||
|
||||
# both -> expect two callbacks (input + output)
|
||||
created.clear()
|
||||
params_both = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="both")
|
||||
cb = initialize_presidio(params_both, guardrail_dict)
|
||||
assert len(created) == 2
|
||||
assert any(not c.apply_to_output for c in created)
|
||||
assert any(c.apply_to_output for c in created)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_content_handling(presidio_guardrail, mock_user_api_key, mock_cache):
|
||||
"""
|
||||
|
|
@ -856,21 +962,175 @@ async def test_tool_calling_complete_scenario(presidio_guardrail, mock_user_api_
|
|||
print("✓ Tool calling complete scenario test passed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests
|
||||
asyncio.run(
|
||||
test_multimodal_message_format_completion_call_type(
|
||||
_OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
output_parse_pii=False,
|
||||
pii_entities_config={
|
||||
PiiEntityType.CREDIT_CARD: PiiAction.MASK,
|
||||
PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK,
|
||||
PiiEntityType.PHONE_NUMBER: PiiAction.MASK,
|
||||
},
|
||||
),
|
||||
UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
|
||||
MagicMock(spec=DualCache),
|
||||
)
|
||||
def test_filter_drops_low_score_detection():
|
||||
"""
|
||||
Detections below the configured score threshold should be removed.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8},
|
||||
)
|
||||
print("\n✅ All Presidio tests passed!")
|
||||
analyze_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}
|
||||
]
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(analyze_results)
|
||||
assert filtered == []
|
||||
|
||||
|
||||
def test_filter_preserves_high_score_detection():
|
||||
"""
|
||||
Detections meeting the score threshold should be preserved.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8},
|
||||
)
|
||||
analyze_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.9, "start": 0, "end": 4}
|
||||
]
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(analyze_results)
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["entity_type"] == PiiEntityType.CREDIT_CARD
|
||||
|
||||
|
||||
def test_no_thresholds_returns_all():
|
||||
"""
|
||||
With no thresholds configured, all detections are kept.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True)
|
||||
analyze_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.1, "start": 0, "end": 4},
|
||||
{"entity_type": PiiEntityType.EMAIL_ADDRESS, "score": 0.2, "start": 5, "end": 9},
|
||||
]
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(analyze_results)
|
||||
assert len(filtered) == 2
|
||||
|
||||
|
||||
def test_entity_specific_threshold_only_applies_to_that_entity():
|
||||
"""
|
||||
Entity-specific thresholds do not affect other entity types.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8},
|
||||
)
|
||||
analyze_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4},
|
||||
{"entity_type": PiiEntityType.EMAIL_ADDRESS, "score": 0.1, "start": 5, "end": 9},
|
||||
]
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(analyze_results)
|
||||
# CREDIT_CARD is filtered, EMAIL_ADDRESS is kept because no threshold
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["entity_type"] == PiiEntityType.EMAIL_ADDRESS
|
||||
|
||||
|
||||
def test_filter_uses_default_all_threshold():
|
||||
"""
|
||||
Default ALL threshold applies to any entity without a specific override.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
presidio_score_thresholds={"ALL": 0.75},
|
||||
)
|
||||
analyze_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4},
|
||||
{"entity_type": PiiEntityType.EMAIL_ADDRESS, "score": 0.8, "start": 5, "end": 9},
|
||||
]
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(analyze_results)
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["entity_type"] == PiiEntityType.EMAIL_ADDRESS
|
||||
|
||||
|
||||
def test_entity_specific_overrides_default_threshold():
|
||||
"""
|
||||
Entity-specific threshold should override the ALL default.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
presidio_score_thresholds={
|
||||
"ALL": 0.8,
|
||||
PiiEntityType.CREDIT_CARD: 0.6,
|
||||
},
|
||||
)
|
||||
analyze_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.65, "start": 0, "end": 4},
|
||||
{"entity_type": PiiEntityType.EMAIL_ADDRESS, "score": 0.75, "start": 5, "end": 9},
|
||||
]
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(analyze_results)
|
||||
# CREDIT_CARD passes due to override, EMAIL_ADDRESS dropped by ALL threshold
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["entity_type"] == PiiEntityType.CREDIT_CARD
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anonymize_skips_when_no_detections_after_filter():
|
||||
"""
|
||||
When all detections are filtered out, anonymize_text should return the original text.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8},
|
||||
)
|
||||
masked_entity_count = {}
|
||||
text = "4111"
|
||||
|
||||
filtered = guardrail.filter_analyze_results_by_score(
|
||||
[{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}]
|
||||
)
|
||||
|
||||
result = await guardrail.anonymize_text(
|
||||
text=text,
|
||||
analyze_results=filtered,
|
||||
output_parse_pii=False,
|
||||
masked_entity_count=masked_entity_count,
|
||||
)
|
||||
|
||||
assert result == text
|
||||
assert masked_entity_count == {}
|
||||
|
||||
|
||||
def test_blocking_respects_threshold_filter():
|
||||
"""
|
||||
Entities filtered out by score should not trigger blocking, but high-score detections should.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.BLOCK},
|
||||
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.9},
|
||||
)
|
||||
|
||||
low_score_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}
|
||||
]
|
||||
filtered = guardrail.filter_analyze_results_by_score(low_score_results)
|
||||
guardrail.raise_exception_if_blocked_entities_detected(filtered)
|
||||
|
||||
high_score_results = [
|
||||
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.95, "start": 0, "end": 4}
|
||||
]
|
||||
filtered_high = guardrail.filter_analyze_results_by_score(high_score_results)
|
||||
with pytest.raises(Exception):
|
||||
guardrail.raise_exception_if_blocked_entities_detected(filtered_high)
|
||||
|
||||
|
||||
def test_update_in_memory_applies_score_thresholds():
|
||||
"""
|
||||
update_in_memory_litellm_params should refresh score thresholds.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True)
|
||||
assert guardrail.presidio_score_thresholds == {}
|
||||
|
||||
params = LitellmParams(
|
||||
guardrail="presidio",
|
||||
mode="pre_call",
|
||||
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.85},
|
||||
)
|
||||
guardrail.update_in_memory_litellm_params(params)
|
||||
|
||||
assert guardrail.presidio_score_thresholds == {PiiEntityType.CREDIT_CARD: 0.85}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue