[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:
Ishaan Jaff 2025-10-28 16:41:32 -07:00 • committed by GitHub
parent 95dd216150
commit ab8a3a5d9e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 319 additions and 58 deletions

View file

@ -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

View file

@ -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

View file

@ -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)}")

View file

@ -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"])

View file

@ -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}"