|
20 | 20 | from typing import AsyncGenerator |
21 | 21 | from unittest import mock |
22 | 22 |
|
| 23 | +from google.adk import platform as adk_platform |
23 | 24 | from google.adk.agents.context import Context |
24 | 25 | from google.adk.events.event import Event |
25 | 26 | from google.adk.runners import Runner |
|
29 | 30 | from google.adk.workflow import START |
30 | 31 | from google.adk.workflow._errors import NodeTimeoutError |
31 | 32 | from google.adk.workflow._graph import Graph |
32 | | -from google.adk.workflow._node import node |
33 | 33 | from google.adk.workflow._node import Node |
| 34 | +from google.adk.workflow._node import node |
34 | 35 | from google.adk.workflow._node_status import NodeStatus |
35 | 36 | from google.adk.workflow._retry_config import RetryConfig |
36 | 37 | from google.adk.workflow._workflow import Workflow |
@@ -715,20 +716,23 @@ async def test_retry_applies_random_jitter(request: pytest.FixtureRequest): |
715 | 716 | session = await ss.create_session(app_name=agent.name, user_id='u') |
716 | 717 | msg = types.Content(parts=[types.Part(text='start')], role='user') |
717 | 718 |
|
718 | | - with ( |
719 | | - mock.patch('asyncio.sleep', new_callable=mock.AsyncMock) as mock_sleep, |
720 | | - mock.patch('random.uniform', return_value=-1.0) as mock_random, |
721 | | - ): |
722 | | - events = [] |
723 | | - async for event in runner.run_async( |
724 | | - user_id='u', session_id=session.id, new_message=msg |
725 | | - ): |
726 | | - events.append(event) |
727 | | - |
728 | | - # 4.0 + (-1.0) = 3.0 |
729 | | - mock_sleep.assert_any_await(3.0) |
730 | | - # Called with -0.5 * 4.0, 0.5 * 4.0 |
731 | | - mock_random.assert_called_once_with(-2.0, 2.0) |
| 719 | + mock_random = mock.Mock() |
| 720 | + mock_random.uniform = mock.Mock(return_value=-1.0) |
| 721 | + adk_platform.set_random_provider(lambda: mock_random) |
| 722 | + try: |
| 723 | + with mock.patch('asyncio.sleep', new_callable=mock.AsyncMock) as mock_sleep: |
| 724 | + events = [] |
| 725 | + async for event in runner.run_async( |
| 726 | + user_id='u', session_id=session.id, new_message=msg |
| 727 | + ): |
| 728 | + events.append(event) |
| 729 | + |
| 730 | + # 4.0 + (-1.0) = 3.0 |
| 731 | + mock_sleep.assert_any_await(3.0) |
| 732 | + # Called with -0.5 * 4.0, 0.5 * 4.0 |
| 733 | + mock_random.uniform.assert_called_once_with(-2.0, 2.0) |
| 734 | + finally: |
| 735 | + adk_platform.reset_random_provider() |
732 | 736 |
|
733 | 737 | results = simplify_events_with_node(events) |
734 | 738 | filtered_results = [ |
|
0 commit comments