mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
add azure blob cache support (#12587)
* add support for Azure Blob caching * add integration tests * address feedback
This commit is contained in:
parent
720b94fd2b
commit
45605f8362
10 changed files with 375 additions and 6 deletions
|
|
@ -88,6 +88,37 @@ response2 = completion(
|
|||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="azureblob" label="azure-blob-cache">
|
||||
|
||||
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
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
|
||||
<TabItem value="redis-sem" label="redis-semantic cache">
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
103
litellm/caching/azure_blob_cache.py
Normal file
103
litellm/caching/azure_blob_cache.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ class LiteLLMCacheType(str, Enum):
|
|||
S3 = "s3"
|
||||
DISK = "disk"
|
||||
QDRANT_SEMANTIC = "qdrant-semantic"
|
||||
AZURE_BLOB = "azure-blob"
|
||||
|
||||
|
||||
CachingSupportedCallTypes = Literal[
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
30
poetry.lock
generated
30
poetry.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
198
tests/test_litellm/caching/test_azure_blob_cache.py
Normal file
198
tests/test_litellm/caching/test_azure_blob_cache.py
Normal file
|
|
@ -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])
|
||||
Loading…
Add table
Reference in a new issue