mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
[Fix] SQS Logger - Add Base64 handling (#16028)
* Enable base64 stripping from sqs (#15927) * Add sqs logger * Add sqs logger * Add sqs strp base64 * Add sqs strp base64 * Add sqs strp base64 * strip base64 * Add sqs strp base64 * strip base64 * Add sqs strp base64 * Add max depth recursion * Add max depth recursion --------- Co-authored-by: deepanshu <deepanshu.lulla@hq.bill.com> * refactor _strip_base64_from_messages * test fixes SQS logger * fix SQS linting --------- Co-authored-by: Deepanshu Lulla <deepanshu.lulla@gmail.com> Co-authored-by: deepanshu <deepanshu.lulla@hq.bill.com>
This commit is contained in:
parent
95dd216150
commit
ab8a3a5d9e
5 changed files with 319 additions and 58 deletions
|
|
@ -1415,18 +1415,26 @@ AWS_REGION_NAME = ""
|
|||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: gpt-4o
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["aws_sqs"]
|
||||
|
||||
aws_sqs_callback_params:
|
||||
sqs_queue_url: https://sqs.us-west-2.amazonaws.com/123456789012/my-queue # AWS SQS Queue URL
|
||||
sqs_region_name: us-west-2 # AWS Region Name for SQS
|
||||
sqs_aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID # use os.environ/<variable name> to pass environment variables. This is AWS Access Key ID for SQS
|
||||
sqs_aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY # AWS Secret Access Key for SQS
|
||||
sqs_batch_size: 10 # [OPTIONAL] Number of messages to batch before sending (default: 10)
|
||||
sqs_flush_interval: 30 # [OPTIONAL] Time in seconds to wait before flushing batch (default: 30)
|
||||
# --- 🧱 Required Parameters ---
|
||||
sqs_queue_url: https://sqs.us-west-2.amazonaws.com/123456789012/my-queue
|
||||
# The AWS SQS Queue URL to which LiteLLM will send log events.
|
||||
|
||||
sqs_region_name: us-west-2
|
||||
# AWS Region for your SQS queue (e.g., us-east-1, eu-central-1, etc.)
|
||||
|
||||
# --- Logging Controls ---
|
||||
sqs_strip_base64_files: true
|
||||
# If true, LiteLLM will remove or redact base64-encoded binary data (e.g., PDFs, images, audio)
|
||||
# from logged messages to avoid large payloads. SQS has a 1 MB payload size limit.
|
||||
|
||||
```
|
||||
|
||||
**Step 3**: Start the proxy, make a test request
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
#### What this does ####
|
||||
# On success, logs events to Promptlayer
|
||||
import re
|
||||
import traceback
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
|
|
@ -15,7 +16,9 @@ from typing import (
|
|||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
|
||||
from litellm.types.integrations.argilla import ArgillaItem
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionRequest
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -53,6 +56,12 @@ else:
|
|||
PreRoutingHookResponse = Any
|
||||
|
||||
|
||||
_BASE64_INLINE_PATTERN = re.compile(
|
||||
r"data:(?:application|image|audio|video)/[a-zA-Z0-9.+-]+;base64,[A-Za-z0-9+/=\s]+",
|
||||
re.MULTILINE,
|
||||
)
|
||||
|
||||
|
||||
class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callback#callback-class
|
||||
# Class variables or attributes
|
||||
def __init__(
|
||||
|
|
@ -567,3 +576,91 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
Get the proxy server request from cold storage using the object key directly.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
async def _strip_base64_from_messages(
|
||||
self, payload: "StandardLoggingPayload", max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
|
||||
) -> "StandardLoggingPayload":
|
||||
"""
|
||||
Removes or redacts base64-encoded file data (e.g., PDFs, images, audio)
|
||||
from messages and responses before sending to SQS.
|
||||
|
||||
Behavior:
|
||||
• Drop entries with a 'file' key.
|
||||
• Drop entries with type == 'file' or any non-text type.
|
||||
• Keep untyped or text content.
|
||||
• Recursively redact inline base64 blobs in *any* string field, at any depth.
|
||||
"""
|
||||
raw_messages: Any = payload.get("messages", [])
|
||||
messages: list[Any] = raw_messages if isinstance(raw_messages, list) else []
|
||||
verbose_logger.debug(f"[CustomLogger] Stripping base64 from {len(messages)} messages")
|
||||
|
||||
if messages:
|
||||
payload["messages"] = self._process_messages(messages=messages, max_depth=max_depth)
|
||||
|
||||
total_items = 0
|
||||
for m in payload.get("messages", []) or []:
|
||||
if isinstance(m, dict):
|
||||
content = m.get("content", [])
|
||||
if isinstance(content, list):
|
||||
total_items += len(content)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"[CustomLogger] Completed base64 strip; retained {total_items} content items"
|
||||
)
|
||||
return payload
|
||||
|
||||
|
||||
def _redact_base64(self, value: Any, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER) -> Any:
|
||||
"""Recursively redact inline base64 from any nested structure with a max recursion depth limit."""
|
||||
if depth > max_depth:
|
||||
verbose_logger.warning(
|
||||
f"[CustomLogger] Max recursion depth {max_depth} reached while redacting base64"
|
||||
)
|
||||
return "[MAX_DEPTH_REACHED]"
|
||||
|
||||
if isinstance(value, str):
|
||||
if _BASE64_INLINE_PATTERN.search(value):
|
||||
verbose_logger.debug(
|
||||
f"[CustomLogger] Redacted inline base64 string: {value[:40]}..."
|
||||
)
|
||||
return _BASE64_INLINE_PATTERN.sub("[BASE64_REDACTED]", value)
|
||||
return value
|
||||
|
||||
if isinstance(value, list):
|
||||
return [self._redact_base64(value=v, depth=depth + 1, max_depth=max_depth) for v in value]
|
||||
|
||||
if isinstance(value, dict):
|
||||
return {k: self._redact_base64(value=v, depth=depth + 1, max_depth=max_depth) for k, v in value.items()}
|
||||
|
||||
return value
|
||||
|
||||
def _should_keep_content(self, content: Any) -> bool:
|
||||
"""Return True if this content item should be retained."""
|
||||
if not isinstance(content, dict):
|
||||
return True
|
||||
if "file" in content:
|
||||
return False
|
||||
ctype = content.get("type")
|
||||
return not (isinstance(ctype, str) and ctype != "text")
|
||||
|
||||
def _process_messages(self, messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER) -> list[dict[str, Any]]:
|
||||
filtered_messages: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
contents: Any = msg.get("content")
|
||||
if isinstance(contents, list):
|
||||
cleaned: list[Any] = []
|
||||
for c in contents:
|
||||
if self._should_keep_content(content=c):
|
||||
cleaned.append(self._redact_base64(value=c, max_depth=max_depth))
|
||||
msg["content"] = cleaned
|
||||
else:
|
||||
msg["content"] = self._redact_base64(value=contents, max_depth=max_depth)
|
||||
|
||||
for key, val in list(msg.items()):
|
||||
if key != "content":
|
||||
msg[key] = self._redact_base64(value=val, max_depth=max_depth)
|
||||
filtered_messages.append(msg)
|
||||
return filtered_messages
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from __future__ import annotations
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import re
|
||||
import traceback
|
||||
from typing import List, Optional
|
||||
|
||||
|
|
@ -30,6 +31,11 @@ from litellm.types.utils import StandardLoggingPayload
|
|||
|
||||
from .custom_batch_logger import CustomBatchLogger
|
||||
|
||||
_BASE64_INLINE_PATTERN = re.compile(
|
||||
r"data:(?:application|image|audio|video)/[a-zA-Z0-9.+-]+;base64,[A-Za-z0-9+/=\s]+",
|
||||
re.MULTILINE,
|
||||
)
|
||||
|
||||
|
||||
class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
||||
"""Batching logger that writes logs to an AWS SQS queue, optionally encrypting the payload."""
|
||||
|
|
@ -54,6 +60,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
|||
sqs_flush_interval: Optional[int] = DEFAULT_SQS_FLUSH_INTERVAL_SECONDS,
|
||||
sqs_batch_size: Optional[int] = DEFAULT_SQS_BATCH_SIZE,
|
||||
sqs_config=None,
|
||||
sqs_strip_base64_files: bool = False,
|
||||
# --- 🔐 Application-level encryption params ---
|
||||
sqs_aws_use_application_level_encryption: bool = False,
|
||||
sqs_app_encryption_key_b64: Optional[str] = None,
|
||||
|
|
@ -84,6 +91,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
|||
sqs_aws_role_name=sqs_aws_role_name,
|
||||
sqs_aws_web_identity_token=sqs_aws_web_identity_token,
|
||||
sqs_aws_sts_endpoint=sqs_aws_sts_endpoint,
|
||||
sqs_strip_base64_files=sqs_strip_base64_files,
|
||||
sqs_aws_use_application_level_encryption=sqs_aws_use_application_level_encryption,
|
||||
sqs_app_encryption_key_b64=sqs_app_encryption_key_b64,
|
||||
sqs_app_encryption_aad=sqs_app_encryption_aad,
|
||||
|
|
@ -113,25 +121,26 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
|||
raise e
|
||||
|
||||
def _init_sqs_params(
|
||||
self,
|
||||
sqs_queue_url: Optional[str] = None,
|
||||
sqs_region_name: Optional[str] = None,
|
||||
sqs_api_version: Optional[str] = None,
|
||||
sqs_use_ssl: bool = True,
|
||||
sqs_verify: Optional[bool] = None,
|
||||
sqs_endpoint_url: Optional[str] = None,
|
||||
sqs_aws_access_key_id: Optional[str] = None,
|
||||
sqs_aws_secret_access_key: Optional[str] = None,
|
||||
sqs_aws_session_token: Optional[str] = None,
|
||||
sqs_aws_session_name: Optional[str] = None,
|
||||
sqs_aws_profile_name: Optional[str] = None,
|
||||
sqs_aws_role_name: Optional[str] = None,
|
||||
sqs_aws_web_identity_token: Optional[str] = None,
|
||||
sqs_aws_sts_endpoint: Optional[str] = None,
|
||||
sqs_aws_use_application_level_encryption: bool = False,
|
||||
sqs_app_encryption_key_b64: Optional[str] = None,
|
||||
sqs_app_encryption_aad: Optional[str] = None,
|
||||
sqs_config=None,
|
||||
self,
|
||||
sqs_queue_url: Optional[str] = None,
|
||||
sqs_region_name: Optional[str] = None,
|
||||
sqs_api_version: Optional[str] = None,
|
||||
sqs_use_ssl: bool = True,
|
||||
sqs_verify: Optional[bool] = None,
|
||||
sqs_endpoint_url: Optional[str] = None,
|
||||
sqs_aws_access_key_id: Optional[str] = None,
|
||||
sqs_aws_secret_access_key: Optional[str] = None,
|
||||
sqs_aws_session_token: Optional[str] = None,
|
||||
sqs_aws_session_name: Optional[str] = None,
|
||||
sqs_aws_profile_name: Optional[str] = None,
|
||||
sqs_aws_role_name: Optional[str] = None,
|
||||
sqs_aws_web_identity_token: Optional[str] = None,
|
||||
sqs_aws_sts_endpoint: Optional[str] = None,
|
||||
sqs_strip_base64_files: bool = False,
|
||||
sqs_aws_use_application_level_encryption: bool = False,
|
||||
sqs_app_encryption_key_b64: Optional[str] = None,
|
||||
sqs_app_encryption_aad: Optional[str] = None,
|
||||
sqs_config=None,
|
||||
) -> None:
|
||||
litellm.aws_sqs_callback_params = litellm.aws_sqs_callback_params or {}
|
||||
|
||||
|
|
@ -141,55 +150,59 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
|||
litellm.aws_sqs_callback_params[key] = litellm.get_secret(value)
|
||||
|
||||
self.sqs_queue_url = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_queue_url") or sqs_queue_url
|
||||
litellm.aws_sqs_callback_params.get("sqs_queue_url") or sqs_queue_url
|
||||
)
|
||||
self.sqs_region_name = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_region_name") or sqs_region_name
|
||||
litellm.aws_sqs_callback_params.get("sqs_region_name") or sqs_region_name
|
||||
)
|
||||
self.sqs_api_version = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_api_version") or sqs_api_version
|
||||
litellm.aws_sqs_callback_params.get("sqs_api_version") or sqs_api_version
|
||||
)
|
||||
self.sqs_use_ssl = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_use_ssl", True) or sqs_use_ssl
|
||||
litellm.aws_sqs_callback_params.get("sqs_use_ssl", True) or sqs_use_ssl
|
||||
)
|
||||
self.sqs_verify = litellm.aws_sqs_callback_params.get("sqs_verify") or sqs_verify
|
||||
self.sqs_endpoint_url = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_endpoint_url") or sqs_endpoint_url
|
||||
litellm.aws_sqs_callback_params.get("sqs_endpoint_url") or sqs_endpoint_url
|
||||
)
|
||||
self.sqs_aws_access_key_id = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_access_key_id")
|
||||
or sqs_aws_access_key_id
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_access_key_id")
|
||||
or sqs_aws_access_key_id
|
||||
)
|
||||
|
||||
self.sqs_aws_secret_access_key = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_secret_access_key")
|
||||
or sqs_aws_secret_access_key
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_secret_access_key")
|
||||
or sqs_aws_secret_access_key
|
||||
)
|
||||
|
||||
self.sqs_aws_session_token = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_session_token")
|
||||
or sqs_aws_session_token
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_session_token")
|
||||
or sqs_aws_session_token
|
||||
)
|
||||
|
||||
self.sqs_aws_session_name = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_session_name") or sqs_aws_session_name
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_session_name") or sqs_aws_session_name
|
||||
)
|
||||
|
||||
self.sqs_aws_profile_name = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_profile_name") or sqs_aws_profile_name
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_profile_name") or sqs_aws_profile_name
|
||||
)
|
||||
|
||||
self.sqs_aws_role_name = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_role_name") or sqs_aws_role_name
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_role_name") or sqs_aws_role_name
|
||||
)
|
||||
|
||||
self.sqs_aws_web_identity_token = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_web_identity_token")
|
||||
or sqs_aws_web_identity_token
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_web_identity_token")
|
||||
or sqs_aws_web_identity_token
|
||||
)
|
||||
|
||||
self.sqs_aws_sts_endpoint = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_sts_endpoint") or sqs_aws_sts_endpoint
|
||||
litellm.aws_sqs_callback_params.get("sqs_aws_sts_endpoint") or sqs_aws_sts_endpoint
|
||||
)
|
||||
self.sqs_strip_base64_files = (
|
||||
litellm.aws_sqs_callback_params.get("sqs_strip_base64_files", False)
|
||||
or sqs_strip_base64_files
|
||||
)
|
||||
|
||||
self.sqs_aws_use_application_level_encryption = (
|
||||
|
|
@ -217,13 +230,15 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
|||
self.sqs_config = litellm.aws_sqs_callback_params.get("sqs_config") or sqs_config
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs, response_obj, start_time, end_time
|
||||
self, kwargs, response_obj, start_time, end_time
|
||||
) -> None:
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
"SQS Logging - Enters logging function for model %s", kwargs
|
||||
)
|
||||
standard_logging_payload = kwargs.get("standard_logging_object")
|
||||
if self.sqs_strip_base64_files:
|
||||
standard_logging_payload = await self._strip_base64_from_messages(standard_logging_payload)
|
||||
if standard_logging_payload is None:
|
||||
raise ValueError("standard_logging_payload is None")
|
||||
|
||||
|
|
@ -337,4 +352,3 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
|
|||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error sending to SQS: {str(e)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -1,16 +1,18 @@
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
|
||||
class AppCrypto:
|
||||
def __init__(self, master_key: bytes):
|
||||
if len(master_key) != 32:
|
||||
raise ValueError("Master key must be 32 bytes for AES-256-GCM")
|
||||
self.key = master_key
|
||||
|
||||
def encrypt_json(self, data: dict, aad: bytes | None = None) -> dict:
|
||||
def encrypt_json(self, data: dict, aad: Optional[bytes] = None) -> dict:
|
||||
aes = AESGCM(self.key)
|
||||
nonce = os.urandom(12)
|
||||
plaintext = json.dumps(data).encode("utf-8")
|
||||
|
|
@ -22,7 +24,7 @@ class AppCrypto:
|
|||
"tag": base64.b64encode(tag).decode(),
|
||||
}
|
||||
|
||||
def decrypt_json(self, enc: dict, aad: bytes | None = None) -> dict:
|
||||
def decrypt_json(self, enc: dict, aad: Optional[bytes] = None) -> dict:
|
||||
aes = AESGCM(self.key)
|
||||
nonce = base64.b64decode(enc["nonce"])
|
||||
ct = base64.b64decode(enc["ciphertext"])
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
from copy import deepcopy
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from urllib.parse import unquote
|
||||
|
||||
|
|
@ -18,18 +19,18 @@ from litellm.litellm_core_utils.app_crypto import AppCrypto
|
|||
async def test_async_sqs_logger_flush():
|
||||
expected_queue_url = "https://sqs.us-east-1.amazonaws.com/123456789012/test-queue"
|
||||
expected_region = "us-east-1"
|
||||
|
||||
|
||||
sqs_logger = SQSLogger(
|
||||
sqs_queue_url=expected_queue_url,
|
||||
sqs_region_name=expected_region,
|
||||
sqs_flush_interval=1,
|
||||
)
|
||||
|
||||
|
||||
# Mock the httpx client
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
sqs_logger.async_httpx_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
|
||||
litellm.callbacks = [sqs_logger]
|
||||
|
||||
await litellm.acompletion(
|
||||
|
|
@ -42,31 +43,31 @@ async def test_async_sqs_logger_flush():
|
|||
|
||||
# Verify that httpx post was called
|
||||
sqs_logger.async_httpx_client.post.assert_called()
|
||||
|
||||
|
||||
# Get the call arguments
|
||||
call_args = sqs_logger.async_httpx_client.post.call_args
|
||||
|
||||
|
||||
# Verify the URL is correct
|
||||
called_url = call_args[0][0] # First positional argument
|
||||
assert called_url == expected_queue_url, f"Expected URL {expected_queue_url}, got {called_url}"
|
||||
|
||||
|
||||
# Verify the payload contains StandardLoggingPayload data
|
||||
called_data = call_args.kwargs['data']
|
||||
|
||||
|
||||
# Extract the MessageBody from the URL-encoded data
|
||||
# Format: "Action=SendMessage&Version=2012-11-05&MessageBody=<url_encoded_json>"
|
||||
assert "Action=SendMessage" in called_data
|
||||
assert "Version=2012-11-05" in called_data
|
||||
assert "MessageBody=" in called_data
|
||||
|
||||
|
||||
# Extract and decode the message body
|
||||
message_body_start = called_data.find("MessageBody=") + len("MessageBody=")
|
||||
message_body_encoded = called_data[message_body_start:]
|
||||
message_body_json = unquote(message_body_encoded)
|
||||
|
||||
|
||||
# Parse the JSON to verify it's a StandardLoggingPayload
|
||||
payload_data = json.loads(message_body_json)
|
||||
|
||||
|
||||
# Verify it has the expected StandardLoggingPayload structure
|
||||
assert "model" in payload_data
|
||||
assert "messages" in payload_data
|
||||
|
|
@ -291,4 +292,143 @@ async def test_async_send_batch_triggers_tasks(monkeypatch):
|
|||
|
||||
await logger.async_send_batch()
|
||||
# It uses asyncio.create_task() so direct await count = 0 is expected
|
||||
asyncio.create_task.assert_called()
|
||||
asyncio.create_task.assert_called()
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strip_base64_removes_file_and_nontext_entries():
|
||||
logger = SQSLogger(sqs_strip_base64_files=True)
|
||||
|
||||
payload = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Hello world"},
|
||||
{"type": "image", "file": {"file_data": "data:image/png;base64,AAAA"}},
|
||||
{"type": "file", "file": {"file_data": "data:application/pdf;base64,BBBB"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "Response"},
|
||||
{"type": "audio", "file": {"file_data": "data:audio/wav;base64,CCCC"}},
|
||||
],
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
stripped = await logger._strip_base64_from_messages(payload)
|
||||
|
||||
# 1️⃣ All file/image/audio entries removed
|
||||
assert len(stripped["messages"][0]["content"]) == 1
|
||||
assert stripped["messages"][0]["content"][0]["text"] == "Hello world"
|
||||
|
||||
assert len(stripped["messages"][1]["content"]) == 1
|
||||
assert stripped["messages"][1]["content"][0]["text"] == "Response"
|
||||
|
||||
# 2️⃣ No residual 'file' keys left
|
||||
for msg in stripped["messages"]:
|
||||
for content in msg["content"]:
|
||||
assert "file" not in content
|
||||
assert content.get("type") == "text"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strip_base64_keeps_non_file_content():
|
||||
logger = SQSLogger(sqs_strip_base64_files=True)
|
||||
|
||||
payload = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Just text"},
|
||||
{"type": "text", "text": "Another message"},
|
||||
],
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
stripped = await logger._strip_base64_from_messages(payload)
|
||||
|
||||
# Should not modify normal text messages
|
||||
assert stripped["messages"][0]["content"] == payload["messages"][0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strip_base64_handles_empty_or_missing_messages():
|
||||
logger = SQSLogger(sqs_strip_base64_files=True)
|
||||
|
||||
payload_no_messages = {}
|
||||
stripped1 = await logger._strip_base64_from_messages(payload_no_messages)
|
||||
assert stripped1 == payload_no_messages
|
||||
|
||||
payload_empty = {"messages": []}
|
||||
stripped2 = await logger._strip_base64_from_messages(payload_empty)
|
||||
assert stripped2 == payload_empty
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strip_base64_mixed_nested_objects():
|
||||
"""
|
||||
Handles weird/nested content structures gracefully.
|
||||
"""
|
||||
logger = SQSLogger(sqs_strip_base64_files=True)
|
||||
|
||||
payload = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "Keep me"},
|
||||
{"type": "custom", "metadata": "ignore but non-text"},
|
||||
{"foo": "bar"},
|
||||
{"file": {"file_data": "data:application/pdf;base64,XXX"}},
|
||||
],
|
||||
"extra": {"trace_id": "123"},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
stripped = await logger._strip_base64_from_messages(payload)
|
||||
|
||||
# 'custom' (non-text) and 'file' entries removed
|
||||
content = stripped["messages"][0]["content"]
|
||||
assert len(content) == 2
|
||||
assert {"type": "text", "text": "Keep me"} in content
|
||||
assert {"foo": "bar"} in content
|
||||
# Other metadata stays
|
||||
assert stripped["messages"][0]["extra"]["trace_id"] == "123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strip_base64_recursive_redaction():
|
||||
logger = SQSLogger(sqs_strip_base64_files=True)
|
||||
payload = {
|
||||
"messages": [
|
||||
{
|
||||
"content": [
|
||||
{"type": "text", "text": "normal text"},
|
||||
{"type": "text", "text": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUg"},
|
||||
{"type": "text", "text": "Nested: {'data': 'data:application/pdf;base64,AAA...'}"},
|
||||
{"file": {"file_data": "data:application/pdf;base64,AAAA"}},
|
||||
{"metadata": {"preview": "data:audio/mp3;base64,AAAAA=="}},
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
result = await logger._strip_base64_from_messages(payload)
|
||||
content = result["messages"][0]["content"]
|
||||
|
||||
# Dropped file-type entry
|
||||
assert not any("file" in c for c in content)
|
||||
# Base64 redacted globally
|
||||
for c in content:
|
||||
if isinstance(c, dict):
|
||||
s = json.dumps(c).lower()
|
||||
# allow "[base64_redacted]" but nothing else
|
||||
assert "base64," not in s, f"Found real base64 blob in: {s}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue