From 45605f8362dccbb96e78b386690459f86df3bf0f Mon Sep 17 00:00:00 2001 From: Brian Caswell Date: Tue, 15 Jul 2025 14:47:38 -0400 Subject: [PATCH] add azure blob cache support (#12587) * add support for Azure Blob caching * add integration tests * address feedback --- docs/my-website/docs/caching/all_caches.md | 31 +++ litellm/caching/Readme.md | 3 +- litellm/caching/__init__.py | 3 +- litellm/caching/azure_blob_cache.py | 103 +++++++++ litellm/caching/caching.py | 8 + litellm/types/caching.py | 1 + litellm/utils.py | 1 + poetry.lock | 30 ++- pyproject.toml | 3 + .../caching/test_azure_blob_cache.py | 198 ++++++++++++++++++ 10 files changed, 375 insertions(+), 6 deletions(-) create mode 100644 litellm/caching/azure_blob_cache.py create mode 100644 tests/test_litellm/caching/test_azure_blob_cache.py diff --git a/docs/my-website/docs/caching/all_caches.md b/docs/my-website/docs/caching/all_caches.md index b331646d5dc..a6be3396291 100644 --- a/docs/my-website/docs/caching/all_caches.md +++ b/docs/my-website/docs/caching/all_caches.md @@ -88,6 +88,37 @@ response2 = completion( + + +Install azure-storage-blob and azure-identity +```shell +pip install azure-storage-blob azure-identity +``` + +```python +import litellm +from litellm import completion +from litellm.caching.caching import Cache +from azure.identity import DefaultAzureCredential + +# pass Azure Blob Storage account URL and container name +litellm.cache = Cache(type="azure-blob", azure_account_url="https://example.blob.core.windows.net", azure_blob_container="litellm") + +# Make completion calls +response1 = completion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Tell me a joke."}] +) +response2 = completion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Tell me a joke."}] +) + +# response1 == response2, response 1 is cached +``` + + + diff --git a/litellm/caching/Readme.md b/litellm/caching/Readme.md index 6b0210a6696..1d920219830 100644 --- a/litellm/caching/Readme.md +++ b/litellm/caching/Readme.md @@ -10,7 +10,8 @@ The following caching mechanisms are supported: 4. **InMemoryCache** 5. **DiskCache** 6. **S3Cache** -7. **DualCache** (updates both Redis and an in-memory cache simultaneously) +7. **AzureBlobCache** +8. **DualCache** (updates both Redis and an in-memory cache simultaneously) ## Folder Structure diff --git a/litellm/caching/__init__.py b/litellm/caching/__init__.py index e10d01ff022..badc462e09b 100644 --- a/litellm/caching/__init__.py +++ b/litellm/caching/__init__.py @@ -1,3 +1,4 @@ +from .azure_blob_cache import AzureBlobCache from .caching import Cache, LiteLLMCacheType from .disk_cache import DiskCache from .dual_cache import DualCache @@ -6,4 +7,4 @@ from .qdrant_semantic_cache import QdrantSemanticCache from .redis_cache import RedisCache from .redis_cluster_cache import RedisClusterCache from .redis_semantic_cache import RedisSemanticCache -from .s3_cache import S3Cache +from .s3_cache import S3Cache \ No newline at end of file diff --git a/litellm/caching/azure_blob_cache.py b/litellm/caching/azure_blob_cache.py new file mode 100644 index 00000000000..45e551bdae9 --- /dev/null +++ b/litellm/caching/azure_blob_cache.py @@ -0,0 +1,103 @@ +""" +Azure Blob Cache implementation + +Has 4 methods: + - set_cache + - get_cache + - async_set_cache + - async_get_cache +""" + +import asyncio +import json +from contextlib import suppress + +from litellm._logging import print_verbose, verbose_logger + +from .base_cache import BaseCache + + +class AzureBlobCache(BaseCache): + def __init__(self, account_url, container) -> None: + from azure.storage.blob import BlobServiceClient + from azure.core.exceptions import ResourceExistsError + from azure.identity import DefaultAzureCredential + from azure.identity.aio import DefaultAzureCredential as AsyncDefaultAzureCredential + from azure.storage.blob.aio import BlobServiceClient as AsyncBlobServiceClient + + self.container_client = BlobServiceClient( + account_url=account_url, + credential=DefaultAzureCredential(), + ).get_container_client(container) + self.async_container_client = AsyncBlobServiceClient( + account_url=account_url, + credential=AsyncDefaultAzureCredential(), + ).get_container_client(container) + + with suppress(ResourceExistsError): + self.container_client.create_container() + + 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) + try: + self.container_client.upload_blob(key, serialized_value) + except Exception as e: + # NON blocking - notify users Azure Blob is throwing an exception + print_verbose(f"LiteLLM set_cache() - Got exception from Azure Blob: {e}") + + 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) + try: + await self.async_container_client.upload_blob(key, serialized_value, overwrite=True) + except Exception as e: + # NON blocking - notify users Azure Blob is throwing an exception + print_verbose(f"LiteLLM set_cache() - Got exception from Azure Blob: {e}") + + def get_cache(self, key, **kwargs): + from azure.core.exceptions import ResourceNotFoundError + + try: + print_verbose(f"Get Azure Blob Cache: key: {key}") + as_bytes = self.container_client.download_blob(key).readall() + as_str = as_bytes.decode("utf-8") + cached_response = json.loads(as_str) + + verbose_logger.debug( + f"Got Azure Blob Cache: key: {key}, cached_response {cached_response}. Type Response {type(cached_response)}" + ) + + return cached_response + except ResourceNotFoundError: + return None + + async def async_get_cache(self, key, **kwargs): + from azure.core.exceptions import ResourceNotFoundError + + try: + print_verbose(f"Get Azure Blob Cache: key: {key}") + blob = await self.async_container_client.download_blob(key) + as_bytes = await blob.readall() + as_str = as_bytes.decode("utf-8") + cached_response = json.loads(as_str) + verbose_logger.debug( + f"Got Azure Blob Cache: key: {key}, cached_response {cached_response}. Type Response {type(cached_response)}" + ) + return cached_response + except ResourceNotFoundError: + return None + + def flush_cache(self) -> None: + for blob in self.container_client.walk_blobs(): + self.container_client.delete_blob(blob.name) + + async def disconnect(self) -> None: + self.container_client.close() + await self.async_container_client.close() + + async def async_set_cache_pipeline(self, cache_list, **kwargs) -> None: + tasks = [] + for val in cache_list: + tasks.append(self.async_set_cache(val[0], val[1], **kwargs)) + await asyncio.gather(*tasks) diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 7adede79619..9165fec1e3f 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -24,6 +24,7 @@ from litellm.litellm_core_utils.model_param_helper import ModelParamHelper from litellm.types.caching import * from litellm.types.utils import EmbeddingResponse, all_litellm_params +from .azure_blob_cache import AzureBlobCache from .base_cache import BaseCache from .disk_cache import DiskCache from .dual_cache import DualCache # noqa @@ -78,6 +79,8 @@ class Cache: "rerank", ], # s3 Bucket, boto3 configuration + azure_account_url: Optional[str] = None, + azure_blob_container: Optional[str] = None, s3_bucket_name: Optional[str] = None, s3_region_name: Optional[str] = None, s3_api_version: Optional[str] = None, @@ -201,6 +204,11 @@ class Cache: s3_path=s3_path, **kwargs, ) + elif type == LiteLLMCacheType.AZURE_BLOB: + self.cache = AzureBlobCache( + account_url=azure_account_url, + container=azure_blob_container, + ) elif type == LiteLLMCacheType.DISK: self.cache = DiskCache(disk_cache_dir=disk_cache_dir) if "cache" not in litellm.input_callback: diff --git a/litellm/types/caching.py b/litellm/types/caching.py index 4893524998e..e457fe8a127 100644 --- a/litellm/types/caching.py +++ b/litellm/types/caching.py @@ -11,6 +11,7 @@ class LiteLLMCacheType(str, Enum): S3 = "s3" DISK = "disk" QDRANT_SEMANTIC = "qdrant-semantic" + AZURE_BLOB = "azure-blob" CachingSupportedCallTypes = Literal[ diff --git a/litellm/utils.py b/litellm/utils.py index a92f75e079f..82cea70e524 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -252,6 +252,7 @@ from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreCon from ._logging import _is_debugging_on, verbose_logger from .caching.caching import ( + AzureBlobCache, Cache, QdrantSemanticCache, RedisCache, diff --git a/poetry.lock b/poetry.lock index dc4ef0fcbcc..dfeba30efe5 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.1.2 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.1.3 and should not be changed by hand. [[package]] name = "aiohappyeyeballs" @@ -312,6 +312,28 @@ azure-core = ">=1.31.0" isodate = ">=0.6.1" typing-extensions = ">=4.0.1" +[[package]] +name = "azure-storage-blob" +version = "12.25.1" +description = "Microsoft Azure Blob Storage Client Library for Python" +optional = true +python-versions = ">=3.8" +groups = ["main"] +markers = "extra == \"proxy\"" +files = [ + {file = "azure_storage_blob-12.25.1-py3-none-any.whl", hash = "sha256:1f337aab12e918ec3f1b638baada97550673911c4ceed892acc8e4e891b74167"}, + {file = "azure_storage_blob-12.25.1.tar.gz", hash = "sha256:4f294ddc9bc47909ac66b8934bd26b50d2000278b10ad82cc109764fdc6e0e3b"}, +] + +[package.dependencies] +azure-core = ">=1.30.0" +cryptography = ">=2.1.4" +isodate = ">=0.6.1" +typing-extensions = ">=4.6.0" + +[package.extras] +aio = ["azure-core[aio] (>=1.30.0)"] + [[package]] name = "babel" version = "2.17.0" @@ -1642,7 +1664,7 @@ description = "An ISO 8601 date/time/duration parser and formatter" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "extra == \"extra-proxy\"" +markers = "extra == \"extra-proxy\" or extra == \"proxy\"" files = [ {file = "isodate-0.7.2-py3-none-any.whl", hash = "sha256:28009937d8031054830160fce6d409ed342816b543597cece116d966c6d99e15"}, {file = "isodate-0.7.2.tar.gz", hash = "sha256:4cd1aa0f43ca76f4a6c6c0292a85f40b35ec2e43e315b59f06e6d32171a953e6"}, @@ -4986,10 +5008,10 @@ type = ["pytest-mypy"] [extras] caching = ["diskcache"] extra-proxy = ["azure-identity", "azure-keyvault-secrets", "google-cloud-kms", "prisma", "redisvl", "resend"] -proxy = ["PyJWT", "apscheduler", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "uvicorn", "uvloop", "websockets"] +proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "uvicorn", "uvloop", "websockets"] utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.8.1,<4.0, !=3.9.7" -content-hash = "75c4a45aee2e0c9b75f85ff562d61c18c8a80312fdfd3371eb8ad5ec1fcd04bf" +content-hash = "e669716cba43fbb590c821859cf317c3c52d843c19a9a5c852f64efc08e17cad" diff --git a/pyproject.toml b/pyproject.toml index a11b959f020..395a8e04a73 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -49,6 +49,7 @@ cryptography = {version = "^43.0.1", optional = true} prisma = {version = "0.11.0", optional = true} azure-identity = {version = "^1.15.0", optional = true} azure-keyvault-secrets = {version = "^4.8.0", optional = true} +azure-storage-blob = {version="^12.25.1", optional=true} google-cloud-kms = {version = "^2.21.3", optional = true} resend = {version = "^0.8.0", optional = true} pynacl = {version = "^1.5.0", optional = true} @@ -79,6 +80,8 @@ proxy = [ "pynacl", "websockets", "boto3", + "azure-identity", + "azure-storage-blob", "mcp", "litellm-proxy-extras", "litellm-enterprise", diff --git a/tests/test_litellm/caching/test_azure_blob_cache.py b/tests/test_litellm/caching/test_azure_blob_cache.py new file mode 100644 index 00000000000..42b5eeb3d33 --- /dev/null +++ b/tests/test_litellm/caching/test_azure_blob_cache.py @@ -0,0 +1,198 @@ +import os +import sys +from unittest.mock import MagicMock, patch, AsyncMock + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +from litellm.caching.azure_blob_cache import AzureBlobCache + + +@pytest.fixture +def mock_azure_dependencies(): + """Mock all Azure dependencies to avoid requiring actual Azure credentials""" + + # Create mock container clients that will be assigned to the cache instance + mock_container_client = MagicMock() + mock_async_container_client = AsyncMock() + + # Mock credentials + mock_credential = MagicMock() + mock_async_credential = AsyncMock() + + # Create mock blob service clients that return the container clients + mock_blob_service_client = MagicMock() + mock_blob_service_client.get_container_client.return_value = mock_container_client + + mock_async_blob_service_client = AsyncMock() + # For AsyncMock, we need to make get_container_client return the mock directly, not a coroutine + mock_async_blob_service_client.get_container_client = MagicMock(return_value=mock_async_container_client) + + # Patch Azure dependencies at their source locations + with patch("azure.identity.DefaultAzureCredential", return_value=mock_credential), \ + patch("azure.identity.aio.DefaultAzureCredential", return_value=mock_async_credential), \ + patch("azure.storage.blob.BlobServiceClient", return_value=mock_blob_service_client), \ + patch("azure.storage.blob.aio.BlobServiceClient", return_value=mock_async_blob_service_client), \ + patch("azure.core.exceptions.ResourceExistsError"): + + yield { + "container_client": mock_container_client, + "async_container_client": mock_async_container_client, + "blob_service_client": mock_blob_service_client, + "async_blob_service_client": mock_async_blob_service_client, + "credential": mock_credential, + "async_credential": mock_async_credential, + } + + +@pytest.mark.asyncio +async def test_blob_cache_async_get_cache(mock_azure_dependencies): + """Test async_get_cache method with mocked Azure dependencies""" + + # Create cache instance (this will use the mocked dependencies) + cache = AzureBlobCache("https://my-test-host", "test-container") + + # Mock the download_blob response + mock_blob = AsyncMock() + mock_blob.readall.return_value = b'{"test_key": "test_value"}' + + # Set up the mock for download_blob on the actual container client instance + cache.async_container_client.download_blob.return_value = mock_blob + + # Test successful cache retrieval + result = await cache.async_get_cache("test_key") + + # Verify the call was made correctly + cache.async_container_client.download_blob.assert_called_once_with("test_key") + mock_blob.readall.assert_called_once() + + # Check the result + assert result == {"test_key": "test_value"} + + +@pytest.mark.asyncio +async def test_blob_cache_async_get_cache_not_found(mock_azure_dependencies): + """Test async_get_cache method when blob is not found""" + + # Import the exception inside the test to avoid import issues + from azure.core.exceptions import ResourceNotFoundError + + cache = AzureBlobCache("https://my-test-host", "test-container") + + # Mock ResourceNotFoundError + cache.async_container_client.download_blob.side_effect = ResourceNotFoundError("Blob not found") + + # Test cache miss + result = await cache.async_get_cache("nonexistent_key") + + # Verify the call was made and result is None + cache.async_container_client.download_blob.assert_called_once_with("nonexistent_key") + assert result is None + + +@pytest.mark.asyncio +async def test_blob_cache_async_set_cache(mock_azure_dependencies): + """Test async_set_cache method with mocked Azure dependencies""" + + cache = AzureBlobCache("https://my-test-host", "test-container") + + test_value = {"key": "value", "number": 42} + + # Test setting cache + await cache.async_set_cache("test_key", test_value) + + # Verify the call was made correctly + cache.async_container_client.upload_blob.assert_called_once_with( + "test_key", + '{"key": "value", "number": 42}', + overwrite=True + ) + + +def test_blob_cache_sync_get_cache(mock_azure_dependencies): + """Test sync get_cache method with mocked Azure dependencies""" + + cache = AzureBlobCache("https://my-test-host", "test-container") + + # Mock the download_blob response + mock_blob = MagicMock() + mock_blob.readall.return_value = b'{"sync_key": "sync_value"}' + + cache.container_client.download_blob.return_value = mock_blob + + # Test successful cache retrieval + result = cache.get_cache("sync_key") + + # Verify the call was made correctly + cache.container_client.download_blob.assert_called_once_with("sync_key") + mock_blob.readall.assert_called_once() + + # Check the result + assert result == {"sync_key": "sync_value"} + + +def test_blob_cache_sync_set_cache(mock_azure_dependencies): + """Test sync set_cache method with mocked Azure dependencies""" + + cache = AzureBlobCache("https://my-test-host", "test-container") + + test_value = {"sync_key": "sync_value", "number": 123} + + # Test setting cache + cache.set_cache("sync_test_key", test_value) + + # Verify the call was made correctly + cache.container_client.upload_blob.assert_called_once_with( + "sync_test_key", + '{"sync_key": "sync_value", "number": 123}' + ) + + +def test_blob_cache_sync_get_cache_not_found(mock_azure_dependencies): + """Test sync get_cache method when blob is not found""" + + from azure.core.exceptions import ResourceNotFoundError + + cache = AzureBlobCache("https://my-test-host", "test-container") + + # Mock ResourceNotFoundError + cache.container_client.download_blob.side_effect = ResourceNotFoundError("Blob not found") + + # Test cache miss + result = cache.get_cache("nonexistent_key") + + # Verify the call was made and result is None + cache.container_client.download_blob.assert_called_once_with("nonexistent_key") + assert result is None + + +@pytest.mark.asyncio +async def test_blob_cache_async_set_cache_pipeline(mock_azure_dependencies): + """Test async_set_cache_pipeline method with mocked Azure dependencies""" + + cache = AzureBlobCache("https://my-test-host", "test-container") + + # Test data for pipeline + cache_list = [ + ("key1", {"value": "data1"}), + ("key2", {"value": "data2"}), + ("key3", {"value": "data3"}), + ] + + # Test pipeline cache setting + await cache.async_set_cache_pipeline(cache_list) + + # Verify all calls were made correctly + expected_calls = [ + (("key1", '{"value": "data1"}'), {"overwrite": True}), + (("key2", '{"value": "data2"}'), {"overwrite": True}), + (("key3", '{"value": "data3"}'), {"overwrite": True}), + ] + + assert cache.async_container_client.upload_blob.call_count == 3 + for expected_call in expected_calls: + cache.async_container_client.upload_blob.assert_any_call(*expected_call[0], **expected_call[1])