|
2 | 2 |
|
3 | 3 | from __future__ import annotations |
4 | 4 |
|
5 | | -from dataclasses import dataclass |
6 | | -from unittest.mock import AsyncMock, MagicMock |
7 | | - |
8 | | -import pytest |
| 5 | +from unittest.mock import AsyncMock |
9 | 6 |
|
| 7 | +from tests.sdk.conftest import MockScenarioView, MockScenarioRunView |
10 | 8 | from runloop_api_client.sdk import AsyncScenario |
11 | 9 |
|
12 | 10 |
|
13 | | -@dataclass |
14 | | -class MockScenarioView: |
15 | | - """Mock ScenarioView for testing.""" |
16 | | - |
17 | | - id: str = "scn_123" |
18 | | - name: str = "test-scenario" |
19 | | - metadata: dict = None |
20 | | - |
21 | | - def __post_init__(self): |
22 | | - if self.metadata is None: |
23 | | - self.metadata = {} |
24 | | - |
25 | | - |
26 | | -@dataclass |
27 | | -class MockScenarioRunView: |
28 | | - """Mock ScenarioRunView for testing.""" |
29 | | - |
30 | | - id: str = "run_123" |
31 | | - devbox_id: str = "dev_123" |
32 | | - scenario_id: str = "scn_123" |
33 | | - state: str = "running" |
34 | | - metadata: dict = None |
35 | | - |
36 | | - def __post_init__(self): |
37 | | - if self.metadata is None: |
38 | | - self.metadata = {} |
39 | | - |
40 | | - |
41 | | -@pytest.fixture |
42 | | -def mock_async_client() -> MagicMock: |
43 | | - """Create a mock AsyncRunloop client with proper async returns.""" |
44 | | - client = MagicMock() |
45 | | - # Set up scenarios resource |
46 | | - client.scenarios = MagicMock() |
47 | | - client.scenarios.retrieve = AsyncMock() |
48 | | - client.scenarios.update = AsyncMock() |
49 | | - client.scenarios.start_run = AsyncMock() |
50 | | - client.scenarios.start_run_and_await_env_ready = AsyncMock() |
51 | | - return client |
52 | | - |
53 | | - |
54 | | -@pytest.fixture |
55 | | -def scenario_view() -> MockScenarioView: |
56 | | - """Create a mock ScenarioView.""" |
57 | | - return MockScenarioView() |
58 | | - |
59 | | - |
60 | | -@pytest.fixture |
61 | | -def scenario_run_view() -> MockScenarioRunView: |
62 | | - """Create a mock ScenarioRunView.""" |
63 | | - return MockScenarioRunView() |
64 | | - |
65 | | - |
66 | 11 | class TestAsyncScenario: |
67 | 12 | """Tests for AsyncScenario class.""" |
68 | 13 |
|
69 | | - def test_init(self, mock_async_client: MagicMock) -> None: |
| 14 | + def test_init(self, mock_async_client: AsyncMock) -> None: |
70 | 15 | """Test AsyncScenario initialization.""" |
71 | 16 | scenario = AsyncScenario(mock_async_client, "scn_123") |
72 | 17 | assert scenario.id == "scn_123" |
73 | 18 |
|
74 | | - def test_repr(self, mock_async_client: MagicMock) -> None: |
| 19 | + def test_repr(self, mock_async_client: AsyncMock) -> None: |
75 | 20 | """Test AsyncScenario string representation.""" |
76 | 21 | scenario = AsyncScenario(mock_async_client, "scn_123") |
77 | 22 | assert repr(scenario) == "<AsyncScenario id='scn_123'>" |
78 | 23 |
|
79 | | - async def test_get_info(self, mock_async_client: MagicMock, scenario_view: MockScenarioView) -> None: |
| 24 | + async def test_get_info(self, mock_async_client: AsyncMock, scenario_view: MockScenarioView) -> None: |
80 | 25 | """Test get_info method.""" |
81 | | - mock_async_client.scenarios.retrieve.return_value = scenario_view |
| 26 | + mock_async_client.scenarios.retrieve = AsyncMock(return_value=scenario_view) |
82 | 27 |
|
83 | 28 | scenario = AsyncScenario(mock_async_client, "scn_123") |
84 | 29 | result = await scenario.get_info() |
85 | 30 |
|
86 | 31 | assert result == scenario_view |
87 | | - mock_async_client.scenarios.retrieve.assert_called_once_with("scn_123") |
| 32 | + mock_async_client.scenarios.retrieve.assert_awaited_once_with("scn_123") |
88 | 33 |
|
89 | | - async def test_update(self, mock_async_client: MagicMock, scenario_view: MockScenarioView) -> None: |
| 34 | + async def test_update(self, mock_async_client: AsyncMock, scenario_view: MockScenarioView) -> None: |
90 | 35 | """Test update method.""" |
91 | | - mock_async_client.scenarios.update.return_value = scenario_view |
| 36 | + mock_async_client.scenarios.update = AsyncMock(return_value=scenario_view) |
92 | 37 |
|
93 | 38 | scenario = AsyncScenario(mock_async_client, "scn_123") |
94 | 39 | result = await scenario.update(name="new-name", metadata={"key": "value"}) |
95 | 40 |
|
96 | 41 | assert result == scenario_view |
97 | | - mock_async_client.scenarios.update.assert_called_once_with( |
| 42 | + mock_async_client.scenarios.update.assert_awaited_once_with( |
98 | 43 | "scn_123", |
99 | 44 | name="new-name", |
100 | 45 | metadata={"key": "value"}, |
101 | 46 | ) |
102 | 47 |
|
103 | | - async def test_run(self, mock_async_client: MagicMock, scenario_run_view: MockScenarioRunView) -> None: |
| 48 | + async def test_run(self, mock_async_client: AsyncMock, scenario_run_view: MockScenarioRunView) -> None: |
104 | 49 | """Test run method returns AsyncScenarioRun wrapper.""" |
105 | | - mock_async_client.scenarios.start_run.return_value = scenario_run_view |
| 50 | + mock_async_client.scenarios.start_run = AsyncMock(return_value=scenario_run_view) |
106 | 51 |
|
107 | 52 | scenario = AsyncScenario(mock_async_client, "scn_123") |
108 | 53 | run = await scenario.run(run_name="test-run") |
109 | 54 |
|
110 | 55 | assert run.id == "run_123" |
111 | 56 | assert run.devbox_id == "dev_123" |
112 | | - mock_async_client.scenarios.start_run.assert_called_once_with( |
| 57 | + mock_async_client.scenarios.start_run.assert_awaited_once_with( |
113 | 58 | scenario_id="scn_123", |
114 | 59 | run_name="test-run", |
115 | 60 | ) |
116 | 61 |
|
117 | 62 | async def test_run_and_await_env_ready( |
118 | | - self, mock_async_client: MagicMock, scenario_run_view: MockScenarioRunView |
| 63 | + self, mock_async_client: AsyncMock, scenario_run_view: MockScenarioRunView |
119 | 64 | ) -> None: |
120 | 65 | """Test run_and_await_env_ready method.""" |
121 | | - mock_async_client.scenarios.start_run_and_await_env_ready.return_value = scenario_run_view |
| 66 | + mock_async_client.scenarios.start_run_and_await_env_ready = AsyncMock(return_value=scenario_run_view) |
122 | 67 |
|
123 | 68 | scenario = AsyncScenario(mock_async_client, "scn_123") |
124 | 69 | run = await scenario.run_and_await_env_ready(run_name="test-run") |
125 | 70 |
|
126 | 71 | assert run.id == "run_123" |
127 | 72 | assert run.devbox_id == "dev_123" |
128 | | - mock_async_client.scenarios.start_run_and_await_env_ready.assert_called_once_with( |
| 73 | + mock_async_client.scenarios.start_run_and_await_env_ready.assert_awaited_once_with( |
129 | 74 | scenario_id="scn_123", |
130 | 75 | run_name="test-run", |
131 | 76 | ) |
0 commit comments