Skip to content

Commit 3e0910d

Browse files
committed
address type check errors in scenario builder unit tests
1 parent 0ea9baa commit 3e0910d

2 files changed

Lines changed: 22 additions & 18 deletions

File tree

tests/sdk/test_async_scenario_builder.py

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -158,14 +158,16 @@ def test_build_params_with_all_options(self, builder: AsyncScenarioBuilder, mock
158158

159159
assert params["name"] == "test-scenario"
160160
assert params["input_context"]["problem_statement"] == "Fix the bug"
161-
assert params["input_context"]["additional_context"] == {"hint": "line 42"}
162-
assert params["environment_parameters"]["blueprint_id"] == "bp-123"
163-
assert params["environment_parameters"]["working_directory"] == "/app"
164-
assert params["metadata"] == {"team": "infra"}
165-
assert params["reference_output"] == "diff content"
166-
assert params["required_environment_variables"] == ["API_KEY"]
167-
assert params["required_secret_names"] == ["db_pass"]
168-
assert params["validation_type"] == "FORWARD"
161+
assert params["input_context"].get("additional_context") == {"hint": "line 42"}
162+
env_params = params.get("environment_parameters")
163+
assert env_params is not None
164+
assert env_params.get("blueprint_id") == "bp-123"
165+
assert env_params.get("working_directory") == "/app"
166+
assert params.get("metadata") == {"team": "infra"}
167+
assert params.get("reference_output") == "diff content"
168+
assert params.get("required_environment_variables") == ["API_KEY"]
169+
assert params.get("required_secret_names") == ["db_pass"]
170+
assert params.get("validation_type") == "FORWARD"
169171

170172
def test_build_params_normalizes_weights(self, builder: AsyncScenarioBuilder) -> None:
171173
"""Test that _build_params normalizes scorer weights to sum to 1.0."""
@@ -175,7 +177,7 @@ def test_build_params_normalizes_weights(self, builder: AsyncScenarioBuilder) ->
175177
builder.add_bash_script_scorer("scorer3", bash_script="echo 3", weight=3.0)
176178

177179
params = builder._build_params()
178-
scorers = params["scoring_contract"]["scoring_function_parameters"]
180+
scorers = list(params["scoring_contract"]["scoring_function_parameters"])
179181

180182
# Weights 1, 2, 3 should normalize to 1/6, 2/6, 3/6
181183
assert len(scorers) == 3

tests/sdk/test_scenario_builder.py

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -156,14 +156,16 @@ def test_build_params_with_all_options(self, builder: ScenarioBuilder, mock_blue
156156

157157
assert params["name"] == "test-scenario"
158158
assert params["input_context"]["problem_statement"] == "Fix the bug"
159-
assert params["input_context"]["additional_context"] == {"hint": "line 42"}
160-
assert params["environment_parameters"]["blueprint_id"] == "bp-123"
161-
assert params["environment_parameters"]["working_directory"] == "/app"
162-
assert params["metadata"] == {"team": "infra"}
163-
assert params["reference_output"] == "diff content"
164-
assert params["required_environment_variables"] == ["API_KEY"]
165-
assert params["required_secret_names"] == ["db_pass"]
166-
assert params["validation_type"] == "FORWARD"
159+
assert params["input_context"].get("additional_context") == {"hint": "line 42"}
160+
env_params = params.get("environment_parameters")
161+
assert env_params is not None
162+
assert env_params.get("blueprint_id") == "bp-123"
163+
assert env_params.get("working_directory") == "/app"
164+
assert params.get("metadata") == {"team": "infra"}
165+
assert params.get("reference_output") == "diff content"
166+
assert params.get("required_environment_variables") == ["API_KEY"]
167+
assert params.get("required_secret_names") == ["db_pass"]
168+
assert params.get("validation_type") == "FORWARD"
167169

168170
def test_build_params_normalizes_weights(self, builder: ScenarioBuilder) -> None:
169171
"""Test that _build_params normalizes scorer weights to sum to 1.0."""
@@ -173,7 +175,7 @@ def test_build_params_normalizes_weights(self, builder: ScenarioBuilder) -> None
173175
builder.add_bash_script_scorer("scorer3", bash_script="echo 3", weight=3.0)
174176

175177
params = builder._build_params()
176-
scorers = params["scoring_contract"]["scoring_function_parameters"]
178+
scorers = list(params["scoring_contract"]["scoring_function_parameters"])
177179

178180
# Weights 1, 2, 3 should normalize to 1/6, 2/6, 3/6
179181
assert len(scorers) == 3

0 commit comments

Comments
 (0)