diff --git a/tests/models/test_retrying_lite_llm.py b/tests/models/test_retrying_lite_llm.py index 03fa0fc8c..c8a70716f 100644 --- a/tests/models/test_retrying_lite_llm.py +++ b/tests/models/test_retrying_lite_llm.py @@ -14,12 +14,14 @@ from __future__ import annotations +import asyncio from types import SimpleNamespace import pytest from google.adk.models.lite_llm import LiteLlm from google.adk.models.llm_request import LlmRequest from google.adk.models.llm_response import LlmResponse +from google.adk.tools.base_tool import BaseTool from google.genai import types from veadk.models.retrying_lite_llm import RetryingLiteLlm @@ -42,6 +44,48 @@ def _request() -> LlmRequest: ) +@pytest.mark.asyncio +@pytest.mark.parametrize("retry", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +async def test_snapshot_preserves_live_tools_and_isolates_request_data( + monkeypatch: pytest.MonkeyPatch, retry: bool, stream: bool +) -> None: + tool = BaseTool(name="live_tool", description="Tool with asynchronous state") + tool.pending = asyncio.get_running_loop().create_future() + request = _request() + request.tools_dict = {tool.name: tool} + request.config.temperature = 0.5 + seen: list[LlmRequest] = [] + + async def generate(_self, llm_request, stream=False): + seen.append(llm_request) + assert llm_request.tools_dict[tool.name] is tool + assert llm_request.contents[0].parts[0].text == "hello" + assert llm_request.config.temperature == 0.5 + if retry and len(seen) == 1: + llm_request.contents[0].parts[0].text = "mutated" + llm_request.config.temperature = 0.1 + llm_request.tools_dict.clear() + raise _RateLimitError() + yield LlmResponse(content=types.Content(role="model", parts=[])) + + monkeypatch.setattr(LiteLlm, "generate_content_async", generate) + model = RetryingLiteLlm(model="openai/test-model") + try: + responses = [ + response + async for response in model.generate_content_async(request, stream=stream) + ] + finally: + tool.pending.cancel() + + assert len(responses) == 1 + assert len(seen) == (2 if retry else 1) + if retry: + assert seen[1] is not request + assert seen[1].tools_dict is not request.tools_dict + + @pytest.mark.asyncio async def test_retries_one_pre_output_429_from_pristine_request( monkeypatch: pytest.MonkeyPatch, diff --git a/veadk/models/retrying_lite_llm.py b/veadk/models/retrying_lite_llm.py index cec14d234..3a7b3b510 100644 --- a/veadk/models/retrying_lite_llm.py +++ b/veadk/models/retrying_lite_llm.py @@ -97,7 +97,13 @@ async def generate_content_async( llm_request: LlmRequest, stream: bool = False, ) -> AsyncGenerator[LlmResponse, None]: - retry_request = copy.deepcopy(llm_request) + # Tools are live runtime objects and may own sessions, locks or Futures. + # Preserve their identity while isolating request data and the tool map + # from mutations made by the first attempt. + retry_request = copy.deepcopy( + llm_request, + memo={id(tool): tool for tool in llm_request.tools_dict.values()}, + ) emitted = False try: self._refresh_fallbacks()