diff --git a/litellm/caching/azure_blob_cache.py b/litellm/caching/azure_blob_cache.py index 45e551bdae9..e6f27351b6f 100644 --- a/litellm/caching/azure_blob_cache.py +++ b/litellm/caching/azure_blob_cache.py @@ -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: diff --git a/litellm/caching/gcs_cache.py b/litellm/caching/gcs_cache.py index 88857ba0e70..a128595049d 100644 --- a/litellm/caching/gcs_cache.py +++ b/litellm/caching/gcs_cache.py @@ -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}") diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 47bc0222ed5..85b323f91b5 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -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: diff --git a/litellm/caching/s3_cache.py b/litellm/caching/s3_cache.py index 180964605f6..64637483764 100644 --- a/litellm/caching/s3_cache.py +++ b/litellm/caching/s3_cache.py @@ -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: diff --git a/tests/test_litellm/caching/test_timedelta_serialization.py b/tests/test_litellm/caching/test_timedelta_serialization.py new file mode 100644 index 00000000000..fca0a8f5850 --- /dev/null +++ b/tests/test_litellm/caching/test_timedelta_serialization.py @@ -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]