Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions complete.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
print("completed")
3 changes: 3 additions & 0 deletions finish_step.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
def plan_step_complete():
print("Moving on")
plan_step_complete()
121 changes: 121 additions & 0 deletions models/tests/test_public_policy_shock_plausibility.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,13 @@
from models.primarycare_model.calibration.public_policy_shock_plausibility import (
REQUIRED_NUMERIC_COMPARISON_COLUMNS,
NumericComparisonContract,
_numeric_comparison_contract,
_numeric_comparison_readiness,
OBSERVED_DELTA_TOLERANCE,
_numeric_value,
_required_text,
_direction_from_delta,
NumericComparisonReadiness,
build_public_policy_shock_evidence,
main,
policy_shock_gate_blockers,
Expand Down Expand Up @@ -211,3 +217,118 @@ def test_passed_numeric_comparison_requires_direction_agreement() -> None:
assert readiness.status == "artifact_invalid"
assert readiness.rows_checked == 1
assert any("comparison_result=passed requires observed_direction to match modelled_direction" in issue for issue in readiness.issues)

def test_numeric_comparison_readiness_to_json_dict() -> None:
readiness = NumericComparisonReadiness(
status="artifact_invalid",
artifact_path="test",
rows_checked=0,
issues=("issue 1", "issue 2")
)
payload = readiness.to_json_dict()
assert payload["status"] == "artifact_invalid"
assert payload["issues"] == ["issue 1", "issue 2"]

def test_required_text() -> None:
assert _required_text(None) is None
assert _required_text(" ") is None
assert _required_text(" test ") == "test"

def test_direction_from_delta() -> None:
assert _direction_from_delta(OBSERVED_DELTA_TOLERANCE + 0.1) == "increase"
assert _direction_from_delta(-(OBSERVED_DELTA_TOLERANCE + 0.1)) == "decrease"
assert _direction_from_delta(0.0) == "no_change"

def test_numeric_comparison_contract_default() -> None:
contract = _numeric_comparison_contract({})
assert contract.required_columns == REQUIRED_NUMERIC_COMPARISON_COLUMNS
assert "Public numeric pre/post comparison requires" in contract.readiness_rule

def test_numeric_comparison_readiness_missing_artifact() -> None:
contract = NumericComparisonContract(
required_columns=REQUIRED_NUMERIC_COMPARISON_COLUMNS,
readiness_rule="test",
pass_rule="test"
)
readiness = _numeric_comparison_readiness(shock_id="shock_id", comparison_artifact="does_not_exist.csv", contract=contract)
assert readiness.status == "artifact_missing"

def test_numeric_comparison_readiness_no_rows() -> None:
import tempfile
import csv
with tempfile.NamedTemporaryFile(mode="w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=REQUIRED_NUMERIC_COMPARISON_COLUMNS)
writer.writeheader()
f.flush()

contract = NumericComparisonContract(
required_columns=REQUIRED_NUMERIC_COMPARISON_COLUMNS,
readiness_rule="test",
pass_rule="test"
)
readiness = _numeric_comparison_readiness(shock_id="shock_id", comparison_artifact=f.name, contract=contract)
assert readiness.status == "artifact_invalid"
assert any("No numeric comparison rows found" in issue for issue in readiness.issues)

def test_numeric_comparison_readiness_all_validation_issues() -> None:
import tempfile
import csv
with tempfile.NamedTemporaryFile(mode="w", newline="") as f:
fields = REQUIRED_NUMERIC_COMPARISON_COLUMNS + ("shock_id", "source_row_index", "source_column_label", "claim_boundary", "source_table_index_pre", "source_table_index_post", "metric_id", "pre_period", "post_period", "pre_value", "post_value", "observed_delta", "observed_direction", "modelled_direction", "comparison_result")
writer = csv.DictWriter(f, fieldnames=fields)
writer.writeheader()

# Write a row missing everything to trigger all issues
writer.writerow({
"shock_id": "shock_id",
"metric_id": " ",
"pre_period": "",
"post_period": " ",
"pre_value": "not_num",
"post_value": "not_num",
"observed_delta": "not_num",
"observed_direction": "invalid",
"modelled_direction": "invalid",
"comparison_result": "invalid",
"source_table_index_pre": "",
"source_table_index_post": "",
"source_row_index": "",
"source_column_label": "",
"claim_boundary": ""
})
f.flush()

contract = NumericComparisonContract(
required_columns=REQUIRED_NUMERIC_COMPARISON_COLUMNS,
readiness_rule="test",
pass_rule="test"
)
readiness = _numeric_comparison_readiness(shock_id="shock_id", comparison_artifact=f.name, contract=contract)
assert readiness.status == "artifact_invalid"
issues_str = " ".join(readiness.issues)
assert "metric_id is required" in issues_str
assert "pre_period is required" in issues_str
assert "post_period is required" in issues_str
assert "pre_value is not numeric" in issues_str
assert "post_value is not numeric" in issues_str
assert "observed_delta is not numeric" in issues_str
assert "observed_direction must be one of" in issues_str
assert "modelled_direction must be one of" in issues_str
assert "comparison_result must be one of" in issues_str

def test_numeric_value_parsing_edge_cases() -> None:
# Test valid numeric values
assert _numeric_value("123.45") == 123.45
assert _numeric_value("$1,234.56") == 1234.56

# Test empty or None values
assert _numeric_value(None) is None
assert _numeric_value("") is None
assert _numeric_value(" ") is None
assert _numeric_value("$") is None
assert _numeric_value(",") is None
assert _numeric_value("$,") is None

# Test invalid string that triggers ValueError
assert _numeric_value("not-a-number") is None
assert _numeric_value("12.34.56") is None
Loading