|
25 | 25 | from google.genai import types |
26 | 26 | from google.genai.types import Content |
27 | 27 | from google.genai.types import Part |
| 28 | +from openai import AsyncOpenAI |
28 | 29 | import pytest |
29 | 30 |
|
30 | 31 |
|
@@ -465,3 +466,46 @@ async def mock_create(*args, **kwargs): |
465 | 466 | ] |
466 | 467 |
|
467 | 468 | assert responses[0].usage_metadata.cached_content_token_count is None |
| 469 | + |
| 470 | + |
| 471 | +@pytest.mark.asyncio |
| 472 | +async def test_generate_content_async_routes_through_provided_client(): |
| 473 | + """Requests reach the pre-configured client, not a default one.""" |
| 474 | + client = AsyncOpenAI(base_url="https://compatible.example/v1", api_key="k") |
| 475 | + openai_llm = OpenAILlm(model="my-model", client=client) |
| 476 | + llm_request = LlmRequest( |
| 477 | + model="my-model", |
| 478 | + contents=[Content(role="user", parts=[Part.from_text(text="Hello")])], |
| 479 | + ) |
| 480 | + |
| 481 | + mock_response = mock.MagicMock() |
| 482 | + mock_choice = mock.MagicMock() |
| 483 | + mock_message = mock.MagicMock() |
| 484 | + mock_message.content = "Hello there!" |
| 485 | + mock_message.tool_calls = None |
| 486 | + mock_choice.message = mock_message |
| 487 | + mock_response.choices = [mock_choice] |
| 488 | + mock_response.usage.prompt_tokens = 10 |
| 489 | + mock_response.usage.completion_tokens = 5 |
| 490 | + mock_response.usage.total_tokens = 15 |
| 491 | + mock_response.usage.prompt_tokens_details = None |
| 492 | + |
| 493 | + async def mock_create(*args, **kwargs): |
| 494 | + return mock_response |
| 495 | + |
| 496 | + with mock.patch.object( |
| 497 | + client.chat.completions, "create", side_effect=mock_create |
| 498 | + ) as mock_client_create: |
| 499 | + with mock.patch( |
| 500 | + "google.adk.labs.openai._openai_llm.AsyncOpenAI" |
| 501 | + ) as mock_client_class: |
| 502 | + responses = [ |
| 503 | + resp |
| 504 | + async for resp in openai_llm.generate_content_async( |
| 505 | + llm_request, stream=False |
| 506 | + ) |
| 507 | + ] |
| 508 | + |
| 509 | + mock_client_class.assert_not_called() |
| 510 | + mock_client_create.assert_called_once() |
| 511 | + assert responses[0].content.parts[0].text == "Hello there!" |
0 commit comments