mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
azure_sentinel_truncate Azure Sentinel truncation changes to comply with limits
This commit is contained in:
parent
98a9005c76
commit
3c8e941041
2 changed files with 447 additions and 75 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue