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
93 changes: 93 additions & 0 deletions src/summarize_deaths.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
import pandas as pd


def summarize_deaths(*args, by: list[str] | str, **kwargs) -> pd.DataFrame:
"""
Summarize deaths by one or more grouping columns.

Groups a CDC mortality dataset by the specified column(s) and returns a
DataFrame with death counts for each group, sorted in descending order.
Optionally accepts multiple named datasets to compare across years.

Parameters
----------
*args : pd.DataFrame
A single unnamed DataFrame for single-year usage.
by : str or list of str
Column name(s) to group by, e.g. "sex" or ["sex", "race_recode3"].
**kwargs : pd.DataFrame
Named DataFrames for multi-year comparison,
e.g. mort1969=mort1969, mort1970=mort1970.

Returns
-------
pd.DataFrame
A DataFrame with the grouping columns, a `year` column (when multiple
datasets are provided), and an `n` column containing death counts,
sorted descending by `n`.

Examples
--------
# Single year
summarize_deaths(mort1969, by="sex")

# Multiple years
summarize_deaths(mort1969=mort1969, mort1970=mort1970, by="sex")
"""
if isinstance(by, str):
by = [by]

if len(args) == 0 and len(kwargs) == 0:
raise ValueError("At least one DataFrame must be provided.")

if len(args) > 0 and len(kwargs) > 0:
raise ValueError(
"Provide either a single positional DataFrame or named DataFrames, not both."
)

# Single dataset — original behavior
if len(args) == 1:
df = args[0]
missing_cols = [col for col in by if col not in df.columns]
if missing_cols:
raise ValueError(
f"The following columns were not found in the data: {', '.join(missing_cols)}"
)

return (
df.groupby(by)
.size()
.reset_index(name="n")
.sort_values("n", ascending=False)
.reset_index(drop=True)
)

# Multiple datasets — compare across years
if len(args) > 1:
raise ValueError(
"When providing multiple datasets, all must be named, "
"e.g. summarize_deaths(mort1969=mort1969, mort1970=mort1970, by='sex')."
)

frames = []
for name, df in kwargs.items():
missing_cols = [col for col in by if col not in df.columns]
if missing_cols:
raise ValueError(
f"The following columns were not found in '{name}': {', '.join(missing_cols)}"
)

summary = (
df.groupby(by)
.size()
.reset_index(name="n")
)
summary.insert(0, "year", name)
frames.append(summary)

combined = pd.concat(frames, ignore_index=True)
return (
combined
.sort_values(["year", "n"], ascending=[True, False])
.reset_index(drop=True)
)
138 changes: 138 additions & 0 deletions tests/test_summarize_deaths.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
import pytest
import pandas as pd
from summarize_deaths import summarize_deaths


# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------

@pytest.fixture
def mort1969():
return pd.DataFrame({
"sex": ["M", "M", "F", "F", "F", "M"],
"race_recode3": ["White", "Black", "White", "White", "Black", "White"],
"age": [45, 60, 32, 78, 55, 90],
})


@pytest.fixture
def mort1970():
return pd.DataFrame({
"sex": ["M", "F", "F", "M", "M"],
"race_recode3": ["White", "Black", "White", "Black", "White"],
"age": [50, 40, 88, 70, 33],
})


# ---------------------------------------------------------------------------
# Single dataset
# ---------------------------------------------------------------------------

class TestSingleDataset:

def test_group_by_single_column(self, mort1969):
result = summarize_deaths(mort1969, by="sex")
assert list(result.columns) == ["sex", "n"]
assert result["n"].iloc[0] >= result["n"].iloc[1], "Should be sorted descending"
assert set(result["sex"]) == {"M", "F"}
assert result["n"].sum() == len(mort1969)

def test_group_by_multiple_columns(self, mort1969):
result = summarize_deaths(mort1969, by=["sex", "race_recode3"])
assert "sex" in result.columns
assert "race_recode3" in result.columns
assert "n" in result.columns
assert result["n"].sum() == len(mort1969)

def test_by_as_string_and_list_equivalent(self, mort1969):
result_str = summarize_deaths(mort1969, by="sex")
result_list = summarize_deaths(mort1969, by=["sex"])
pd.testing.assert_frame_equal(result_str, result_list)

def test_sorted_descending(self, mort1969):
result = summarize_deaths(mort1969, by="sex")
assert result["n"].is_monotonic_decreasing

def test_counts_are_correct(self, mort1969):
result = summarize_deaths(mort1969, by="sex")
counts = dict(zip(result["sex"], result["n"]))
assert counts["M"] == 3
assert counts["F"] == 3

def test_single_group_value(self):
df = pd.DataFrame({"sex": ["M", "M", "M"]})
result = summarize_deaths(df, by="sex")
assert len(result) == 1
assert result["n"].iloc[0] == 3


# ---------------------------------------------------------------------------
# Multiple datasets
# ---------------------------------------------------------------------------

class TestMultipleDatasets:

def test_returns_year_column(self, mort1969, mort1970):
result = summarize_deaths(mort1969=mort1969, mort1970=mort1970, by="sex")
assert "year" in result.columns
assert set(result["year"]) == {"mort1969", "mort1970"}

def test_year_is_first_column(self, mort1969, mort1970):
result = summarize_deaths(mort1969=mort1969, mort1970=mort1970, by="sex")
assert result.columns[0] == "year"

def test_total_counts_correct(self, mort1969, mort1970):
result = summarize_deaths(mort1969=mort1969, mort1970=mort1970, by="sex")
assert result["n"].sum() == len(mort1969) + len(mort1970)

def test_sorted_by_year_then_desc_n(self, mort1969, mort1970):
result = summarize_deaths(mort1969=mort1969, mort1970=mort1970, by="sex")
for year, group in result.groupby("year", sort=False):
assert group["n"].is_monotonic_decreasing, f"Year '{year}' not sorted descending by n"

def test_multi_column_grouping(self, mort1969, mort1970):
result = summarize_deaths(
mort1969=mort1969, mort1970=mort1970, by=["sex", "race_recode3"]
)
assert "year" in result.columns
assert "sex" in result.columns
assert "race_recode3" in result.columns
assert "n" in result.columns


# ---------------------------------------------------------------------------
# Error handling
# ---------------------------------------------------------------------------

class TestErrors:

def test_no_datasets_raises(self):
with pytest.raises(ValueError, match="At least one DataFrame"):
summarize_deaths(by="sex")

def test_missing_column_single_dataset(self, mort1969):
with pytest.raises(ValueError, match="not found in the data"):
summarize_deaths(mort1969, by="nonexistent_col")

def test_missing_column_multi_dataset(self, mort1969, mort1970):
with pytest.raises(ValueError, match="not found in 'mort1970'"):
summarize_deaths(mort1969=mort1969, mort1970=mort1970, by="nonexistent_col")

def test_mixed_positional_and_named_raises(self, mort1969, mort1970):
with pytest.raises(ValueError, match="not both"):
summarize_deaths(mort1969, mort1970=mort1970, by="sex")

def test_multiple_positional_raises(self, mort1969, mort1970):
with pytest.raises(ValueError, match="must be named"):
summarize_deaths(mort1969, mort1970, by="sex")

def test_missing_one_of_multiple_by_columns(self, mort1969):
with pytest.raises(ValueError, match="not found in the data"):
summarize_deaths(mort1969, by=["sex", "does_not_exist"])

def test_empty_dataframe(self):
df = pd.DataFrame({"sex": pd.Series([], dtype=str)})
result = summarize_deaths(df, by="sex")
assert result.empty
assert list(result.columns) == ["sex", "n"]
Loading