Skip to content

Commit 45d0ab3

Browse files
committed
fixed redundancy and formatting in scenario and scenario run unit tests
1 parent ed187e7 commit 45d0ab3

5 files changed

Lines changed: 53 additions & 273 deletions

File tree

tests/sdk/conftest.py

Lines changed: 4 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,8 @@
44

55
import asyncio
66
import threading
7-
from typing import Any
8-
from dataclasses import dataclass
7+
from typing import Any, Dict
8+
from dataclasses import field, dataclass
99
from unittest.mock import Mock, AsyncMock
1010

1111
import httpx
@@ -114,11 +114,7 @@ class MockScenarioView:
114114

115115
id: str = "scn_123"
116116
name: str = "test-scenario"
117-
metadata: dict = None
118-
119-
def __post_init__(self):
120-
if self.metadata is None:
121-
self.metadata = {}
117+
metadata: Dict[str, str] = field(default_factory=dict)
122118

123119

124120
@dataclass
@@ -129,13 +125,9 @@ class MockScenarioRunView:
129125
devbox_id: str = "dev_123"
130126
scenario_id: str = "scn_123"
131127
state: str = "running"
132-
metadata: dict = None
128+
metadata: Dict[str, str] = field(default_factory=dict)
133129
scoring_contract_result: object = None
134130

135-
def __post_init__(self):
136-
if self.metadata is None:
137-
self.metadata = {}
138-
139131

140132
def create_mock_httpx_client(methods: dict[str, Any] | None = None) -> AsyncMock:
141133
"""

tests/sdk/test_async_scenario.py

Lines changed: 16 additions & 71 deletions
Original file line numberDiff line numberDiff line change
@@ -2,130 +2,75 @@
22

33
from __future__ import annotations
44

5-
from dataclasses import dataclass
6-
from unittest.mock import AsyncMock, MagicMock
7-
8-
import pytest
5+
from unittest.mock import AsyncMock
96

7+
from tests.sdk.conftest import MockScenarioView, MockScenarioRunView
108
from runloop_api_client.sdk import AsyncScenario
119

1210

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-
6611
class TestAsyncScenario:
6712
"""Tests for AsyncScenario class."""
6813

69-
def test_init(self, mock_async_client: MagicMock) -> None:
14+
def test_init(self, mock_async_client: AsyncMock) -> None:
7015
"""Test AsyncScenario initialization."""
7116
scenario = AsyncScenario(mock_async_client, "scn_123")
7217
assert scenario.id == "scn_123"
7318

74-
def test_repr(self, mock_async_client: MagicMock) -> None:
19+
def test_repr(self, mock_async_client: AsyncMock) -> None:
7520
"""Test AsyncScenario string representation."""
7621
scenario = AsyncScenario(mock_async_client, "scn_123")
7722
assert repr(scenario) == "<AsyncScenario id='scn_123'>"
7823

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:
8025
"""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)
8227

8328
scenario = AsyncScenario(mock_async_client, "scn_123")
8429
result = await scenario.get_info()
8530

8631
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")
8833

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:
9035
"""Test update method."""
91-
mock_async_client.scenarios.update.return_value = scenario_view
36+
mock_async_client.scenarios.update = AsyncMock(return_value=scenario_view)
9237

9338
scenario = AsyncScenario(mock_async_client, "scn_123")
9439
result = await scenario.update(name="new-name", metadata={"key": "value"})
9540

9641
assert result == scenario_view
97-
mock_async_client.scenarios.update.assert_called_once_with(
42+
mock_async_client.scenarios.update.assert_awaited_once_with(
9843
"scn_123",
9944
name="new-name",
10045
metadata={"key": "value"},
10146
)
10247

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:
10449
"""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)
10651

10752
scenario = AsyncScenario(mock_async_client, "scn_123")
10853
run = await scenario.run(run_name="test-run")
10954

11055
assert run.id == "run_123"
11156
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(
11358
scenario_id="scn_123",
11459
run_name="test-run",
11560
)
11661

11762
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
11964
) -> None:
12065
"""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)
12267

12368
scenario = AsyncScenario(mock_async_client, "scn_123")
12469
run = await scenario.run_and_await_env_ready(run_name="test-run")
12570

12671
assert run.id == "run_123"
12772
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(
12974
scenario_id="scn_123",
13075
run_name="test-run",
13176
)

0 commit comments

Comments
 (0)