Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .semversioner/next-release/patch-20260811050438120824.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
{
"type": "patch",
"description": "Preserve caller event loops in synchronous LLM cache middleware."
}
71 changes: 36 additions & 35 deletions packages/graphrag-llm/graphrag_llm/middleware/with_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,41 +67,42 @@ def _cache_middleware(
cache_key = cache_key_creator(kwargs)

event_loop = asyncio.new_event_loop()
asyncio.set_event_loop(event_loop)
cached_response = event_loop.run_until_complete(cache.get(cache_key))
if (
cached_response is not None
and isinstance(cached_response, dict)
and "response" in cached_response
and cached_response["response"] is not None
and isinstance(cached_response["response"], dict)
):
try:
if (
metrics is not None
and "metrics" in cached_response
and cached_response["metrics"] is not None
and isinstance(cached_response["metrics"], dict)
):
metrics.update(cached_response["metrics"])
metrics["cached_responses"] = 1

if request_type == "chat":
return LLMCompletionResponse(**cached_response["response"])
return LLMEmbeddingResponse(**cached_response["response"])
except Exception: # noqa: BLE001
# Try to retrieve value from cache but if it fails, continue
# to make the request.
...

response = sync_middleware(**kwargs)
cache_value = {
"response": response.model_dump(), # type: ignore
"metrics": metrics if metrics is not None else {},
}
event_loop.run_until_complete(cache.set(cache_key, cache_value))
event_loop.close()
return response
try:
cached_response = event_loop.run_until_complete(cache.get(cache_key))
if (
cached_response is not None
and isinstance(cached_response, dict)
and "response" in cached_response
and cached_response["response"] is not None
and isinstance(cached_response["response"], dict)
):
try:
if (
metrics is not None
and "metrics" in cached_response
and cached_response["metrics"] is not None
and isinstance(cached_response["metrics"], dict)
):
metrics.update(cached_response["metrics"])
metrics["cached_responses"] = 1

if request_type == "chat":
return LLMCompletionResponse(**cached_response["response"])
return LLMEmbeddingResponse(**cached_response["response"])
except Exception: # noqa: BLE001
# Try to retrieve value from cache but if it fails, continue
# to make the request.
...

response = sync_middleware(**kwargs)
cache_value = {
"response": response.model_dump(), # type: ignore
"metrics": metrics if metrics is not None else {},
}
event_loop.run_until_complete(cache.set(cache_key, cache_value))
return response
finally:
event_loop.close()

async def _cache_middleware_async(
**kwargs: Any,
Expand Down
2 changes: 2 additions & 0 deletions tests/unit/language_model/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
# Copyright (c) 2024 Microsoft Corporation.
# Licensed under the MIT License
118 changes: 118 additions & 0 deletions tests/unit/language_model/test_cache_middleware.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
# Copyright (c) 2024 Microsoft Corporation.
# Licensed under the MIT License

"""Unit tests for the LLM cache middleware."""

import asyncio
from collections.abc import Callable
from typing import Any

import pytest
from graphrag_cache.memory_cache import MemoryCache
from graphrag_llm.middleware.with_cache import with_cache
from graphrag_llm.types import LLMCompletionResponse
from graphrag_llm.utils import create_completion_response


@pytest.fixture
def tracked_event_loops(monkeypatch: pytest.MonkeyPatch):
"""Install a caller-owned loop and track loops created by the middleware."""
original_loop = asyncio.new_event_loop()
create_event_loop = asyncio.new_event_loop
created_loops: list[asyncio.AbstractEventLoop] = []

def _create_event_loop() -> asyncio.AbstractEventLoop:
event_loop = create_event_loop()
created_loops.append(event_loop)
return event_loop

asyncio.set_event_loop(original_loop)
monkeypatch.setattr(asyncio, "new_event_loop", _create_event_loop)

yield original_loop, created_loops

asyncio.set_event_loop(None)
original_loop.close()
for event_loop in created_loops:
if not event_loop.is_closed():
event_loop.close()


def _with_sync_cache(
cache: MemoryCache,
sync_middleware: Callable[..., LLMCompletionResponse],
):
async def _async_middleware(**kwargs: Any) -> LLMCompletionResponse:
return await asyncio.to_thread(sync_middleware, **kwargs)

def _cache_key(input_args: dict[str, Any]) -> str:
return "cache-key"

cached_middleware, _ = with_cache(
sync_middleware=sync_middleware,
async_middleware=_async_middleware,
request_type="chat",
cache=cache,
cache_key_creator=_cache_key,
)
return cached_middleware


def test_sync_cache_preserves_event_loop_on_miss(tracked_event_loops) -> None:
"""The sync cache should not replace the caller's loop on a cache miss."""
original_loop, created_loops = tracked_event_loops
response = create_completion_response("uncached")
cached_middleware = _with_sync_cache(MemoryCache(), lambda **_: response)

cached_response = cached_middleware(messages=[])

assert isinstance(cached_response, LLMCompletionResponse)
assert cached_response.content == "uncached"
assert asyncio.get_event_loop() is original_loop
assert len(created_loops) == 1
assert created_loops[0].is_closed()


def test_sync_cache_preserves_event_loop_on_hit(tracked_event_loops) -> None:
"""The sync cache should close its loop before returning a cached response."""
original_loop, created_loops = tracked_event_loops
response = create_completion_response("cached")
cache = MemoryCache()
asyncio.run(
cache.set(
"cache-key",
{"response": response.model_dump(), "metrics": {}},
)
)
asyncio.set_event_loop(original_loop)

def _unexpected_request(**_: Any) -> LLMCompletionResponse:
pytest.fail("The wrapped middleware should not run on a cache hit.")

cached_middleware = _with_sync_cache(cache, _unexpected_request)

cached_response = cached_middleware(messages=[])

assert isinstance(cached_response, LLMCompletionResponse)
assert cached_response.content == "cached"
assert asyncio.get_event_loop() is original_loop
assert len(created_loops) == 1
assert created_loops[0].is_closed()


def test_sync_cache_closes_event_loop_on_error(tracked_event_loops) -> None:
"""The sync cache should close its loop when the wrapped middleware fails."""
original_loop, created_loops = tracked_event_loops

def _raise_error(**_: Any) -> LLMCompletionResponse:
msg = "request failed"
raise RuntimeError(msg)

cached_middleware = _with_sync_cache(MemoryCache(), _raise_error)

with pytest.raises(RuntimeError, match="request failed"):
cached_middleware(messages=[])

assert asyncio.get_event_loop() is original_loop
assert len(created_loops) == 1
assert created_loops[0].is_closed()