mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
PR fixes
This commit is contained in:
parent
b473344f70
commit
b83b497d38
2 changed files with 72 additions and 21 deletions
|
|
@ -28,6 +28,7 @@ from litellm.types.utils import EmbeddingResponse, ImageResponse
|
|||
# Constants
|
||||
USER_ROLE: Final[Literal["user"]] = "user"
|
||||
ASSISTANT_ROLE: Final[Literal["assistant"]] = "assistant"
|
||||
SENSITIVE_DATA_DETECTOR_KEYS: Final[list[str]] = ["sensitiveData", "dataDetector"]
|
||||
|
||||
# Type aliases
|
||||
MessageRole = Literal["user", "assistant"]
|
||||
|
|
@ -96,7 +97,7 @@ class NomaBlockedMessage(HTTPException):
|
|||
if filtered_topics:
|
||||
result[key] = filtered_topics
|
||||
|
||||
elif key in ["sensitiveData", "dataDetector"] and isinstance(value, dict):
|
||||
elif key in SENSITIVE_DATA_DETECTOR_KEYS and isinstance(value, dict):
|
||||
filtered_sensitive = {}
|
||||
for data_type, data_result in value.items():
|
||||
if self._is_result_true(data_result):
|
||||
|
|
@ -308,7 +309,7 @@ class NomaGuardrail(CustomGuardrail):
|
|||
sensitive_data_detected = False
|
||||
|
||||
for key, value in classification_obj.items():
|
||||
if key in ["sensitiveData", "dataDetector"] and isinstance(value, dict):
|
||||
if key in SENSITIVE_DATA_DETECTOR_KEYS and isinstance(value, dict):
|
||||
# Check if any sensitive data detector has result=true
|
||||
for data_type, data_result in value.items():
|
||||
if self._is_result_true(data_result):
|
||||
|
|
@ -411,20 +412,17 @@ class NomaGuardrail(CustomGuardrail):
|
|||
|
||||
def _replace_user_message_content(
|
||||
self, request_data: dict, anonymized_content: str
|
||||
) -> dict:
|
||||
):
|
||||
"""
|
||||
Replace the user message content in request data with anonymized version.
|
||||
|
||||
Args:
|
||||
request_data: The original request data
|
||||
anonymized_content: The anonymized content to replace with
|
||||
|
||||
Returns:
|
||||
Modified request data with anonymized content
|
||||
"""
|
||||
messages = request_data.get("messages", [])
|
||||
if not messages:
|
||||
return request_data
|
||||
return
|
||||
|
||||
# Find and replace the last user message
|
||||
for i in range(len(messages) - 1, -1, -1):
|
||||
|
|
@ -432,31 +430,24 @@ class NomaGuardrail(CustomGuardrail):
|
|||
messages[i]["content"] = anonymized_content
|
||||
break
|
||||
|
||||
return request_data
|
||||
|
||||
def _replace_llm_response_content(
|
||||
self, response: LLMResponse, anonymized_content: str
|
||||
) -> LLMResponse:
|
||||
):
|
||||
"""
|
||||
Replace the LLM response content with anonymized version.
|
||||
|
||||
Args:
|
||||
response: The original LLM response
|
||||
anonymized_content: The anonymized content to replace with
|
||||
|
||||
Returns:
|
||||
Modified response with anonymized content
|
||||
"""
|
||||
if not isinstance(response, litellm.ModelResponse):
|
||||
return response
|
||||
return
|
||||
|
||||
# Replace content in all choices
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, litellm.Choices) and choice.message.content:
|
||||
choice.message.content = anonymized_content
|
||||
|
||||
return response
|
||||
|
||||
async def _check_user_message_background(
|
||||
self,
|
||||
request_data: dict,
|
||||
|
|
|
|||
|
|
@ -1052,13 +1052,13 @@ class TestNomaAnonymizationLogic:
|
|||
]
|
||||
}
|
||||
|
||||
result = anonymize_guardrail._replace_user_message_content(
|
||||
anonymize_guardrail._replace_user_message_content(
|
||||
request_data, "My phone is *******"
|
||||
)
|
||||
|
||||
# Should replace the last user message
|
||||
assert result["messages"][-1]["content"] == "My phone is *******"
|
||||
assert result["messages"][1]["content"] == "My email is test@example.com" # Unchanged
|
||||
assert request_data["messages"][-1]["content"] == "My phone is *******"
|
||||
assert request_data["messages"][1]["content"] == "My email is test@example.com" # Unchanged
|
||||
|
||||
def test_replace_llm_response_content(self, anonymize_guardrail):
|
||||
"""Test _replace_llm_response_content"""
|
||||
|
|
@ -1080,11 +1080,11 @@ class TestNomaAnonymizationLogic:
|
|||
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
||||
)
|
||||
|
||||
result = anonymize_guardrail._replace_llm_response_content(
|
||||
anonymize_guardrail._replace_llm_response_content(
|
||||
response, "Your email is *******"
|
||||
)
|
||||
|
||||
assert result.choices[0].message.content == "Your email is *******"
|
||||
assert response.choices[0].message.content == "Your email is *******"
|
||||
|
||||
|
||||
class TestNomaAnonymizationFlow:
|
||||
|
|
@ -1446,3 +1446,63 @@ class TestNomaAnonymizationFlow:
|
|||
data=request_data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anonymization_llm_response_no_anonymized_content_available(
|
||||
self, anonymize_guardrail, mock_user_api_key_dict
|
||||
):
|
||||
"""Test behavior when LLM response has no anonymized content available"""
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "What's your email?"}],
|
||||
"litellm_call_id": "test-call-id",
|
||||
}
|
||||
|
||||
# Create LLM response with test data
|
||||
llm_response = ModelResponse(
|
||||
id="test-response-id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(
|
||||
content="My email is admin@company.com", role="assistant"
|
||||
),
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-3.5-turbo",
|
||||
object="chat.completion",
|
||||
system_fingerprint=None,
|
||||
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
||||
)
|
||||
|
||||
# Mock Noma API response with no anonymized content available
|
||||
noma_response = {
|
||||
"originalResponse": {
|
||||
"response": {
|
||||
"dataDetector": {
|
||||
"dataType1": {"result": True},
|
||||
},
|
||||
"contentDetector": {"result": False},
|
||||
},
|
||||
},
|
||||
"verdict": False,
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = noma_response
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
# Update guardrail to use post_call event hook
|
||||
anonymize_guardrail.event_hook = "post_call"
|
||||
|
||||
with patch.object(
|
||||
anonymize_guardrail.async_handler, "post", return_value=mock_response
|
||||
):
|
||||
# Should raise NomaBlockedMessage because no anonymized content available for LLM response
|
||||
with pytest.raises(NomaBlockedMessage):
|
||||
await anonymize_guardrail.async_post_call_success_hook(
|
||||
data=request_data,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
response=llm_response,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue