mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: Redis caching timedelta serialization error
This commit is contained in:
parent
6a47ac15ab
commit
c83b0337f2
5 changed files with 292 additions and 8 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
240
tests/test_litellm/caching/test_timedelta_serialization.py
Normal file
240
tests/test_litellm/caching/test_timedelta_serialization.py
Normal 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]
|
||||
Loading…
Add table
Reference in a new issue