From 3c8e9410414e63dfc09538f382c3ede045a72a3f Mon Sep 17 00:00:00 2001 From: alihacks Date: Fri, 24 Apr 2026 16:09:33 -0400 Subject: [PATCH] azure_sentinel_truncate Azure Sentinel truncation changes to comply with limits --- .../azure_sentinel/azure_sentinel.py | 204 +++++++++-- .../integrations/test_azure_sentinel.py | 318 +++++++++++++++--- 2 files changed, 447 insertions(+), 75 deletions(-) diff --git a/litellm/integrations/azure_sentinel/azure_sentinel.py b/litellm/integrations/azure_sentinel/azure_sentinel.py index dd508e6c6c2..f70c40280a9 100644 --- a/litellm/integrations/azure_sentinel/azure_sentinel.py +++ b/litellm/integrations/azure_sentinel/azure_sentinel.py @@ -13,9 +13,11 @@ For batching specific details see CustomBatchLogger class """ import asyncio +import gzip import os import traceback -from typing import List, Optional +from copy import deepcopy +from typing import Any, Dict, List, Optional from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger @@ -119,6 +121,14 @@ class AzureSentinelLogger(CustomBatchLogger): asyncio.create_task(self.periodic_flush()) self.log_queue: List[StandardLoggingPayload] = [] + # When True, string fields (messages, response) are truncated to the + # Azure Log Analytics column limit (256 KB / 262,144 chars). Azure + # silently truncates at this limit anyway; doing it ourselves lets us + # keep the tail (most recent content) and record metadata. + # Controlled by AZURE_SENTINEL_TRUNCATE_CONTENT env var (default: false). + truncate_env = os.getenv("AZURE_SENTINEL_TRUNCATE_CONTENT", "false") + self.truncate_content = truncate_env.lower() in ("true", "1", "yes") + async def _get_oauth_token(self) -> str: """ Get OAuth2 Bearer token for Azure Monitor Logs Ingestion API @@ -189,9 +199,6 @@ class AzureSentinelLogger(CustomBatchLogger): Raises a NON Blocking verbose_logger.exception if an error occurs """ try: - verbose_logger.debug( - "Azure Sentinel: Logging - Enters logging function for model %s", kwargs - ) standard_logging_payload = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: @@ -223,10 +230,6 @@ class AzureSentinelLogger(CustomBatchLogger): Raises a NON Blocking verbose_logger.exception if an error occurs """ try: - verbose_logger.debug( - "Azure Sentinel: Logging - Enters failure logging function for model %s", - kwargs, - ) standard_logging_payload = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: @@ -246,9 +249,141 @@ class AzureSentinelLogger(CustomBatchLogger): ) pass + # Azure DCR Logs Ingestion API has a 1MB request size limit. + # We target a conservative threshold to stay safely under the limit. + MAX_BATCH_SIZE_BYTES = 950_000 # ~950KB uncompressed target per batch + + # Azure Log Analytics silently truncates string column values at 256 KB. + # We enforce this limit ourselves so we can keep the tail (most recent + # content) and record truncation metadata. + MAX_COLUMN_CHARS = 262_144 # 256 KB Azure Log Analytics column limit + + def _enforce_column_limits( + self, payload: StandardLoggingPayload + ) -> StandardLoggingPayload: + """ + Truncate messages/response string fields to the Azure Log Analytics + column limit (262,144 chars). Keeps the *tail* of each field so that + the most recent conversation turns and response text are preserved. + + Only called when ``self.truncate_content`` is True. + + Returns the original payload unchanged if neither field exceeds the + limit, otherwise returns a deep copy with truncated fields and + truncation metadata added. + """ + limit = self.MAX_COLUMN_CHARS + messages = payload.get("messages") + response = payload.get("response") + msg_str = str(messages) if messages is not None else "" + resp_str = str(response) if response is not None else "" + + needs_truncation = len(msg_str) > limit or len(resp_str) > limit + if not needs_truncation: + return payload + + entry = deepcopy(payload) + truncated_fields: List[str] = [] + + if len(msg_str) > limit: + entry["messages"] = "[truncated by litellm]..." + msg_str[-limit:] + truncated_fields.append("messages") + + if len(resp_str) > limit: + entry["response"] = "[truncated by litellm]..." + resp_str[-limit:] + truncated_fields.append("response") + + truncation_info: Dict[str, Any] = { + "truncated": True, + "truncate_reason": "azure_column_limit", + "truncated_fields": truncated_fields, + "original_messages_chars": len(msg_str), + "original_response_chars": len(resp_str), + "max_column_chars": limit, + } + + if "metadata" in entry and isinstance(entry.get("metadata"), dict): + entry["metadata"]["litellm_content_truncated"] = truncation_info # type: ignore + else: + entry["litellm_content_truncated"] = truncation_info # type: ignore + + verbose_logger.debug( + "Azure Sentinel: column-level truncation (id=%s). " + "messages: %d→%d chars, response: %d→%d chars", + entry.get("id", "?"), + len(msg_str), + min(len(msg_str), limit), + len(resp_str), + min(len(resp_str), limit), + ) + + return entry + + def _split_into_batches( + self, payloads: List[StandardLoggingPayload] + ) -> List[bytes]: + """ + Splits payloads into gzip-compressed batches that stay under + MAX_BATCH_SIZE_BYTES (uncompressed) per batch. + + When truncate_content is enabled, enforces the Azure Log Analytics + column limit (256 KB) on messages/response fields before batching. + + Returns a list of gzip-compressed byte strings, each representing + a JSON array of log entries. + """ + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + + batches: List[bytes] = [] + current_batch: list = [] + current_size = 2 # account for JSON array brackets '[]' + + for payload in payloads: + # Enforce Azure column-level string limit when enabled + if self.truncate_content: + payload = self._enforce_column_limits(payload) + + entry_json = safe_dumps(payload) + entry_size = len(entry_json.encode("utf-8")) + + # If a single entry exceeds the uncompressed batch limit, + # send it alone in its own batch + if entry_size + 2 > self.MAX_BATCH_SIZE_BYTES: + # Flush any accumulated batch first + if current_batch: + batch_body = safe_dumps(current_batch) + batches.append(gzip.compress(batch_body.encode("utf-8"))) + current_batch = [] + current_size = 2 + + single_body = safe_dumps([payload]) + batches.append(gzip.compress(single_body.encode("utf-8"))) + continue + + # +1 for comma separator between entries + separator = 1 if current_batch else 0 + if current_size + separator + entry_size > self.MAX_BATCH_SIZE_BYTES: + # Current batch is full — flush it + batch_body = safe_dumps(current_batch) + batches.append(gzip.compress(batch_body.encode("utf-8"))) + current_batch = [] + current_size = 2 + + current_batch.append(payload) + current_size += entry_size + (1 if len(current_batch) > 1 else 0) + + # Flush remaining + if current_batch: + batch_body = safe_dumps(current_batch) + batches.append(gzip.compress(batch_body.encode("utf-8"))) + + return batches + async def async_send_batch(self): """ Sends the batch of logs to Azure Monitor Logs Ingestion API + with gzip compression. Splits into multiple requests if the + batch exceeds ~1MB uncompressed. Raises: Raises a NON Blocking verbose_logger.exception if an error occurs @@ -261,40 +396,49 @@ class AzureSentinelLogger(CustomBatchLogger): "Azure Sentinel - about to flush %s events", len(self.log_queue) ) - from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - # Get OAuth2 token bearer_token = await self._get_oauth_token() - # Convert log queue to JSON array format expected by Logs Ingestion API - # Each log entry should be a JSON object in the array - body = safe_dumps(self.log_queue) + # Split into size-limited, gzip-compressed batches + compressed_batches = self._split_into_batches(self.log_queue) - # Set headers for Logs Ingestion API + verbose_logger.debug( + "Azure Sentinel - split into %s batch(es)", len(compressed_batches) + ) + + # Set headers for Logs Ingestion API with gzip encoding headers = { "Authorization": f"Bearer {bearer_token}", "Content-Type": "application/json", + "Content-Encoding": "gzip", } - # Send the request - response = await self.async_httpx_client.post( - url=self.api_endpoint, data=body.encode("utf-8"), headers=headers - ) + # Send each batch + for i, compressed_body in enumerate(compressed_batches): + response = await self.async_httpx_client.post( + url=self.api_endpoint, + content=compressed_body, + headers=headers, + ) - if response.status_code not in [200, 204]: - verbose_logger.error( - "Azure Sentinel API error: status_code=%s, response=%s", + if response.status_code not in [200, 204]: + verbose_logger.error( + "Azure Sentinel API error on batch %s/%s: status_code=%s, response=%s", + i + 1, + len(compressed_batches), + response.status_code, + response.text, + ) + raise Exception( + f"Failed to send logs to Azure Sentinel: {response.status_code} - {response.text}" + ) + + verbose_logger.debug( + "Azure Sentinel: batch %s/%s sent, status_code: %s", + i + 1, + len(compressed_batches), response.status_code, - response.text, ) - raise Exception( - f"Failed to send logs to Azure Sentinel: {response.status_code} - {response.text}" - ) - - verbose_logger.debug( - "Azure Sentinel: Response from API status_code: %s", - response.status_code, - ) except Exception as e: verbose_logger.exception( diff --git a/tests/test_litellm/integrations/test_azure_sentinel.py b/tests/test_litellm/integrations/test_azure_sentinel.py index 031b85211f7..de710c72fac 100644 --- a/tests/test_litellm/integrations/test_azure_sentinel.py +++ b/tests/test_litellm/integrations/test_azure_sentinel.py @@ -2,8 +2,11 @@ Test Azure Sentinel logging integration """ -import datetime -from unittest.mock import AsyncMock, patch +import gzip +import json +import os + +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -11,26 +14,22 @@ from litellm.integrations.azure_sentinel.azure_sentinel import AzureSentinelLogg from litellm.types.utils import StandardLoggingPayload -@pytest.mark.asyncio -async def test_azure_sentinel_oauth_and_send_batch(): - """Test that Azure Sentinel logger gets OAuth token and sends batch to API""" - test_dcr_id = "dcr-test123456789" - test_endpoint = "https://test-dce.eastus-1.ingest.monitor.azure.com" - test_tenant_id = "test-tenant-id" - test_client_id = "test-client-id" - test_client_secret = "test-client-secret" +def _make_logger(**overrides): + """Helper to create an AzureSentinelLogger with mocked asyncio.create_task""" + defaults = dict( + dcr_immutable_id="dcr-test123456789", + endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com", + tenant_id="test-tenant-id", + client_id="test-client-id", + client_secret="test-client-secret", + ) + defaults.update(overrides) + with patch("asyncio.create_task", side_effect=lambda coro: coro.close()): + return AzureSentinelLogger(**defaults) - with patch("asyncio.create_task"): - logger = AzureSentinelLogger( - dcr_immutable_id=test_dcr_id, - endpoint=test_endpoint, - tenant_id=test_tenant_id, - client_id=test_client_id, - client_secret=test_client_secret, - ) - # Create test payload - standard_payload = StandardLoggingPayload( +def _make_payload(**overrides): + defaults = dict( id="test_id", call_type="completion", model="gpt-3.5-turbo", @@ -38,56 +37,285 @@ async def test_azure_sentinel_oauth_and_send_batch(): messages=[{"role": "user", "content": "Hello"}], response={"choices": [{"message": {"content": "Hi"}}]}, ) + defaults.update(overrides) + return StandardLoggingPayload(**defaults) - # Add to queue - logger.log_queue.append(standard_payload) - - # Mock OAuth token response - from unittest.mock import MagicMock +def _mock_http_client(logger): + """Wire up mocked token + API responses and return the mock.""" mock_token_response = MagicMock() mock_token_response.status_code = 200 mock_token_response.json = MagicMock( - return_value={ - "access_token": "test-bearer-token", - "expires_in": 3600, - } + return_value={"access_token": "test-bearer-token", "expires_in": 3600} ) mock_token_response.text = "Success" - # Mock API response mock_api_response = MagicMock() mock_api_response.status_code = 204 mock_api_response.text = "Success" - # Mock HTTP client - first call for token, second for API async def mock_post(*args, **kwargs): if "oauth2/v2.0/token" in kwargs.get("url", ""): return mock_token_response return mock_api_response logger.async_httpx_client.post = AsyncMock(side_effect=mock_post) + return logger.async_httpx_client.post + + +@pytest.mark.asyncio +async def test_azure_sentinel_oauth_and_send_batch(): + """Test that Azure Sentinel logger gets OAuth token and sends batch to API with gzip""" + logger = _make_logger() + logger.log_queue.append(_make_payload()) + mock_post = _mock_http_client(logger) - # Send batch await logger.async_send_batch() - # Verify OAuth token request was made - assert logger.async_httpx_client.post.called - - # Verify API request was made - call_count = logger.async_httpx_client.post.call_count - assert call_count >= 2 # At least token + API call + # Verify OAuth token + at least one API call + assert mock_post.called + assert mock_post.call_count >= 2 # Get the API call (last call) - api_call_args = logger.async_httpx_client.post.call_args_list[-1] - assert test_dcr_id in api_call_args.kwargs["url"] - assert test_endpoint in api_call_args.kwargs["url"] + api_call = mock_post.call_args_list[-1] + assert "dcr-test123456789" in api_call.kwargs["url"] - # Verify headers - headers = api_call_args.kwargs["headers"] + # Verify gzip Content-Encoding header + headers = api_call.kwargs["headers"] assert headers["Content-Type"] == "application/json" - assert "Authorization" in headers + assert headers["Content-Encoding"] == "gzip" assert headers["Authorization"].startswith("Bearer ") - # Verify queue is cleared + # Verify body is valid gzip containing JSON array + compressed_body = api_call.kwargs["content"] + decompressed = gzip.decompress(compressed_body) + parsed = json.loads(decompressed) + assert isinstance(parsed, list) + assert len(parsed) == 1 + + # Queue should be cleared assert len(logger.log_queue) == 0 + + +@pytest.mark.asyncio +async def test_azure_sentinel_batch_splitting(): + """Test that large batches are split into multiple requests under 1MB""" + logger = _make_logger() + + # Create payloads with large content to force splitting. + # Each payload will be ~100KB so 15 of them (~1.5MB) should trigger a split. + large_content = "x" * 100_000 + for i in range(15): + logger.log_queue.append( + _make_payload( + id=f"test_{i}", + messages=[{"role": "user", "content": large_content}], + ) + ) + + mock_post = _mock_http_client(logger) + await logger.async_send_batch() + + # Should have token call + multiple API calls (more than 1 batch) + api_calls = [ + c for c in mock_post.call_args_list if "oauth2/v2.0/token" not in c.kwargs.get("url", "") + ] + assert len(api_calls) >= 2, f"Expected multiple batches, got {len(api_calls)}" + + # Each batch body should decompress to a valid JSON array + total_events = 0 + for call in api_calls: + compressed = call.kwargs["content"] + decompressed = gzip.decompress(compressed) + parsed = json.loads(decompressed) + assert isinstance(parsed, list) + total_events += len(parsed) + + # All 15 events accounted for + assert total_events == 15 + assert len(logger.log_queue) == 0 + + +@pytest.mark.asyncio +async def test_azure_sentinel_split_into_batches_single_oversized_entry(): + """Test that a single entry larger than MAX_BATCH_SIZE_BYTES is sent alone""" + with patch.dict(os.environ, {"AZURE_SENTINEL_TRUNCATE_CONTENT": "false"}): + logger = _make_logger() + + # One very large payload that exceeds the batch size on its own. + # Use varied content so JSON serialization keeps it large. + import hashlib + chunks = [hashlib.sha256(str(i).encode()).hexdigest() for i in range(20_000)] + huge_content = " ".join(chunks) # ~1.3MB of hex digests + batches = logger._split_into_batches( + [ + _make_payload( + id="small_1", + messages=[{"role": "user", "content": "hi"}], + ), + _make_payload( + id="huge", + messages=[{"role": "user", "content": huge_content}], + ), + _make_payload( + id="small_2", + messages=[{"role": "user", "content": "hello"}], + ), + ] + ) + + # Should produce at least 2 batches (the oversized one isolated, plus the smalls) + assert len(batches) >= 2 + + # Verify all entries present across batches + total_ids = [] + for compressed in batches: + parsed = json.loads(gzip.decompress(compressed)) + for entry in parsed: + total_ids.append(entry.get("id")) + assert set(total_ids) == {"small_1", "huge", "small_2"} + + +def test_column_limit_truncates_large_fields(): + """Test that fields exceeding 262,144 chars are truncated (keeping the tail)""" + with patch.dict(os.environ, {"AZURE_SENTINEL_TRUNCATE_CONTENT": "true"}): + logger = _make_logger() + + # Content larger than 256 KB column limit + big_messages = "A" * 300_000 + big_response = "B" * 300_000 + + payload = _make_payload( + id="big_entry", + messages=[{"role": "user", "content": big_messages}], + response=big_response, + ) + + result = logger._enforce_column_limits(payload) + + # Should be a new object (deep copy) + assert result is not payload + + # Messages field should be truncated, keeping tail + msg_str = str(result["messages"]) + assert msg_str.startswith("[truncated by litellm]...") + assert len(msg_str) <= logger.MAX_COLUMN_CHARS + len("[truncated by litellm]...") + + # Response field should be truncated, keeping tail + resp_str = str(result["response"]) + assert resp_str.startswith("[truncated by litellm]...") + + # Truncation metadata present + metadata = result.get("metadata", {}) + trunc_info = metadata.get("litellm_content_truncated") if isinstance(metadata, dict) else None + if trunc_info is None: + trunc_info = result.get("litellm_content_truncated") + assert trunc_info is not None + assert trunc_info["truncated"] is True + assert trunc_info["truncate_reason"] == "azure_column_limit" + assert "messages" in trunc_info["truncated_fields"] + assert "response" in trunc_info["truncated_fields"] + assert trunc_info["original_messages_chars"] == len(str(payload["messages"])) + assert trunc_info["max_column_chars"] == 262_144 + + # Original payload not mutated + assert len(str(payload["messages"])) > 262_144 + + +def test_column_limit_preserves_small_payloads(): + """Test that payloads under the column limit are returned unchanged""" + with patch.dict(os.environ, {"AZURE_SENTINEL_TRUNCATE_CONTENT": "true"}): + logger = _make_logger() + + payload = _make_payload( + id="small_entry", + messages=[{"role": "user", "content": "hello world"}], + ) + + result = logger._enforce_column_limits(payload) + + # Should be the exact same object (no copy needed) + assert result is payload + metadata = result.get("metadata", {}) + if isinstance(metadata, dict): + assert "litellm_content_truncated" not in metadata + + +def test_truncate_disabled_via_env_var(): + """Test that truncation is skipped when AZURE_SENTINEL_TRUNCATE_CONTENT=false""" + with patch.dict(os.environ, {"AZURE_SENTINEL_TRUNCATE_CONTENT": "false"}): + logger = _make_logger() + assert logger.truncate_content is False + + # Create payload with content exceeding 256 KB column limit + huge_content = "z" * 300_000 + payloads = [ + _make_payload( + id="huge_no_truncate", + messages=[{"role": "user", "content": huge_content}], + ), + ] + + batches = logger._split_into_batches(payloads) + assert len(batches) >= 1 + # Collect all entries + all_entries = [] + for compressed in batches: + all_entries.extend(json.loads(gzip.decompress(compressed))) + entry = all_entries[0] + # The messages should be the full original content (not truncated) + assert "[truncated by litellm]" not in str(entry.get("messages", "")) + + +def test_truncate_enabled_in_split_batches(): + """Test that _split_into_batches truncates large fields when enabled""" + with patch.dict(os.environ, {"AZURE_SENTINEL_TRUNCATE_CONTENT": "true"}): + logger = _make_logger() + assert logger.truncate_content is True + + # Content exceeding 256 KB column limit + huge_content = "X" * 400_000 + payloads = [ + _make_payload( + id="small_before", + messages=[{"role": "user", "content": "hi"}], + ), + _make_payload( + id="huge_truncated", + messages=[{"role": "user", "content": huge_content}], + ), + _make_payload( + id="small_after", + messages=[{"role": "user", "content": "bye"}], + ), + ] + + batches = logger._split_into_batches(payloads) + + # Collect all entries + all_entries = {} + for compressed in batches: + for entry in json.loads(gzip.decompress(compressed)): + all_entries[entry["id"]] = entry + + assert set(all_entries.keys()) == {"small_before", "huge_truncated", "small_after"} + + # The huge entry should have truncation metadata + huge_entry = all_entries["huge_truncated"] + metadata = huge_entry.get("metadata", {}) + trunc_info = metadata.get("litellm_content_truncated") if isinstance(metadata, dict) else None + if trunc_info is None: + trunc_info = huge_entry.get("litellm_content_truncated") + assert trunc_info is not None + assert trunc_info["truncated"] is True + assert trunc_info["truncate_reason"] == "azure_column_limit" + # Messages field should be capped near 262,144 chars + msg_str = str(huge_entry["messages"]) + assert msg_str.startswith("[truncated by litellm]...") + + # Small entries should have no truncation metadata + for entry_id in ("small_before", "small_after"): + entry = all_entries[entry_id] + meta = entry.get("metadata", {}) + if isinstance(meta, dict): + assert "litellm_content_truncated" not in meta