mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(caching): retain references to background cache-write tasks
LLMCachingHandler.async_set_cache scheduled cache writes with bare asyncio.create_task(...) whose results were discarded. asyncio only keeps a weak reference to a task, so these writes could be garbage-collected before completing, silently dropping cache entries (and wasting spend on the resulting cache misses). Track the tasks in a module-level set and release them via a done-callback, matching the pattern already used elsewhere in litellm (e.g. the responses streaming iterator's cache_write_task, the akto guardrail, and the logging worker). Adds unit tests covering the success and exception paths.
This commit is contained in:
parent
5699a06413
commit
d953270259
2 changed files with 82 additions and 9 deletions
|
|
@ -27,6 +27,7 @@ from typing import (
|
|||
Generator,
|
||||
List,
|
||||
Optional,
|
||||
Set,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
|
|
@ -109,6 +110,20 @@ def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bo
|
|||
return kwargs.get("stream", False) is True
|
||||
|
||||
|
||||
# Strong references to fire-and-forget cache-write tasks. asyncio only keeps a
|
||||
# weak reference to a task created with create_task, so without this a cache
|
||||
# write could be garbage-collected before it completes, silently dropping the
|
||||
# entry. Tasks remove themselves from the set on completion.
|
||||
_background_cache_tasks: Set["asyncio.Task[Any]"] = set()
|
||||
|
||||
|
||||
def _track_cache_task(task: "asyncio.Task[Any]") -> "asyncio.Task[Any]":
|
||||
"""Retain a strong reference to a background cache-write task until done."""
|
||||
_background_cache_tasks.add(task)
|
||||
task.add_done_callback(_background_cache_tasks.discard)
|
||||
return task
|
||||
|
||||
|
||||
class LLMCachingHandler:
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -1020,21 +1035,29 @@ class LLMCachingHandler:
|
|||
litellm.cache.cache, S3Cache
|
||||
) # s3 doesn't support bulk writing. Exclude.
|
||||
):
|
||||
asyncio.create_task(
|
||||
litellm.cache.async_add_cache_pipeline(
|
||||
result, dynamic_cache_object=self.dual_cache, **new_kwargs
|
||||
_track_cache_task(
|
||||
asyncio.create_task(
|
||||
litellm.cache.async_add_cache_pipeline(
|
||||
result, dynamic_cache_object=self.dual_cache, **new_kwargs
|
||||
)
|
||||
)
|
||||
)
|
||||
else:
|
||||
asyncio.create_task(
|
||||
litellm.cache.async_add_cache(
|
||||
result.model_dump_json(),
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
**new_kwargs,
|
||||
_track_cache_task(
|
||||
asyncio.create_task(
|
||||
litellm.cache.async_add_cache(
|
||||
result.model_dump_json(),
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
**new_kwargs,
|
||||
)
|
||||
)
|
||||
)
|
||||
else:
|
||||
asyncio.create_task(litellm.cache.async_add_cache(result, **new_kwargs))
|
||||
_track_cache_task(
|
||||
asyncio.create_task(
|
||||
litellm.cache.async_add_cache(result, **new_kwargs)
|
||||
)
|
||||
)
|
||||
|
||||
def sync_set_cache(
|
||||
self,
|
||||
|
|
|
|||
50
tests/litellm/caching/test_background_cache_tasks.py
Normal file
50
tests/litellm/caching/test_background_cache_tasks.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
"""Tests for background cache-write task tracking in the caching handler.
|
||||
|
||||
asyncio only keeps a weak reference to a task created with create_task, so a
|
||||
fire-and-forget cache write can be garbage-collected before it finishes,
|
||||
silently dropping the entry. _track_cache_task retains a strong reference until
|
||||
the task completes.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.caching.caching_handler import (
|
||||
_background_cache_tasks,
|
||||
_track_cache_task,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_track_cache_task_keeps_reference_until_done():
|
||||
async def _write():
|
||||
await asyncio.sleep(0)
|
||||
return "written"
|
||||
|
||||
task = _track_cache_task(asyncio.create_task(_write()))
|
||||
|
||||
# While pending, a strong reference is held in the tracking set.
|
||||
assert task in _background_cache_tasks
|
||||
|
||||
result = await task
|
||||
|
||||
# Once done, the task delivers its result and the reference is released.
|
||||
assert result == "written"
|
||||
await asyncio.sleep(0) # allow the done-callback to run
|
||||
assert task not in _background_cache_tasks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_track_cache_task_releases_reference_on_exception():
|
||||
async def _boom():
|
||||
raise RuntimeError("cache backend down")
|
||||
|
||||
task = _track_cache_task(asyncio.create_task(_boom()))
|
||||
assert task in _background_cache_tasks
|
||||
|
||||
with pytest.raises(RuntimeError, match="cache backend down"):
|
||||
await task
|
||||
|
||||
await asyncio.sleep(0)
|
||||
assert task not in _background_cache_tasks
|
||||
Loading…
Add table
Reference in a new issue