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
44 changes: 44 additions & 0 deletions tests/models/test_retrying_lite_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down
8 changes: 7 additions & 1 deletion veadk/models/retrying_lite_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading