diff --git a/docs/my-website/docs/proxy/logging.md b/docs/my-website/docs/proxy/logging.md index 497e6e95a52..baf3d42234d 100644 --- a/docs/my-website/docs/proxy/logging.md +++ b/docs/my-website/docs/proxy/logging.md @@ -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/ 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 diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index d631c466865..88ad74aa5af 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -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 diff --git a/litellm/integrations/sqs.py b/litellm/integrations/sqs.py index ad4df97f98e..545aebbec6d 100644 --- a/litellm/integrations/sqs.py +++ b/litellm/integrations/sqs.py @@ -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)}") - diff --git a/litellm/litellm_core_utils/app_crypto.py b/litellm/litellm_core_utils/app_crypto.py index bb5043b2b7b..5ce6d8d77f9 100644 --- a/litellm/litellm_core_utils/app_crypto.py +++ b/litellm/litellm_core_utils/app_crypto.py @@ -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"]) diff --git a/tests/logging_callback_tests/test_sqs_logger.py b/tests/logging_callback_tests/test_sqs_logger.py index 3f9be7c3b30..ed963342285 100644 --- a/tests/logging_callback_tests/test_sqs_logger.py +++ b/tests/logging_callback_tests/test_sqs_logger.py @@ -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=" 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() \ No newline at end of file + 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}"