From 2a25a28e3fe93e57c998f71a052b5afbfab567af Mon Sep 17 00:00:00 2001 From: Matthew Spah Date: Wed, 23 Sep 2026 21:29:16 -0700 Subject: [PATCH] refactor: share analytics schema consistency validators --- app/schemas/analytics.py | 229 +++++++++------------ tests/test_analytics_schema_consistency.py | 211 +++++++++++++++++++ 2 files changed, 304 insertions(+), 136 deletions(-) create mode 100644 tests/test_analytics_schema_consistency.py diff --git a/app/schemas/analytics.py b/app/schemas/analytics.py index 5d3b8eb..c04c21f 100644 --- a/app/schemas/analytics.py +++ b/app/schemas/analytics.py @@ -20,6 +20,45 @@ from app.schemas.games import HomeAway +def _validate_prior_window_pair( + prior_value: float | None, + change: float | None, + *, + value_field_name: str, +) -> None: + """Require a prior-window value and its change to be known together.""" + if (prior_value is not None) != (change is not None): + raise ValueError( + f"{value_field_name} and change_vs_prior_window must both be " + "present or both be None" + ) + + +def _validate_summary_point_count(games_played: int, point_count: int) -> None: + """Require the summary to count exactly the games represented by points.""" + if games_played != point_count: + raise ValueError("summary.games_played must equal the number of chart points") + + +def _validate_teams_within_records( + teams_represented: int, team_game_records: int +) -> None: + """Every represented team must have at least one counted record.""" + if teams_represented > team_game_records: + raise ValueError( + f"teams_represented ({teams_represented}) cannot exceed " + f"team_game_records ({team_game_records})" + ) + + +def _validate_league_season(season: int, league_season: int) -> None: + """Require a team comparison and its league context to describe one season.""" + if season != league_season: + raise ValueError( + f"season ({season}) must match the league context season ({league_season})" + ) + + class TeamHitsPoint(BaseModel): """One completed game plotted on the team hits chart.""" @@ -77,13 +116,11 @@ class TeamHitsSummary(BaseModel): @model_validator(mode="after") def _prior_window_fields_agree(self) -> TeamHitsSummary: - has_prior = self.prior_window_average is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_average and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_average, + self.change_vs_prior_window, + value_field_name="prior_window_average", + ) return self @@ -103,10 +140,7 @@ class TeamHitsAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamHitsAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) return self @property @@ -180,13 +214,11 @@ class TeamStrikeoutsSummary(BaseModel): @model_validator(mode="after") def _prior_window_fields_agree(self) -> TeamStrikeoutsSummary: - has_prior = self.prior_window_average is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_average and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_average, + self.change_vs_prior_window, + value_field_name="prior_window_average", + ) return self @@ -206,10 +238,7 @@ class TeamStrikeoutsAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamStrikeoutsAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) return self @property @@ -259,11 +288,7 @@ def _hits_per_game_matches_the_totals(self) -> LeagueHitsContext: f"hits_per_game ({self.hits_per_game}) must equal total_hits / " f"team_game_records ({expected})" ) - if self.teams_represented > self.team_game_records: - raise ValueError( - f"teams_represented ({self.teams_represented}) cannot exceed " - f"team_game_records ({self.team_game_records})" - ) + _validate_teams_within_records(self.teams_represented, self.team_game_records) return self @@ -294,11 +319,7 @@ class TeamHitsLeagueComparison(BaseModel): @model_validator(mode="after") def _comparison_is_internally_consistent(self) -> TeamHitsLeagueComparison: - if self.season != self.league.season: - raise ValueError( - f"season ({self.season}) must match the league context season " - f"({self.league.season})" - ) + _validate_league_season(self.season, self.league.season) expected = self.team_hits_per_game - self.league.hits_per_game if not isclose(self.difference_vs_mlb, expected, rel_tol=1e-9, abs_tol=1e-9): raise ValueError( @@ -357,11 +378,7 @@ def _strikeouts_per_game_matches_the_totals(self) -> LeagueStrikeoutsContext: f"strikeouts_per_game ({self.strikeouts_per_game}) must equal " f"total_strikeouts / team_game_records ({expected})" ) - if self.teams_represented > self.team_game_records: - raise ValueError( - f"teams_represented ({self.teams_represented}) cannot exceed " - f"team_game_records ({self.team_game_records})" - ) + _validate_teams_within_records(self.teams_represented, self.team_game_records) return self @@ -397,11 +414,7 @@ class TeamStrikeoutsLeagueComparison(BaseModel): @model_validator(mode="after") def _comparison_is_internally_consistent(self) -> TeamStrikeoutsLeagueComparison: - if self.season != self.league.season: - raise ValueError( - f"season ({self.season}) must match the league context season " - f"({self.league.season})" - ) + _validate_league_season(self.season, self.league.season) expected = self.team_strikeouts_per_game - self.league.strikeouts_per_game if not isclose(self.difference_vs_mlb, expected, rel_tol=1e-9, abs_tol=1e-9): raise ValueError( @@ -504,10 +517,7 @@ def _comparison_is_internally_consistent( ) -> TeamHittingComparisonAnalysis: if not isclose(self.baseline_index, 100.0, rel_tol=0.0, abs_tol=1e-12): raise ValueError("baseline_index must equal 100") - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) recent = self.points[-1] if not isclose( @@ -599,13 +609,11 @@ class TeamRunsSummary(BaseModel): @model_validator(mode="after") def _prior_window_fields_agree(self) -> TeamRunsSummary: - has_prior = self.prior_window_average is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_average and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_average, + self.change_vs_prior_window, + value_field_name="prior_window_average", + ) return self @@ -625,10 +633,7 @@ class TeamRunsAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamRunsAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) return self @property @@ -681,11 +686,7 @@ def _runs_per_game_matches_the_totals(self) -> LeagueRunsContext: f"runs_per_game ({self.runs_per_game}) must equal total_runs / " f"team_game_records ({expected})" ) - if self.teams_represented > self.team_game_records: - raise ValueError( - f"teams_represented ({self.teams_represented}) cannot exceed " - f"team_game_records ({self.team_game_records})" - ) + _validate_teams_within_records(self.teams_represented, self.team_game_records) return self @@ -717,11 +718,7 @@ class TeamRunsLeagueComparison(BaseModel): @model_validator(mode="after") def _comparison_is_internally_consistent(self) -> TeamRunsLeagueComparison: - if self.season != self.league.season: - raise ValueError( - f"season ({self.season}) must match the league context season " - f"({self.league.season})" - ) + _validate_league_season(self.season, self.league.season) expected = self.team_runs_per_game - self.league.runs_per_game if not isclose(self.difference_vs_mlb, expected, rel_tol=1e-9, abs_tol=1e-9): raise ValueError( @@ -796,13 +793,11 @@ class TeamBaserunnersSummary(BaseModel): @model_validator(mode="after") def _prior_window_fields_agree(self) -> TeamBaserunnersSummary: - has_prior = self.prior_window_average is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_average and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_average, + self.change_vs_prior_window, + value_field_name="prior_window_average", + ) return self @@ -822,10 +817,7 @@ class TeamBaserunnersAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamBaserunnersAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) return self @property @@ -883,11 +875,7 @@ def _baserunners_per_game_matches_the_totals(self) -> LeagueBaserunnersContext: f"baserunners_per_game ({self.baserunners_per_game}) must equal " f"total_baserunners / team_game_records ({expected})" ) - if self.teams_represented > self.team_game_records: - raise ValueError( - f"teams_represented ({self.teams_represented}) cannot exceed " - f"team_game_records ({self.team_game_records})" - ) + _validate_teams_within_records(self.teams_represented, self.team_game_records) return self @@ -923,11 +911,7 @@ class TeamBaserunnersLeagueComparison(BaseModel): @model_validator(mode="after") def _comparison_is_internally_consistent(self) -> TeamBaserunnersLeagueComparison: - if self.season != self.league.season: - raise ValueError( - f"season ({self.season}) must match the league context season " - f"({self.league.season})" - ) + _validate_league_season(self.season, self.league.season) expected = self.team_baserunners_per_game - self.league.baserunners_per_game if not isclose(self.difference_vs_mlb, expected, rel_tol=1e-9, abs_tol=1e-9): raise ValueError( @@ -1130,13 +1114,11 @@ def _totals_and_windows_agree(self) -> TeamRunDifferentialSummary: f"season_average ({self.season_average}) must equal " f"total_run_differential / games_played ({expected_average})" ) - has_prior = self.prior_window_average is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_average and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_average, + self.change_vs_prior_window, + value_field_name="prior_window_average", + ) return self @@ -1157,10 +1139,7 @@ class TeamRunDifferentialAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamRunDifferentialAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) decided = self.pythagorean.actual_wins + self.pythagorean.actual_losses if decided != len(self.points): raise ValueError( @@ -1371,13 +1350,11 @@ class TeamPitchingSummary(BaseModel): @model_validator(mode="after") def _prior_window_fields_agree(self) -> TeamPitchingSummary: - has_prior = self.prior_window_era is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_era and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_era, + self.change_vs_prior_window, + value_field_name="prior_window_era", + ) return self @@ -1397,10 +1374,7 @@ class TeamPitchingAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamPitchingAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) charted_outs = sum(point.outs for point in self.points) if charted_outs != self.summary.season.outs: raise ValueError( @@ -1476,11 +1450,7 @@ def _rates_match_the_totals(self) -> LeaguePitchingContext: f"era ({self.era}) must equal total_earned_runs * 27 / outs " f"({expected_era})" ) - if self.teams_represented > self.team_game_records: - raise ValueError( - f"teams_represented ({self.teams_represented}) cannot exceed " - f"team_game_records ({self.team_game_records})" - ) + _validate_teams_within_records(self.teams_represented, self.team_game_records) return self @@ -1514,11 +1484,7 @@ class TeamPitchingLeagueComparison(BaseModel): @model_validator(mode="after") def _comparison_is_internally_consistent(self) -> TeamPitchingLeagueComparison: - if self.season != self.league.season: - raise ValueError( - f"season ({self.season}) must match the league context season " - f"({self.league.season})" - ) + _validate_league_season(self.season, self.league.season) for name, team_value, league_value in ( ("era_difference_vs_mlb", self.team_era, self.league.era), ("whip_difference_vs_mlb", self.team_whip, self.league.whip), @@ -1622,13 +1588,11 @@ def _totals_and_windows_agree(self) -> TeamHitsAllowedSummary: f"hits_per_nine ({self.hits_per_nine}) must equal " f"total_hits_allowed * 27 / total_outs ({expected_rate})" ) - has_prior = self.prior_window_average is not None - has_change = self.change_vs_prior_window is not None - if has_prior != has_change: - raise ValueError( - "prior_window_average and change_vs_prior_window must both be " - "present or both be None" - ) + _validate_prior_window_pair( + self.prior_window_average, + self.change_vs_prior_window, + value_field_name="prior_window_average", + ) return self @@ -1648,10 +1612,7 @@ class TeamHitsAllowedAnalysis(BaseModel): @model_validator(mode="after") def _summary_matches_points(self) -> TeamHitsAllowedAnalysis: - if self.summary.games_played != len(self.points): - raise ValueError( - "summary.games_played must equal the number of chart points" - ) + _validate_summary_point_count(self.summary.games_played, len(self.points)) charted = sum(point.hits_allowed for point in self.points) if charted != self.summary.total_hits_allowed: raise ValueError( @@ -1707,11 +1668,7 @@ class TeamHitsAllowedLeagueComparison(BaseModel): @model_validator(mode="after") def _comparison_is_internally_consistent(self) -> TeamHitsAllowedLeagueComparison: - if self.season != self.league.season: - raise ValueError( - f"season ({self.season}) must match the league context season " - f"({self.league.season})" - ) + _validate_league_season(self.season, self.league.season) expected = self.team_hits_allowed_per_game - self.league.hits_per_game if not isclose(self.difference_vs_mlb, expected, rel_tol=1e-9, abs_tol=1e-9): raise ValueError( diff --git a/tests/test_analytics_schema_consistency.py b/tests/test_analytics_schema_consistency.py new file mode 100644 index 0000000..c9d50d9 --- /dev/null +++ b/tests/test_analytics_schema_consistency.py @@ -0,0 +1,211 @@ +"""Characterize shared consistency rules through the public analytics models.""" + +import pytest +from pydantic import BaseModel, ValidationError + +from app.analytics.team_baserunners import build_team_baserunners_analysis +from app.analytics.team_hits_allowed import build_team_hits_allowed_analysis +from app.analytics.team_hitting import build_team_hits_analysis +from app.analytics.team_hitting_comparison import build_team_hitting_comparison_analysis +from app.analytics.team_pitching import build_team_pitching_analysis +from app.analytics.team_run_differential import build_team_run_differential_analysis +from app.analytics.team_runs import build_team_runs_analysis +from app.analytics.team_strikeouts import build_team_strikeouts_analysis +from app.schemas.analytics import ( + LeaguePitchingContext, + TeamBaserunnersLeagueComparison, + TeamHitsAllowedLeagueComparison, + TeamHitsLeagueComparison, + TeamPitchingLeagueComparison, + TeamRunsLeagueComparison, + TeamStrikeoutsLeagueComparison, +) +from tests.factories import ( + make_batting_line, + make_league_baserunners_context, + make_league_hits_context, + make_league_runs_context, + make_league_strikeouts_context, + make_pitching_line, + make_run_result, +) + +# Builders supply valid starting objects; each test revalidates a modified dump +# through the public schema rather than using model_copy (which skips validation). +_batting = [make_batting_line(strikeouts=9, base_on_balls=2, hit_by_pitch=0)] +_pitching = [make_pitching_line()] +_hits = build_team_hits_analysis(_batting) +_strikeouts = build_team_strikeouts_analysis(_batting) +_analyses = [ + _hits, + _strikeouts, + build_team_runs_analysis(_batting), + build_team_baserunners_analysis(_batting), + build_team_pitching_analysis(_pitching), + build_team_run_differential_analysis([make_run_result()]), + build_team_hits_allowed_analysis(_pitching), + build_team_hitting_comparison_analysis( + _hits, + _strikeouts, + make_league_hits_context(), + make_league_strikeouts_context(), + ), +] + + +@pytest.mark.parametrize("analysis", _analyses, ids=lambda model: type(model).__name__) +@pytest.mark.parametrize("point_copies", [1, 2]) +def test_summary_game_count_matches_points( + analysis: BaseModel, point_copies: int +) -> None: + data = analysis.model_dump() + data["points"] = data["points"] * point_copies + if point_copies == 1: + assert type(analysis).model_validate(data) == analysis + else: + with pytest.raises(ValidationError) as caught: + type(analysis).model_validate(data) + assert caught.value.errors()[0]["msg"] == ( + "Value error, summary.games_played must equal the number of chart points" + ) + + +@pytest.mark.parametrize( + "summary", + [analysis.summary for analysis in _analyses[:-1]], + ids=lambda model: type(model).__name__, +) +@pytest.mark.parametrize( + ("prior", "change"), [(None, None), (0.0, 0.0), (0.0, None), (None, 0.0)] +) +def test_prior_window_fields_are_present_or_absent_together( + summary: BaseModel, prior: float | None, change: float | None +) -> None: + data = summary.model_dump() + prior_field = ( + "prior_window_era" if "prior_window_era" in data else "prior_window_average" + ) + data[prior_field] = prior + data["change_vs_prior_window"] = change + if (prior is None) == (change is None): + validated = type(summary).model_validate(data).model_dump() + assert validated[prior_field] == prior + assert validated["change_vs_prior_window"] == change + else: + with pytest.raises(ValidationError) as caught: + type(summary).model_validate(data) + assert caught.value.errors()[0]["msg"] == ( + f"Value error, {prior_field} and change_vs_prior_window must both be " + "present or both be None" + ) + + +_league_pitching = LeaguePitchingContext( + season=2025, + teams_represented=2, + team_game_records=10, + outs=270, + innings_pitched=90.0, + total_earned_runs=30, + era=3.0, + whip=1.0, + strikeouts_per_nine=9.0, + walks_per_nine=2.0, +) +_leagues = [ + make_league_hits_context(), + make_league_strikeouts_context(), + make_league_runs_context(), + make_league_baserunners_context(), + _league_pitching, +] + + +@pytest.mark.parametrize("league", _leagues, ids=lambda model: type(model).__name__) +@pytest.mark.parametrize("teams", [1, 10, 11]) +def test_represented_teams_cannot_exceed_records(league: BaseModel, teams: int) -> None: + data = league.model_dump() + data["teams_represented"] = teams + if teams <= 10: + assert ( + type(league).model_validate(data).model_dump()["teams_represented"] == teams + ) + else: + with pytest.raises(ValidationError) as caught: + type(league).model_validate(data) + assert caught.value.errors()[0]["msg"] == ( + "Value error, teams_represented (11) cannot exceed team_game_records (10)" + ) + + +@pytest.mark.parametrize( + ("model", "fields"), + [ + ( + TeamHitsLeagueComparison, + dict( + team_hits_per_game=8.0, + league=make_league_hits_context(), + difference_vs_mlb=0.0, + ), + ), + ( + TeamStrikeoutsLeagueComparison, + dict( + team_strikeouts_per_game=8.0, + league=make_league_strikeouts_context(), + difference_vs_mlb=0.0, + ), + ), + ( + TeamRunsLeagueComparison, + dict( + team_runs_per_game=4.5, + league=make_league_runs_context(), + difference_vs_mlb=0.0, + ), + ), + ( + TeamBaserunnersLeagueComparison, + dict( + team_baserunners_per_game=10.0, + league=make_league_baserunners_context(), + difference_vs_mlb=0.0, + ), + ), + ( + TeamHitsAllowedLeagueComparison, + dict( + team_hits_allowed_per_game=8.0, + league=make_league_hits_context(), + difference_vs_mlb=0.0, + ), + ), + ( + TeamPitchingLeagueComparison, + dict( + team_era=3.0, + team_whip=1.0, + team_strikeouts_per_nine=9.0, + team_walks_per_nine=2.0, + league=_league_pitching, + era_difference_vs_mlb=0.0, + whip_difference_vs_mlb=0.0, + ), + ), + ], + ids=lambda value: value.__name__ if isinstance(value, type) else None, +) +@pytest.mark.parametrize("season", [2025, 2026]) +def test_comparison_requires_the_league_season( + model: type[BaseModel], fields: dict[str, object], season: int +) -> None: + data = dict(fields, team_id=136, team_name="Seattle Mariners", season=season) + if season == 2025: + assert model.model_validate(data).model_dump()["season"] == season + else: + with pytest.raises(ValidationError) as caught: + model.model_validate(data) + assert caught.value.errors()[0]["msg"] == ( + "Value error, season (2026) must match the league context season (2025)" + )