From ce533391367eb17081b5d5b12a349e4161048410 Mon Sep 17 00:00:00 2001 From: FU-max-boop <214359569+FU-max-boop@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:07:16 +0800 Subject: [PATCH] fix(llm): preserve event loop in sync cache --- .../patch-20260811050438120824.json | 4 + .../graphrag_llm/middleware/with_cache.py | 71 +++++------ tests/unit/language_model/__init__.py | 2 + .../language_model/test_cache_middleware.py | 118 ++++++++++++++++++ 4 files changed, 160 insertions(+), 35 deletions(-) create mode 100644 .semversioner/next-release/patch-20260811050438120824.json create mode 100644 tests/unit/language_model/__init__.py create mode 100644 tests/unit/language_model/test_cache_middleware.py diff --git a/.semversioner/next-release/patch-20260811050438120824.json b/.semversioner/next-release/patch-20260811050438120824.json new file mode 100644 index 000000000..54b8d46a9 --- /dev/null +++ b/.semversioner/next-release/patch-20260811050438120824.json @@ -0,0 +1,4 @@ +{ + "type": "patch", + "description": "Preserve caller event loops in synchronous LLM cache middleware." +} diff --git a/packages/graphrag-llm/graphrag_llm/middleware/with_cache.py b/packages/graphrag-llm/graphrag_llm/middleware/with_cache.py index 280953807..fad59e7fa 100644 --- a/packages/graphrag-llm/graphrag_llm/middleware/with_cache.py +++ b/packages/graphrag-llm/graphrag_llm/middleware/with_cache.py @@ -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, diff --git a/tests/unit/language_model/__init__.py b/tests/unit/language_model/__init__.py new file mode 100644 index 000000000..0a3e38adf --- /dev/null +++ b/tests/unit/language_model/__init__.py @@ -0,0 +1,2 @@ +# Copyright (c) 2024 Microsoft Corporation. +# Licensed under the MIT License diff --git a/tests/unit/language_model/test_cache_middleware.py b/tests/unit/language_model/test_cache_middleware.py new file mode 100644 index 000000000..a2253cb6e --- /dev/null +++ b/tests/unit/language_model/test_cache_middleware.py @@ -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()