fix: Redis caching timedelta serialization error

This commit is contained in:
swarnabhasinha 2025-09-10 13:31:31 +05:30
parent 6a47ac15ab
commit c83b0337f2
No known key found for this signature in database
GPG key ID: 9B0C3EA24F21722E
5 changed files with 292 additions and 8 deletions

View file

@ -11,12 +11,22 @@ Has 4 methods:
import asyncio
import json
from contextlib import suppress
from datetime import timedelta
from litellm._logging import print_verbose, verbose_logger
from .base_cache import BaseCache
class TimedeltaJSONEncoder(json.JSONEncoder):
"""Custom JSON encoder that handles timedelta objects by converting them to seconds."""
def default(self, obj):
if isinstance(obj, timedelta):
return obj.total_seconds()
return super().default(obj)
class AzureBlobCache(BaseCache):
def __init__(self, account_url, container) -> None:
from azure.storage.blob import BlobServiceClient
@ -39,7 +49,7 @@ class AzureBlobCache(BaseCache):
def set_cache(self, key, value, **kwargs) -> None:
print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}")
serialized_value = json.dumps(value)
serialized_value = json.dumps(value, cls=TimedeltaJSONEncoder)
try:
self.container_client.upload_blob(key, serialized_value)
except Exception as e:
@ -48,7 +58,7 @@ class AzureBlobCache(BaseCache):
async def async_set_cache(self, key, value, **kwargs) -> None:
print_verbose(f"LiteLLM SET Cache - Azure Blob. Key={key}. Value={value}")
serialized_value = json.dumps(value)
serialized_value = json.dumps(value, cls=TimedeltaJSONEncoder)
try:
await self.async_container_client.upload_blob(key, serialized_value, overwrite=True)
except Exception as e:

View file

@ -3,6 +3,7 @@ Supports syncing responses to Google Cloud Storage Buckets using HTTP requests.
"""
import json
import asyncio
from datetime import timedelta
from typing import Optional
from litellm._logging import print_verbose, verbose_logger
@ -15,6 +16,15 @@ from litellm.llms.custom_httpx.http_handler import (
from .base_cache import BaseCache
class TimedeltaJSONEncoder(json.JSONEncoder):
"""Custom JSON encoder that handles timedelta objects by converting them to seconds."""
def default(self, obj):
if isinstance(obj, timedelta):
return obj.total_seconds()
return super().default(obj)
class GCSCache(BaseCache):
def __init__(self, bucket_name: Optional[str] = None, path_service_account: Optional[str] = None, gcs_path: Optional[str] = None) -> None:
super().__init__()
@ -38,7 +48,7 @@ class GCSCache(BaseCache):
object_name = self.key_prefix + key
bucket_name = self.bucket_name
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
data = json.dumps(value)
data = json.dumps(value, cls=TimedeltaJSONEncoder)
self.sync_client.post(url=url, data=data, headers=headers)
except Exception as e:
print_verbose(f"GCS Caching: set_cache() - Got exception from GCS: {e}")
@ -49,7 +59,7 @@ class GCSCache(BaseCache):
object_name = self.key_prefix + key
bucket_name = self.bucket_name
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
data = json.dumps(value)
data = json.dumps(value, cls=TimedeltaJSONEncoder)
await self.async_client.post(url=url, data=data, headers=headers)
except Exception as e:
print_verbose(f"GCS Caching: async_set_cache() - Got exception from GCS: {e}")

View file

@ -24,6 +24,16 @@ from litellm.types.services import ServiceTypes
from .base_cache import BaseCache
class TimedeltaJSONEncoder(json.JSONEncoder):
"""Custom JSON encoder that handles timedelta objects by converting them to seconds."""
def default(self, obj):
if isinstance(obj, timedelta):
return obj.total_seconds()
return super().default(obj)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from redis.asyncio import Redis, RedisCluster
@ -214,7 +224,12 @@ class RedisCache(BaseCache):
key = self.check_and_fix_namespace(key=key)
try:
start_time = time.time()
self.redis_client.set(name=key, value=str(value), ex=ttl)
# Convert value to JSON string to handle complex objects like timedelta
if isinstance(value, (dict, list)) or hasattr(value, '__dict__'):
serialized_value = json.dumps(value, cls=TimedeltaJSONEncoder)
else:
serialized_value = str(value)
self.redis_client.set(name=key, value=serialized_value, ex=ttl)
end_time = time.time()
_duration = end_time - start_time
self.service_logger_obj.service_success_hook(
@ -400,7 +415,7 @@ class RedisCache(BaseCache):
raise Exception("Redis client cannot set cache. Attribute not found.")
result = await _redis_client.set(
name=key,
value=json.dumps(value),
value=json.dumps(value, cls=TimedeltaJSONEncoder),
nx=nx,
ex=ttl,
)
@ -458,7 +473,7 @@ class RedisCache(BaseCache):
print_verbose(
f"Set ASYNC Redis Cache PIPELINE: key: {cache_key}\nValue {cache_value}\nttl={ttl}"
)
json_cache_value = json.dumps(cache_value)
json_cache_value = json.dumps(cache_value, cls=TimedeltaJSONEncoder)
# Set the value with a TTL if it's provided.
_td: Optional[timedelta] = None
if ttl is not None:

View file

@ -20,6 +20,15 @@ from litellm._logging import print_verbose, verbose_logger
from .base_cache import BaseCache
class TimedeltaJSONEncoder(json.JSONEncoder):
"""Custom JSON encoder that handles timedelta objects by converting them to seconds."""
def default(self, obj):
if isinstance(obj, timedelta):
return obj.total_seconds()
return super().default(obj)
class S3Cache(BaseCache):
def __init__(
self,
@ -65,7 +74,7 @@ class S3Cache(BaseCache):
print_verbose(f"LiteLLM SET Cache - S3. Key={key}. Value={value}")
ttl = kwargs.get("ttl", None)
# Convert value to JSON before storing in S3
serialized_value = json.dumps(value)
serialized_value = json.dumps(value, cls=TimedeltaJSONEncoder)
key = self._to_s3_key(key)
if ttl is not None:

View file

@ -0,0 +1,240 @@
"""
Test timedelta serialization in cache implementations.
This test ensures that timedelta objects can be properly serialized
to JSON in all cache implementations without causing serialization errors.
"""
import json
from datetime import timedelta
from unittest.mock import Mock, patch
import pytest
from litellm.caching.redis_cache import RedisCache, TimedeltaJSONEncoder
from litellm.caching.s3_cache import S3Cache, TimedeltaJSONEncoder as S3TimedeltaJSONEncoder
from litellm.caching.gcs_cache import GCSCache, TimedeltaJSONEncoder as GCSTimedeltaJSONEncoder
from litellm.caching.azure_blob_cache import AzureBlobCache, TimedeltaJSONEncoder as AzureTimedeltaJSONEncoder
class TestTimedeltaJSONEncoder:
"""Test the TimedeltaJSONEncoder class."""
def test_timedelta_serialization(self):
"""Test that timedelta objects are properly serialized to seconds."""
test_data = {
"latency": [timedelta(seconds=1.5), timedelta(milliseconds=500)],
"time_to_first_token": [timedelta(milliseconds=200)],
"normal_value": 42,
"nested": {
"response_time": timedelta(seconds=2.3),
"other_data": "test"
}
}
# This should not raise an exception
serialized = json.dumps(test_data, cls=TimedeltaJSONEncoder)
deserialized = json.loads(serialized)
# Verify timedelta objects were converted to seconds
assert deserialized["latency"] == [1.5, 0.5]
assert deserialized["time_to_first_token"] == [0.2]
assert deserialized["normal_value"] == 42
assert deserialized["nested"]["response_time"] == 2.3
assert deserialized["nested"]["other_data"] == "test"
def test_reproduces_original_error(self):
"""Test that reproduces the original error before the fix."""
test_data = {
"latency": [timedelta(seconds=1.5)],
"normal_value": 42
}
# This would have raised: "Object of type timedelta is not JSON serializable"
with pytest.raises(TypeError, match="timedelta is not JSON serializable"):
json.dumps(test_data)
# But with our custom encoder, it should work
serialized = json.dumps(test_data, cls=TimedeltaJSONEncoder)
deserialized = json.loads(serialized)
assert deserialized["latency"] == [1.5]
assert deserialized["normal_value"] == 42
class TestRedisCacheTimedeltaHandling:
"""Test Redis cache handling of timedelta objects."""
def test_redis_cache_sync_with_timedelta(self):
"""Test that RedisCache can handle timedelta objects without errors."""
# Mock Redis client
mock_redis_client = Mock()
mock_redis_client.set = Mock()
mock_redis_client.get = Mock(return_value=None)
mock_redis_client.ping = Mock(return_value=True)
# Create RedisCache instance with mocked client
with patch('litellm.caching.redis_cache.get_redis_client', return_value=mock_redis_client):
cache = RedisCache()
# Test data with timedelta objects (similar to what's stored in latency routing)
test_data = {
"deployment_1": {
"latency": [timedelta(seconds=1.2), timedelta(milliseconds=800)],
"time_to_first_token": [timedelta(milliseconds=150)],
"2024-01-01-10-30": {
"tpm": 1000,
"rpm": 10
}
}
}
# This should not raise an exception
cache.set_cache("test_key", test_data)
# Verify that set was called with JSON-serialized data
mock_redis_client.set.assert_called_once()
call_args = mock_redis_client.set.call_args
serialized_value = call_args[1]['value']
# Verify the serialized value is valid JSON
deserialized = json.loads(serialized_value)
assert deserialized["deployment_1"]["latency"] == [1.2, 0.8]
assert deserialized["deployment_1"]["time_to_first_token"] == [0.15]
@pytest.mark.asyncio
async def test_redis_cache_async_with_timedelta(self):
"""Test that async RedisCache can handle timedelta objects without errors."""
# Mock async Redis client
mock_async_redis_client = Mock()
mock_async_redis_client.set = Mock(return_value=True)
mock_async_redis_client.get = Mock(return_value=None)
mock_async_redis_client.ping = Mock(return_value=True)
# Create RedisCache instance with mocked client
with patch('litellm.caching.redis_cache.get_redis_async_client', return_value=mock_async_redis_client):
cache = RedisCache()
# Test data with timedelta objects
test_data = {
"deployment_1": {
"latency": [timedelta(seconds=1.2), timedelta(milliseconds=800)],
"time_to_first_token": [timedelta(milliseconds=150)],
"2024-01-01-10-30": {
"tpm": 1000,
"rpm": 10
}
}
}
# This should not raise an exception
await cache.async_set_cache("test_key", test_data)
# Verify that set was called with JSON-serialized data
mock_async_redis_client.set.assert_called_once()
call_args = mock_async_redis_client.set.call_args
serialized_value = call_args[1]['value']
# Verify the serialized value is valid JSON
deserialized = json.loads(serialized_value)
assert deserialized["deployment_1"]["latency"] == [1.2, 0.8]
assert deserialized["deployment_1"]["time_to_first_token"] == [0.15]
class TestS3CacheTimedeltaHandling:
"""Test S3 cache handling of timedelta objects."""
def test_s3_cache_with_timedelta(self):
"""Test that S3Cache can handle timedelta objects without errors."""
# Mock S3 client
mock_s3_client = Mock()
mock_s3_client.put_object = Mock()
# Create S3Cache instance with mocked client
with patch('litellm.caching.s3_cache.boto3.client', return_value=mock_s3_client):
cache = S3Cache(s3_bucket_name="test-bucket")
# Test data with timedelta objects
test_data = {
"latency": [timedelta(seconds=1.2), timedelta(milliseconds=800)],
"time_to_first_token": [timedelta(milliseconds=150)]
}
# This should not raise an exception
cache.set_cache("test_key", test_data)
# Verify that put_object was called with JSON-serialized data
mock_s3_client.put_object.assert_called_once()
call_args = mock_s3_client.put_object.call_args
serialized_value = call_args[1]['Body']
# Verify the serialized value is valid JSON
deserialized = json.loads(serialized_value)
assert deserialized["latency"] == [1.2, 0.8]
assert deserialized["time_to_first_token"] == [0.15]
class TestGCSCacheTimedeltaHandling:
"""Test GCS cache handling of timedelta objects."""
def test_gcs_cache_with_timedelta(self):
"""Test that GCSCache can handle timedelta objects without errors."""
# Mock HTTP client
mock_client = Mock()
mock_client.post = Mock()
# Create GCSCache instance with mocked client
with patch('litellm.caching.gcs_cache._get_httpx_client', return_value=mock_client):
cache = GCSCache(bucket_name="test-bucket")
# Test data with timedelta objects
test_data = {
"latency": [timedelta(seconds=1.2), timedelta(milliseconds=800)],
"time_to_first_token": [timedelta(milliseconds=150)]
}
# This should not raise an exception
cache.set_cache("test_key", test_data)
# Verify that post was called with JSON-serialized data
mock_client.post.assert_called_once()
call_args = mock_client.post.call_args
serialized_value = call_args[1]['data']
# Verify the serialized value is valid JSON
deserialized = json.loads(serialized_value)
assert deserialized["latency"] == [1.2, 0.8]
assert deserialized["time_to_first_token"] == [0.15]
class TestAzureBlobCacheTimedeltaHandling:
"""Test Azure Blob cache handling of timedelta objects."""
def test_azure_blob_cache_with_timedelta(self):
"""Test that AzureBlobCache can handle timedelta objects without errors."""
# Mock Azure Blob client
mock_container_client = Mock()
mock_container_client.upload_blob = Mock()
# Create AzureBlobCache instance with mocked client
with patch('litellm.caching.azure_blob_cache.BlobServiceClient'):
cache = AzureBlobCache(account_url="https://test.blob.core.windows.net", container="test-container")
cache.container_client = mock_container_client
# Test data with timedelta objects
test_data = {
"latency": [timedelta(seconds=1.2), timedelta(milliseconds=800)],
"time_to_first_token": [timedelta(milliseconds=150)]
}
# This should not raise an exception
cache.set_cache("test_key", test_data)
# Verify that upload_blob was called with JSON-serialized data
mock_container_client.upload_blob.assert_called_once()
call_args = mock_container_client.upload_blob.call_args
serialized_value = call_args[0][1] # Second positional argument
# Verify the serialized value is valid JSON
deserialized = json.loads(serialized_value)
assert deserialized["latency"] == [1.2, 0.8]
assert deserialized["time_to_first_token"] == [0.15]