Skip to content

Commit 66a363d

Browse files
committed
update generate_data tests
1 parent 5026289 commit 66a363d

2 files changed

Lines changed: 17 additions & 25 deletions

File tree

tests/test_covid_wf/test_generate_data.py

Lines changed: 16 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -9,8 +9,11 @@
99
)
1010
from cfa.scenarios.dataops.workflows.covid.generate_data import (
1111
generate_hospitalization_data,
12-
generate_vaccination_data,
12+
generate_vaccination_data
1313
)
14+
from cfa.scenarios.dataops.datasets.schemas.hospitalization import(
15+
tf_synth_data as tf_synth_hosp_data
16+
)
1417

1518

1619
@patch("cfa.scenarios.dataops.workflows.covid.generate_data.get_data")
@@ -32,29 +35,17 @@ def test_generate_vaccination_data_file(mock_get_data):
3235

3336

3437
@patch("cfa.scenarios.dataops.workflows.covid.generate_data.get_data")
35-
@patch("cfa.scenarios.dataops.workflows.covid.generate_data.requests.get")
36-
def test_generate_hospitalization_data_file(mock_requests_get, mock_get_data):
37-
# Mock region_id DataFrame
38-
mock_get_data.return_value = pd.DataFrame(
39-
{"stusps": ["US", "CA"], "stname": ["United States", "California"]}
40-
)
41-
42-
# Mock CDC API response
43-
fake_api_response = [
44-
{
45-
"week_end_date": "2024-06-01T00:00:00.000",
46-
"jurisdiction": "USA",
47-
"total_admissions_all_covid_confirmed": 100,
48-
},
49-
{
50-
"week_end_date": "2024-06-01T00:00:00.000",
51-
"jurisdiction": "CA",
52-
"total_admissions_all_covid_confirmed": 50,
53-
},
54-
]
55-
mock_requests_get.return_value = MagicMock(
56-
text=pd.io.json.dumps(fake_api_response)
57-
)
38+
def test_generate_hospitalization_data_file(mock_get_data):
39+
# Mock transformed hospitalization data
40+
mock_hosp = tf_synth_hosp_data.copy()
41+
# Mock region_id data
42+
mock_region = pd.DataFrame({
43+
"stusps": ["CA", "TX", "NY", "FL", "IL"],
44+
"stname": ["California" , "Texas", "New York", "Florida", "Illinois"]
45+
})
46+
47+
# get_data is called twice: first for hospitalization, then for region_id
48+
mock_get_data.side_effect = [mock_hosp, mock_region]
5849

5950
with tempfile.TemporaryDirectory() as tmpdir:
6051
df = generate_hospitalization_data(tmpdir, blob=False)
@@ -64,4 +55,4 @@ def test_generate_hospitalization_data_file(mock_requests_get, mock_get_data):
6455
# Check DataFrame structure
6556
assert isinstance(df, pd.DataFrame)
6657
assert set(["date", "state", "total"]).issubset(df.columns)
67-
assert not df.empty
58+
assert not df.empty

tests/test_datasets_catalog.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ def test_list_datasets():
3838
"donor_seroprevalence_2022",
3939
"fips_to_name_improved",
4040
"fips_to_name",
41+
"hospitalization",
4142
"sars_cov2_proportions",
4243
"seroprevalence",
4344
"omicron_variant_regions",

0 commit comments

Comments
 (0)