diff --git a/src/datasure/checks/enumerator.py b/src/datasure/checks/enumerator.py deleted file mode 100644 index a5edb116..00000000 --- a/src/datasure/checks/enumerator.py +++ /dev/null @@ -1,2673 +0,0 @@ -"""Enumerator performance analysis module for survey data quality checks. - -This module provides comprehensive enumerator performance tracking with: -- Enumerator overview metrics and statistics -- Productivity tracking over time (daily, weekly, monthly) -- Summary tables with missing data, duration, consent, and outcome analysis -- Statistical analysis across enumerators -- Time-series analysis of enumerator performance -- Configurable settings with Pydantic validation -- Modular, testable architecture -- Polars-based data processing for performance -""" - -from datetime import date as dt_date -from datetime import timedelta -from typing import Literal - -import polars as pl -import streamlit as st -from pydantic import BaseModel, Field, field_validator - -from datasure.checks import missing -from datasure.utils.dataframe_utils import ColumnByType -from datasure.utils.duckdb_utils import ( - duckdb_get_table, - duckdb_save_table, - load_missing_codes_from_db, -) -from datasure.utils.navigations_utils import demo_callout -from datasure.utils.onboarding_utils import demo_output_onboarding -from datasure.utils.settings_utils import ( - load_check_settings, - save_check_settings, - trigger_save, -) - -TAB_NAME: str = "enumerators" - - -# ============================================================================= -# Pydantic Models for Data Validation -# ============================================================================= - - -class EnumeratorSettings(BaseModel): - """Settings for enumerator report configuration. - - Attributes - ---------- - date : str | None - Column name containing survey submission date. - survey_id : str | None - Column name containing survey ID. - enumerator : str | None - Column name containing enumerator identifier (required). - formdef_version : str | None - Column name containing form version. - duration : str | None - Column name containing survey duration in seconds. - team : str | None - Column name containing team identifier. - consent : str | None - Column name containing consent status. - consent_vals : list[str] | None - List of values indicating valid consent. - outcome : str | None - Column name containing survey outcome status. - outcome_vals : list[str] | None - List of values indicating completed surveys. - """ - - survey_key: str | None = Field(None, description="Survey key column") - survey_id: str | None = Field(..., min_length=1, description="Survey ID column") - survey_date: str | None = Field(None, description="Survey date column") - enumerator: str | None = Field(None, description="Enumerator ID column") - formversion: str | None = Field(None, description="Form version column") - duration: str | None = Field(None, description="Duration column") - duration_unit: str = Field("default='seconds'", description="Duration unit") - team: str | None = Field(None, description="Team identifier column") - - -class ConsentOutcomeSettings(BaseModel): - """Settings for consent and outcome configuration. - - Attributes - ---------- - consent : str | None - Column name containing consent status. - consent_vals : list[str] | None - List of values indicating valid consent. - outcome : str | None - Column name containing survey outcome status. - outcome_vals : list[str] | None - List of values indicating completed surveys. - """ - - consent: str | None = Field(None, description="Consent status column") - consent_vals: list[str] | None = Field(None, description="Valid consent values") - outcome: str | None = Field(None, description="Outcome status column") - outcome_vals: list[str] | None = Field(None, description="Completed survey values") - - -class ProductivitySettings(BaseModel): - """Settings for productivity analysis configuration. - - Attributes - ---------- - view_option : str - Time period for analysis: Daily, Weekly, or Monthly. - weekstartday : str - First day of the week for weekly analysis. - """ - - view_option: str = Field(default="Daily", description="Time period view") - weekstartday: str = Field(default="Monday", description="Week start day") - - @field_validator("view_option") - @classmethod - def validate_view_option(cls, v: str) -> str: - """Validate view option is one of the allowed values.""" - allowed = ["Daily", "Weekly", "Monthly"] - if v not in allowed: - raise ValueError(f"view_option must be one of {allowed}") - return v - - @field_validator("weekstartday") - @classmethod - def validate_weekstartday(cls, v: str) -> str: - """Validate weekstartday is one of the allowed values.""" - allowed = [ - "Monday", - "Tuesday", - "Wednesday", - "Thursday", - "Friday", - "Saturday", - "Sunday", - ] - if v not in allowed: - raise ValueError(f"weekstartday must be one of {allowed}") - return v - - -# Constants for statistics options -ALLOWED_STATISTICS = [ - "count", - "min", - "mean", - "median", - "max", - "std", - "25th percentile", - "75th percentile", -] -ALLOWED_STATISTICS_OVERTIME = ALLOWED_STATISTICS + ["missing"] -ALLOWED_TIME_PERIODS = ["Daily", "Weekly", "Monthly"] -WEEKDAY_NAMES = [ - "Monday", - "Tuesday", - "Wednesday", - "Thursday", - "Friday", - "Saturday", - "Sunday", -] - -# Maps weekday names to offset codes used in computation -WEEKDAY_OFFSET_MAP = { - "Monday": "SUN", - "Tuesday": "MON", - "Wednesday": "TUE", - "Thursday": "WED", - "Friday": "THU", - "Saturday": "FRI", - "Sunday": "SAT", -} - -# Maps offset codes to numeric values for week calculations -WEEKDAY_OFFSET_TO_NUMERIC = { - "SUN": 0, - "MON": 1, - "TUE": 2, - "WED": 3, - "THU": 4, - "FRI": 5, - "SAT": 6, -} - - -class StatisticsSettings(BaseModel): - """Settings for statistics analysis configuration. - - Attributes - ---------- - statscols : list[str] | None - Columns to compute statistics on. - stats : list[str] - Statistics to compute (count, mean, median, etc.). - """ - - statscols: list[str] | None = Field(None, description="Columns for statistics") - stats: list[str] = Field( - default=["count", "mean"], description="Statistics to compute" - ) - - @field_validator("stats") - @classmethod - def validate_stats(cls, v: list[str]) -> list[str]: - """Validate that statistics are from allowed list.""" - for stat in v: - if stat not in ALLOWED_STATISTICS: - raise ValueError( - f"Invalid statistic: {stat}. Must be one of {ALLOWED_STATISTICS}" - ) - return v - - -class StatisticsOvertimeSettings(BaseModel): - """Settings for statistics over time analysis configuration. - - Attributes - ---------- - period : str - Time period for analysis (Daily, Weekly, Monthly). - weekstartday : str - First day of the week for weekly analysis. - stat : str - Statistic to compute over time. - statscol : str | None - Column to compute statistics on. - """ - - period_overtime: str = Field(default="Week", description="Time period for analysis") - weekstartday: str = Field(default="Monday", description="Week start day") - stat: str = Field(default="count", description="Statistic to compute") - statscol: str | None = Field(None, description="Column for statistics") - - @field_validator("period_overtime") - @classmethod - def validate_period(cls, v: str) -> str: - """Validate period is from allowed list.""" - if v not in ALLOWED_TIME_PERIODS: - raise ValueError( - f"Invalid period: {v}. Must be one of {ALLOWED_TIME_PERIODS}" - ) - return v - - @field_validator("weekstartday") - @classmethod - def validate_weekstartday(cls, v: str) -> str: - """Validate weekstartday is from allowed list.""" - if v not in WEEKDAY_NAMES: - raise ValueError( - f"Invalid weekstartday: {v}. Must be one of {WEEKDAY_NAMES}" - ) - return v - - @field_validator("stat") - @classmethod - def validate_stat(cls, v: str) -> str: - """Validate stat is from allowed list.""" - if v not in ALLOWED_STATISTICS_OVERTIME: - raise ValueError( - f"Invalid statistic: {v}. Must be one of {ALLOWED_STATISTICS_OVERTIME}" - ) - return v - - -class EnumeratorOverviewMetrics(BaseModel): - """Metrics for enumerator overview. - - Attributes - ---------- - all_submissions : int - Total number of submissions. - num_active_enumerators : int - Number of enumerators active in past 7 days. - num_enumerators : int - Total number of enumerators. - num_teams : int | str - Number of teams or 'n/a' if not available. - min_submissions : int - Minimum daily submissions. - max_submissions : int - Maximum daily submissions. - avg_submissions : int - Average daily submissions. - pct_active_enumerators : str - Percentage of active enumerators formatted as string. - """ - - all_submissions: int = Field(ge=0) - num_active_enumerators: int = Field(ge=0) - num_enumerators: int = Field(ge=0) - num_teams: int | str - min_submissions: int = Field(ge=0) - max_submissions: int = Field(ge=0) - avg_submissions: int = Field(ge=0) - pct_active_enumerators: str - - -# ============================================================================= -# Settings Management Functions -# ============================================================================= - - -@st.cache_data(ttl=60) -def load_default_enumerator_settings( - settings_file: str, config: EnumeratorSettings -) -> EnumeratorSettings: - """Load and merge saved settings with default configuration. - - Loads previously saved duplicates report settings from the settings file - and merges them with the provided default configuration. Saved settings - take precedence over defaults. - - Cached for 60 seconds to reduce file I/O operations. - - Parameters - ---------- - settings_file : str - Path to the settings file containing saved configurations. - config : DuplicatesSettings - Default configuration to use as fallback for missing settings. - - Returns - ------- - DuplicatesSettings - Merged settings combining saved and default configurations. - """ - saved_settings = load_check_settings(settings_file, TAB_NAME) - - default_settings: dict = dict(config) - default_settings.update(saved_settings) - - return EnumeratorSettings(**default_settings) - - -@demo_output_onboarding(TAB_NAME) -def enumerator_report_settings( - project_id: str, - settings_file: str, - data: pl.DataFrame, - config: EnumeratorSettings, - categorical_columns: list[str], - datetime_columns: list[str], -) -> EnumeratorSettings: - """Create and render the settings UI for duplicates report configuration. - - This function creates a comprehensive Streamlit UI for configuring - duplicates report settings. It includes: - - Survey identifiers (key and ID columns) - - Survey date column selection - - Enumerator ID column - - Filtering conditions for targeted duplicate detection - - Settings are automatically saved to the settings file when changed - and loaded from previous sessions if available. - - Parameters - ---------- - project_id : str - Unique project identifier for database operations. - settings_file : str - Path to settings file for saving/loading configurations. - data : pl.DataFrame - Dataset to analyze for duplicates. - config : DuplicatesSettings - Default configuration used as fallback values. - categorical_columns : list[str] - Available categorical columns for selection (survey key, ID, enumerator). - datetime_columns : list[str] - Available datetime columns for date selection. - - Returns - ------- - DuplicatesSettings - User-configured settings from the UI. - """ - with st.expander("settings", icon=":material/settings:"): - st.markdown("## Configure settings for enumerator report") - st.write("---") - - default_settings = load_default_enumerator_settings(settings_file, config) - - # Survey Identifiers - with st.container(border=True): - st.subheader("Survey Identifiers") - si1, si2, _ = st.columns(3) - - with si1: - default_survey_key = default_settings.survey_key - default_survey_key_index = ( - categorical_columns.index(default_survey_key) - if default_survey_key and default_survey_key in categorical_columns - else None - ) - survey_key = st.selectbox( - "Survey Key", - options=categorical_columns, - key="survey_key_enumerator", - help="Select the column that contains the survey key", - index=default_survey_key_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_survey_key"}, - ) - save_check_settings(settings_file, TAB_NAME, {"survey_key": survey_key}) - - with si2: - default_survey_id = default_settings.survey_id - default_survey_id_index = ( - categorical_columns.index(default_survey_id) - if default_survey_id and default_survey_id in categorical_columns - else None - ) - survey_id = st.selectbox( - "Survey ID", - options=categorical_columns, - help="Select the column that contains the survey ID", - key="survey_id_enumerator", - index=default_survey_id_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_survey_id"}, - ) - save_check_settings(settings_file, TAB_NAME, {"survey_id": survey_id}) - - with st.container(border=True): - st.subheader("Survey Date") - - sd1, _, _ = st.columns(3) - - with sd1: - default_survey_date = default_settings.survey_date - default_survey_date_index = ( - datetime_columns.index(default_survey_date) - if default_survey_date and default_survey_date in datetime_columns - else None - ) - - survey_date = st.selectbox( - "Survey Date", - options=datetime_columns, - help="Select the column that contains the survey date", - key="survey_date_enumerator", - index=default_survey_date_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_survey_date"}, - ) - save_check_settings( - settings_file, TAB_NAME, {"survey_date": survey_date} - ) - - with st.container(border=True): - st.subheader("Enumerator") - ec1, ec2, _ = st.columns(3) - with ec1: - default_enumerator = default_settings.enumerator - default_enumerator_index = ( - categorical_columns.index(default_enumerator) - if default_enumerator and default_enumerator in categorical_columns - else None - ) - enumerator = st.selectbox( - "Enumerator ID", - options=categorical_columns, - key="enumerator_enumerator", - help="Select the column that contains the enumerator ID", - index=default_enumerator_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_enumerator"}, - ) - save_check_settings(settings_file, TAB_NAME, {"enumerator": enumerator}) - - with ec2: - default_team = default_settings.team - default_team_index = ( - categorical_columns.index(default_team) - if default_team and default_team in categorical_columns - else None - ) - team = st.selectbox( - "Team", - options=categorical_columns, - key="team_enumerator", - help="Select the column that contains the team identifier", - index=default_team_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_team"}, - ) - save_check_settings(settings_file, TAB_NAME, {"team": team}) - - with st.container(border=True): - st.subheader("Survey Duration") - dc1, dc2, _ = st.columns(3) - with dc1: - default_duration = default_settings.duration - default_duration_index = ( - categorical_columns.index(default_duration) - if default_duration and default_duration in categorical_columns - else None - ) - duration = st.selectbox( - "Duration Column", - options=categorical_columns, - key="duration_enumerator", - help="Select the column that contains the survey duration in seconds", - index=default_duration_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_duration"}, - ) - save_check_settings(settings_file, TAB_NAME, {"duration": duration}) - - with dc2: - default_duration_unit = default_settings.duration_unit - default_duration_unit_index = ( - ["seconds", "minutes", "hours"].index(default_duration_unit) - if default_duration_unit in ["seconds", "minutes", "hours"] - else 0 - ) - duration_unit = st.selectbox( - "Duration Unit", - options=["seconds", "minutes", "hours"], - key="duration_unit_enumerator", - help="Select the unit for survey duration", - index=default_duration_unit_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_duration_unit"}, - ) - save_check_settings( - settings_file, TAB_NAME, {"duration_unit": duration_unit} - ) - - with st.container(border=True): - st.subheader("Form Version") - fv1, _ = st.columns([1, 2]) - with fv1: - default_formversion = default_settings.formversion - default_formversion_index = ( - categorical_columns.index(default_formversion) - if default_formversion - and default_formversion in categorical_columns - else None - ) - formversion = st.selectbox( - "Form Version Column", - options=categorical_columns, - key="formversion_enumerator", - help="Select the column that contains the form version", - index=default_formversion_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_formversion"}, - ) - save_check_settings( - settings_file, TAB_NAME, {"formversion": formversion} - ) - - with st.container(border=True): - st.subheader("Consent and Outcome Settings") - st.info( - "Configure consent and outcome columns along with their valid values." - ) - - _render_consent_outcome_settings( - project_id, data, categorical_columns, settings_file - ) - if st.session_state.get("st_apply_consent_outcome_enumerator"): - st.success("Consent and outcome settings applied successfully.") - st.session_state["st_apply_consent_outcome_enumerator"] = False - - return EnumeratorSettings( - survey_key=survey_key, - survey_id=survey_id, - survey_date=survey_date, - enumerator=enumerator, - team=team, - formversion=formversion, - duration=duration, - duration_unit=duration_unit, - ) - - -@st.fragment -def _render_consent_outcome_settings( - project_id: str, data: pl.DataFrame, categorical_columns: list, settings_file: str -) -> ConsentOutcomeSettings: - """Render consent and outcome settings UI. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - settings_file : str - Path to settings file for saving/loading configurations. - - Returns - ------- - ConsentOutcomeSettings - User-configured consent and outcome settings. - """ - default_settings = load_check_settings(settings_file, TAB_NAME) - - with st.container(border=True): - st.subheader("Consent Settings") - co1, co2 = st.columns([0.3, 0.7]) - with co1: - default_consent_col = default_settings.get("consent") - default_consent_index = ( - categorical_columns.index(default_consent_col) - if default_consent_col and default_consent_col in categorical_columns - else 0 - ) - consent_col = st.selectbox( - "Consent Column", - options=categorical_columns, - help="Select the column that contains consent status", - key="consent_enumerator", - index=default_consent_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_consent"}, - ) - save_check_settings(settings_file, TAB_NAME, {"consent": consent_col}) - - with co2: - default_consent_vals = default_settings.get("consent_vals", []) - consent_val_options = data[consent_col].unique().to_list() - consent_vals = st.multiselect( - "Valid Consent Values", - options=consent_val_options, - default=default_consent_vals, - help="Select values that indicate valid consent", - key="consent_vals_enumerator", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_consent_vals"}, - ) - save_check_settings(settings_file, TAB_NAME, {"consent_vals": consent_vals}) - - with st.container(border=True): - st.subheader("Outcome Settings") - oo1, oo2 = st.columns([0.3, 0.7]) - with oo1: - default_outcome_col = default_settings.get("outcome") - default_outcome_index = ( - categorical_columns.index(default_outcome_col) - if default_outcome_col and default_outcome_col in categorical_columns - else 0 - ) - outcome_col = st.selectbox( - "Outcome Column", - options=categorical_columns, - help="Select the column that contains survey outcome status", - key="outcome_enumerator", - index=default_outcome_index, - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_outcome"}, - ) - save_check_settings(settings_file, TAB_NAME, {"outcome": outcome_col}) - - with oo2: - default_outcome_vals = default_settings.get("outcome_vals", []) - outcome_val_options = data[outcome_col].unique().to_list() - outcome_vals = st.multiselect( - "Completed Survey Values", - options=outcome_val_options, - default=default_outcome_vals, - help="Select values that indicate completed surveys", - key="outcome_vals_enumerator", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_outcome_vals"}, - ) - save_check_settings(settings_file, TAB_NAME, {"outcome_vals": outcome_vals}) - - config_dict = { - "consent": consent_col, - "consent_vals": consent_vals, - "outcome": outcome_col, - "outcome_vals": outcome_vals, - } - config = ConsentOutcomeSettings(**config_dict) - - if st.button( - "Apply Consent and Outcome Settings", - key="apply_consent_outcome_enumerator", - type="primary", - width="stretch", - ): - _create_enum_data_on_settings(project_id, data, config) - _trigger_success_message("st_apply_consent_outcome_enumerator") - st.rerun() - - -def _trigger_success_message(button_key: str) -> None: - """Trigger a success message after button click. - - Parameters - ---------- - button_key : str - Unique key of the button to associate the success message with. - """ - st.session_state[f"{button_key}"] = True - - -def _create_enum_data_on_settings( - project_id: str, - data: pl.DataFrame, - config: ConsentOutcomeSettings, -) -> None: - """Create enumerator data based on consent and outcome settings. - - Parameters - ---------- - project_id : str - Unique project identifier for database operations. - data : pl.DataFrame - DataFrame containing survey data. - conditions : ConsentOutcomeSettings - Consent and outcome configuration settings. - """ - # If consent and consent values are provided, create a dummy column indicating - # valid consent else set to 1 - - if config.consent and config.consent_vals: - enum_data = data.with_columns( - pl.col(config.consent) - .is_in(config.consent_vals) - .cast(pl.Int32) - .alias("consent_granted_agg_col") - ) - else: - enum_data = data.with_columns( - pl.lit(1).cast(pl.Int32).alias("consent_granted_agg_col") - ) - - if config.outcome and config.outcome_vals: - enum_data = enum_data.with_columns( - pl.col(config.outcome) - .is_in(config.outcome_vals) - .cast(pl.Int32) - .alias("completed_survey_agg_col") - ) - else: - enum_data = enum_data.with_columns( - pl.lit(1).cast(pl.Int32).alias("completed_survey_agg_col") - ) - - # save to database - duckdb_save_table( - project_id, - enum_data, - "enumerator_data_with_consent_outcome", - "intermediate", - ) - - -# ============================================================================= -# Overview Computation Functions -# ============================================================================= - - -def compute_enumerator_overview( - data: pl.DataFrame, date: str, enumerator: str, team: str | None -) -> EnumeratorOverviewMetrics: - """Compute enumerator overview metrics. - - Calculates key metrics including total submissions, active enumerators, - team counts, and submission statistics. - - Cached for 5 minutes to improve performance for repeated calls. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - team : str | None - Team column name (optional). - - Returns - ------- - EnumeratorOverviewMetrics - Overview metrics for enumerators including counts and statistics. - """ - if data.is_empty(): - raise ValueError( - "Input data is empty. Cannot compute enumerator overview metrics." - ) - data = data.sort([enumerator, date]) - - all_submissions = data.height - - # Calculate daily submissions - data = data.with_row_index(name="TOKEN KEY") - daily_submissions_sum = data.group_by([date, enumerator]).agg( - pl.col("TOKEN KEY").count().alias("count") - ) - - # Calculate active enumerators (past 7 days) - from datetime import date as dt_date - from datetime import timedelta - - active_date_cut_off = dt_date.today() - timedelta(weeks=1) - - daily_submissions_sum = daily_submissions_sum.with_columns( - (pl.col(date).cast(pl.Date) > active_date_cut_off).alias("active") - ) - - num_active_enumerators = ( - daily_submissions_sum.filter(pl.col("active")) - .select(pl.col(enumerator).n_unique()) - .item() - ) - - num_enumerators = data[enumerator].n_unique() - num_teams = data[team].n_unique() if team else "n/a" - min_submissions = int(daily_submissions_sum["count"].min()) - max_submissions = int(daily_submissions_sum["count"].max()) - avg_submissions = int(daily_submissions_sum["count"].mean()) - - pct_active_enumerators = f"{(num_active_enumerators / num_enumerators) * 100:.0f}%" - - return EnumeratorOverviewMetrics( - all_submissions=all_submissions, - num_active_enumerators=num_active_enumerators, - num_enumerators=num_enumerators, - num_teams=num_teams, - min_submissions=min_submissions, - max_submissions=max_submissions, - avg_submissions=avg_submissions, - pct_active_enumerators=pct_active_enumerators, - ) - - -def compute_enumerator_missing_table( - data: pl.DataFrame, missing_codes_config: pl.DataFrame, group_by_col: list[str] -) -> pl.DataFrame: - """Compute missing data statistics per enumerator. - - Calculates missing data counts and percentages for each enumerator - based on provided missing codes configuration. - - Cached for 5 minutes to improve performance for repeated calls. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - missing_settings_file : str - Path to missing codes configuration file. - enumerator : str - Enumerator column name. - - Returns - ------- - pl.DataFrame - DataFrame with missing data statistics per enumerator. - """ - # define columns to exclude from missing data stats - columns_to_exclude = ["consent_granted_agg_col", "completed_survey_agg_col"] - - data_for_missing = data.select( - [col for col in data.columns if col not in columns_to_exclude] - ) - - # Metadata for missing data calculation - enum_data_missing = data_for_missing.select(group_by_col) - - # If missing_codes_config is empty, calculate only null missingness - if missing_codes_config.is_empty(): - # Calculate overall null missingness per enumerator - columns_to_check = [ - col for col in data_for_missing.columns if col not in group_by_col - ] - - # Count nulls per row and calculate percentage - missing_summary = data_for_missing.with_columns( - [ - pl.sum_horizontal( - [pl.col(col).is_null().cast(pl.Int32) for col in columns_to_check] - ).alias("_null_count"), - pl.lit(len(columns_to_check)).alias("_total_fields"), - ] - ).with_columns( - (pl.col("_null_count") / pl.col("_total_fields") * 100).alias( - "% Null values" - ) - ) - - # Group by enumerator and calculate mean missingness rate - result_df = missing_summary.group_by(group_by_col, maintain_order=True).agg( - pl.col("% Null values").mean() - ) - - return result_df - - # If missing_codes_config is provided, calculate missingness by category - # Get missing code pairs from the config - missing_code_pairs = missing._get_missing_code_pairs(missing_codes_config) - - # Compute missing data with paired encoding - missing_data_encoded = missing._compute_missing_data_paired( - data_for_missing, missing_codes_config - ) - - # Get columns to check (exclude enumerator column) - columns_to_check = [ - col for col in missing_data_encoded.columns if col not in group_by_col - ] - total_fields = len(columns_to_check) - - # Calculate counts for each missing category per row - agg_expressions = [] - - # Add null values count (encoded as 1) - agg_expressions.append( - pl.sum_horizontal( - [(pl.col(col) == 1).cast(pl.Int32) for col in columns_to_check] - ).alias("_null_count") - ) - - # Add count for each special missing code category (starting from 2) - for idx, label in enumerate(missing_code_pairs.keys(), start=2): - agg_expressions.append( - pl.sum_horizontal( - [(pl.col(col) == idx).cast(pl.Int32) for col in columns_to_check] - ).alias(f"_{label}_count") - ) - - # Add total missing count (any value > 0) - agg_expressions.append( - pl.sum_horizontal( - [(pl.col(col) > 0).cast(pl.Int32) for col in columns_to_check] - ).alias("_total_missing_count") - ) - - # Add total fields count - agg_expressions.append(pl.lit(total_fields).alias("_total_fields")) - - # Apply the aggregations - missing_counts = missing_data_encoded.select( - [pl.col(group_by_col)] + agg_expressions - ) - - # Calculate percentages - percentage_expressions = [] - - # Null values percentage - percentage_expressions.append( - (pl.col("_null_count") / pl.col("_total_fields") * 100).alias("% Null values") - ) - - # Special missing code category percentages - for label in missing_code_pairs: - percentage_expressions.append( - (pl.col(f"_{label}_count") / pl.col("_total_fields") * 100).alias( - f"% {label}" - ) - ) - - # Total missing percentage - percentage_expressions.append( - (pl.col("_total_missing_count") / pl.col("_total_fields") * 100).alias( - "% Total Missing" - ) - ) - - missing_with_percentages = missing_counts.with_columns(percentage_expressions) - - # Group by enumerator and calculate mean percentages - final_agg_expressions = [pl.col("% Null values").mean()] - - for label in missing_code_pairs: - final_agg_expressions.append(pl.col(f"% {label}").mean()) - - final_agg_expressions.append(pl.col("% Total Missing").mean()) - - # drop enumerator column from missing_with_percentages - missing_with_percentages = missing_with_percentages.select( - [col for col in missing_with_percentages.columns if col not in group_by_col] - ) - # merge missing_with_percentages with enumerator column - missing_with_percentages = pl.concat( - [enum_data_missing, missing_with_percentages], how="horizontal" - ) - - result_df = missing_with_percentages.group_by( - group_by_col, maintain_order=True - ).agg(final_agg_expressions) - - return result_df - - -def compute_enumerator_summary( - project_id: str, - data: pl.DataFrame, - date: str, - enumerator: str, - team: str | None, - formversion: str | None, - duration: str | None, -) -> pl.DataFrame: - """Compute comprehensive enumerator summary statistics. - - Calculates submission counts, date ranges, duration statistics, - form version tracking, consent rates, outcome rates, and missing data - patterns for each enumerator. - - Cached for 5 minutes to improve performance for repeated calls. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - missing_settings_file : str - Path to missing codes configuration file. - date : str - Date column name. - enumerator : str - Enumerator column name. - formdef_version : str | None - Form version column name (optional). - duration : str | None - Duration column name (optional). - consent : str | None - Consent column name (optional). - consent_vals : list[str] | None - List of values indicating valid consent (optional). - outcome : str | None - Outcome column name (optional). - outcome_vals : list[str] | None - List of values indicating completed surveys (optional). - - Returns - ------- - pl.DataFrame - Comprehensive summary DataFrame with enumerator statistics. - """ - group_by_cols = [enumerator, team] if team else [enumerator] - # Format date column - df = data.with_columns(pl.col(date).dt.strftime("%b %d, %Y").alias(date)) - - # Basic summary aggregations - summary_df = df.group_by(group_by_cols, maintain_order=True).agg( - [ - pl.col(date).min().alias("first submission"), - pl.col(date).max().alias("last submission"), - pl.col(date).count().alias("# submissions"), - pl.col(date).n_unique().alias("# unique dates"), - ] - ) - - # Calculate time-based submissions - today = dt_date.today() - start_of_week = today - timedelta(days=today.weekday()) - start_of_month = today.replace(day=1) - - today_str = today.strftime("%b %d, %Y") - week_str = start_of_week.strftime("%b %d, %Y") - month_str = start_of_month.strftime("%b %d, %Y") - - df = df.with_columns( - [ - (pl.col(date) == today_str).alias("submitted_today"), - (pl.col(date) >= week_str).alias("submitted_this_week"), - (pl.col(date) >= month_str).alias("submitted_this_month"), - ] - ) - - lagged_df = df.group_by(group_by_cols, maintain_order=True).agg( - [ - pl.col("submitted_today").sum().alias("# submissions today"), - pl.col("submitted_this_week").sum().alias("# submissions this week"), - pl.col("submitted_this_month").sum().alias("# submissions this month"), - ] - ) - - summary_df = summary_df.join(lagged_df, on=group_by_cols, how="left") - - # Add missing data statistics - missing_settings_file = load_missing_codes_from_db(project_id) - enumerator_missing_df = compute_enumerator_missing_table( - data, missing_settings_file, group_by_cols - ) - summary_df = summary_df.join(enumerator_missing_df, on=group_by_cols, how="left") - - # Add duration statistics if available - if duration: - duration_df = df.group_by(group_by_cols, maintain_order=True).agg( - [ - pl.col(duration).min().alias("min duration"), - pl.col(duration).mean().alias("mean duration"), - pl.col(duration).median().alias("median duration"), - pl.col(duration).max().alias("max duration"), - ] - ) - summary_df = summary_df.join(duration_df, on=group_by_cols, how="left") - - # Add form version statistics if available - if formversion: - # Get latest form version per date - formdef_outdated = df.group_by(date, maintain_order=True).agg( - pl.col(formversion).max().alias("latest daily form version") - ) - - df = df.join(formdef_outdated, on=date, how="left") - df = df.with_columns( - (pl.col(formversion) != pl.col("latest daily form version")).alias( - "outdated_form_version" - ) - ) - - formdef_outdated_df = df.group_by(group_by_cols, maintain_order=True).agg( - pl.col("outdated_form_version").sum().alias("# of outdated form versions") - ) - - formdef_df = df.group_by(group_by_cols, maintain_order=True).agg( - [ - pl.col(formversion).n_unique().alias("# form versions"), - pl.col(formversion).max().alias("latest form version"), - ] - ) - - latest_enum_formversion = df.group_by(group_by_cols, maintain_order=True).agg( - pl.col(formversion).max().alias("last form version") - ) - - summary_df = summary_df.join(formdef_df, on=group_by_cols, how="left") - summary_df = summary_df.join(formdef_outdated_df, on=group_by_cols, how="left") - summary_df = summary_df.join( - latest_enum_formversion, on=group_by_cols, how="left" - ) - - # Add consent statistics if available - if "consent_granted_agg_col" in df.columns: - consent_df = df.group_by(group_by_cols, maintain_order=True).agg( - pl.col("consent_granted_agg_col").mean().alias("% consent") - ) - summary_df = summary_df.join(consent_df, on=group_by_cols, how="left") - - # Add outcome statistics if available - if "completed_survey_agg_col" in df.columns: - outcome_df = df.group_by(group_by_cols, maintain_order=True).agg( - pl.col("completed_survey_agg_col").mean().alias("% completed survey") - ) - summary_df = summary_df.join(outcome_df, on=group_by_cols, how="left") - - return summary_df - - -# ============================================================================= -# Productivity Computation Functions -# ============================================================================= - - -def compute_enumerator_productivity( - data: pl.DataFrame, - date: str, - group_by_cols: list[str], - period: str, - weekstartday: str, -) -> pl.DataFrame: - """Compute enumerator productivity over time. - - Analyzes submission counts by enumerator across time periods (daily, - weekly, or monthly). - - Cached for 5 minutes to improve performance for repeated calls. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - period : str - Time period: "Daily", "Weekly", "Monthly", "Day", "Week", or "Month". - weekstartday : str - Start day of the week (e.g., "SUN", "MON") for weekly analysis. - - Returns - ------- - pl.DataFrame - Pivoted DataFrame with enumerators as rows and time periods as columns. - """ - prod_df = data.clone() - - # Normalize period values to handle both old and new formats - period_normalized = period - if period == "Day": - period_normalized = "Daily" - elif period == "Week": - period_normalized = "Weekly" - elif period == "Month": - period_normalized = "Monthly" - - # Create time period column based on selection with user-friendly formatting - if period_normalized == "Daily": - # Format as "Jan 1, 2025" - prod_df = prod_df.with_columns( - pl.col(date).dt.strftime("%b %d, %Y").alias("TIME PERIOD") - ) - elif period_normalized == "Weekly": - # Calculate week start and end dates for user-friendly display - offset = WEEKDAY_OFFSET_TO_NUMERIC.get(weekstartday, 1) - - # Calculate the week start date (beginning of the week containing this date) - # weekday() returns 0=Monday, 6=Sunday - prod_df = prod_df.with_columns( - [ - # Calculate days since the start of the week - ((pl.col(date).dt.weekday() - offset + 7) % 7).alias( - "_days_since_week_start" - ), - ] - ) - - # Calculate week_start_date by subtracting days_since_week_start - prod_df = prod_df.with_columns( - [ - ( - pl.col(date) - pl.duration(days=pl.col("_days_since_week_start")) - ).alias("_week_start"), - ( - pl.col(date) - - pl.duration(days=pl.col("_days_since_week_start")) - + pl.duration(days=6) - ).alias("_week_end"), - ] - ) - - # Format as "Jan 1, 2025 to Jan 7, 2025" - prod_df = prod_df.with_columns( - ( - pl.col("_week_start").dt.strftime("%b %d, %Y") - + " to " - + pl.col("_week_end").dt.strftime("%b %d, %Y") - ).alias("TIME PERIOD") - ) - elif period_normalized == "Monthly": - # Format as "January 2025" - prod_df = prod_df.with_columns( - pl.col(date).dt.strftime("%B %Y").alias("TIME PERIOD") - ) - - # Count submissions per period and enumerator - prod_df = prod_df.with_row_index(name="TOKEN KEY") - prod_res = prod_df.group_by( - ["TIME PERIOD"] + group_by_cols, maintain_order=True - ).agg(pl.col("TOKEN KEY").count().alias("submissions")) - - # Pivot to wide format - prod_res = prod_res.pivot( - index=group_by_cols, - on="TIME PERIOD", - values="submissions", - ).fill_null(0) - - return prod_res - - -# ============================================================================= -# Statistics Computation Functions -# ============================================================================= - - -@st.cache_data(ttl=300) -def compute_enumerator_statistics( - data: pl.DataFrame, - group_by_cols: list[str], - statscols: list[str], - stats: list[str], -) -> pl.DataFrame: - """Compute enumerator statistics across specified columns. - - Calculates summary statistics (mean, median, std, etc.) for numeric - columns grouped by enumerator (and optionally team). - - Cached for 5 minutes to improve performance for repeated calls. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - group_by_cols : list[str] - List of columns to group by (e.g., ["enumerator"] or ["enumerator", "team"]). - statscols : list[str] - List of columns to compute statistics on. - stats : list[str] - List of statistics to compute (e.g., ["mean", "median", "std"]). - - Returns - ------- - pl.DataFrame - DataFrame with enumerators (and teams) and computed statistics. - """ - # Map stat names to Polars expressions - stat_mapping = { - "count": "count", - "min": "min", - "mean": "mean", - "median": "median", - "max": "max", - "std": "std", - "25th percentile": "quantile", - "75th percentile": "quantile", - } - - agg_exprs = [] - for col in statscols: - for stat in stats: - if stat == "25th percentile": - agg_exprs.append(pl.col(col).quantile(0.25).alias(f"{col}_{stat}")) - elif stat == "75th percentile": - agg_exprs.append(pl.col(col).quantile(0.75).alias(f"{col}_{stat}")) - else: - method = stat_mapping.get(stat, stat) - agg_exprs.append(getattr(pl.col(col), method)().alias(f"{col}_{stat}")) - - stats_res = data.group_by(group_by_cols, maintain_order=True).agg(agg_exprs) - - return stats_res - - -def compute_enumerator_statistics_overtime( - data: pl.DataFrame, - date: str, - group_by_cols: list[str], - statscol: str, - stat: str, - period: str, - weekstartday: str, -) -> pl.DataFrame: - """Compute enumerator statistics over time for a specific column. - - Analyzes how a specific statistic changes over time periods for each - enumerator (and optionally team). - - Cached for 5 minutes to improve performance for repeated calls. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - group_by_cols : list[str] - List of columns to group by (e.g., ["enumerator"] or ["enumerator", "team"]). - statscol : str - Column to compute statistics on. - stat : str - Statistic to compute (e.g., "mean", "median", "missing"). - period : str - Time period: "Daily", "Weekly", "Monthly", "Day", "Week", or "Month". - weekstartday : str - Start day of the week for weekly analysis. - - Returns - ------- - pl.DataFrame - Pivoted DataFrame with enumerators (and teams) as rows and time - periods as columns. - """ - stats_overtime_df = data.select([date] + group_by_cols + [statscol]).clone() - - # Normalize period values to handle both old and new formats - period_normalized = period - if period == "Day": - period_normalized = "Daily" - elif period == "Week": - period_normalized = "Weekly" - elif period == "Month": - period_normalized = "Monthly" - - # Create time period column with user-friendly formatting - if period_normalized == "Daily": - # Format as "Jan 1, 2025" - stats_overtime_df = stats_overtime_df.with_columns( - pl.col(date).dt.strftime("%b %d, %Y").alias("TIME PERIOD") - ) - elif period_normalized == "Weekly": - # Calculate week start and end dates for user-friendly display - offset = WEEKDAY_OFFSET_TO_NUMERIC.get(weekstartday, 1) - - # Calculate the week start date (beginning of the week containing this date) - # weekday() returns 0=Monday, 6=Sunday - stats_overtime_df = stats_overtime_df.with_columns( - [ - # Calculate days since the start of the week - ((pl.col(date).dt.weekday() - offset + 7) % 7).alias( - "_days_since_week_start" - ), - ] - ) - - # Calculate week_start_date by subtracting days_since_week_start - stats_overtime_df = stats_overtime_df.with_columns( - [ - ( - pl.col(date) - pl.duration(days=pl.col("_days_since_week_start")) - ).alias("_week_start"), - ( - pl.col(date) - - pl.duration(days=pl.col("_days_since_week_start")) - + pl.duration(days=6) - ).alias("_week_end"), - ] - ) - - # Format as "Jan 1, 2025 to Jan 7, 2025" - stats_overtime_df = stats_overtime_df.with_columns( - ( - pl.col("_week_start").dt.strftime("%b %d, %Y") - + " to " - + pl.col("_week_end").dt.strftime("%b %d, %Y") - ).alias("TIME PERIOD") - ) - elif period_normalized == "Monthly": - # Format as "January 2025" - stats_overtime_df = stats_overtime_df.with_columns( - pl.col(date).dt.strftime("%B %Y").alias("TIME PERIOD") - ) - - # Calculate statistic - if stat == "missing": - stats_overtime_res = stats_overtime_df.group_by( - ["TIME PERIOD"] + group_by_cols, maintain_order=True - ).agg(pl.col(statscol).is_null().mean().alias("_STAT")) - elif stat == "25th percentile": - stats_overtime_res = stats_overtime_df.group_by( - ["TIME PERIOD"] + group_by_cols, maintain_order=True - ).agg(pl.col(statscol).quantile(0.25).alias("_STAT")) - elif stat == "75th percentile": - stats_overtime_res = stats_overtime_df.group_by( - ["TIME PERIOD"] + group_by_cols, maintain_order=True - ).agg(pl.col(statscol).quantile(0.75).alias("_STAT")) - else: - stats_overtime_res = stats_overtime_df.group_by( - ["TIME PERIOD"] + group_by_cols, maintain_order=True - ).agg(getattr(pl.col(statscol), stat)().alias("_STAT")) - - # Pivot to wide format - stats_overtime_res = stats_overtime_res.pivot( - index=group_by_cols, on="TIME PERIOD", values="_STAT" - ) - - return stats_overtime_res - - -# ============================================================================= -# Display Functions - Overview -# ============================================================================= - - -def _render_enumerator_overview_metrics( - data: pl.DataFrame, date: str, enumerator: str, team: str | None -) -> None: - """Display enumerator overview metrics. - - Shows key metrics including total submissions, active enumerators, - team counts, and submission statistics in a grid layout. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - team : str | None - Team column name (optional). - """ - if not (enumerator and date): - st.info( - "Enumerator overview requires a date and enumerator column to be selected. " - "Go to the :material/settings: settings section above to select them." - ) - return - - metrics: EnumeratorOverviewMetrics = compute_enumerator_overview( - data, date, enumerator, team - ) - - tc1, tc2, tc3, tc4 = st.columns(4, border=True) - num_enumerators_formatted = ( - f"{metrics.num_enumerators:,}" - if isinstance(metrics.num_enumerators, int) - else metrics.num_enumerators - ) - tc1.metric( - r"\# of enumerators", - num_enumerators_formatted, - help="Total unique enumerators in the dataset", - ) - num_teams_formatted = ( - f"{metrics.num_teams:,}" - if isinstance(metrics.num_teams, int) - else metrics.num_teams - ) - tc2.metric( - r"\# of teams", num_teams_formatted, help="Total unique teams in the dataset" - ) - num_active_enumerators_formatted = f"{metrics.num_active_enumerators:,}" - tc3.metric( - r"\# of Active enumerators (past 7 days)", - num_active_enumerators_formatted, - help="Number of enumerators with submissions in the past 7 days", - ) - pct_active_enumerators_formatted = f"{metrics.pct_active_enumerators}" - tc4.metric( - "% of active enumerator (past 7 days)", - pct_active_enumerators_formatted, - help="Percentage of enumerators active in the past 7 days", - ) - - bc1, bc2, bc3, bc4 = st.columns(4, border=True) - min_submissions_formatted = f"{metrics.min_submissions:,}" - bc1.metric( - "Fewest enumerator submissions", - min_submissions_formatted, - help="Minimum number of submissions by any enumerator", - ) - max_submissions_formatted = f"{metrics.max_submissions:,}" - bc2.metric( - "Highest enumerator submissions", - max_submissions_formatted, - help="Maximum number of submissions by any enumerator", - ) - avg_submissions_formatted = f"{metrics.avg_submissions:,}" - bc3.metric( - "Average enumerator submissions", - avg_submissions_formatted, - help="Average number of submissions per enumerator", - ) - all_submissions_formatted = f"{metrics.all_submissions:,}" - bc4.metric( - "Total survey submissions", - all_submissions_formatted, - help="Total number of survey submissions in the dataset", - ) - - -@st.fragment -def _render_enumerator_summary_table( - project_id: str, - data: pl.DataFrame, - date: str, - enumerator: str, - team: str | None, - formversion: str | None, - duration: str | None, -) -> None: - """Display enumerator summary table. - - Shows comprehensive enumerator statistics including submission counts, - duration, missing data, consent rates, and outcome rates with styled - formatting. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - missing_settings_file : str - Path to missing codes configuration file. - date : str - Date column name. - enumerator : str - Enumerator column name. - formdef_version : str | None - Form version column name (optional). - duration : str | None - Duration column name (optional). - consent : str | None - Consent column name (optional). - consent_vals : list[str] | None - Valid consent values (optional). - outcome : str | None - Outcome column name (optional). - outcome_vals : list[str] | None - Completed survey values (optional). - """ - if not (enumerator and date): - st.info( - "Enumerator summary requires a date and enumerator column to be selected. " - "Go to the :material/settings: settings section above to select them." - ) - return - - summary_df = compute_enumerator_summary( - project_id, - data, - date, - enumerator, - team, - formversion, - duration, - ) - - options_map = { - "submissions": ":material/arrow_upload_progress: Submissions", - "missing": ":material/incomplete_circle: Missing Data", - "duration": ":material/timer: Duration", - "formversion": ":material/difference: Form Version", - "consent_outcome": ":material/check_circle: Consent & Outcome", - } - with st.container(horizontal_alignment="left"): - show_info = st.pills( - "Select Summary Information to Display", - options=options_map.keys(), - format_func=lambda x: options_map[x], - key="show_info_enumerator", - help="Select which summary information to display in the table", - selection_mode="multi", - ) - - # Define column groups - column_groups = { - "submissions": [ - "first submission", - "last submission", - "# submissions", - "# unique dates", - "# submissions today", - "# submissions this week", - "# submissions this month", - ], - "missing": [ - col - for col in summary_df.columns - if "%" in col - and ( - "Null" in col - or "Missing" in col - or any( - keyword in col - for keyword in ["Don't Know", "Refuse", "Not Applicable"] - ) - ) - ], - "duration": [ - "min duration", - "mean duration", - "median duration", - "max duration", - ], - "formversion": [ - "# form versions", - "latest form version", - "last form version", - "# of outdated form versions", - ], - "consent_outcome": [ - "% consent", - "% completed survey", - ], - } - - # Always include enumerator and # submissions - columns_to_show = ( - [enumerator, team, "# submissions"] if team else [enumerator, "# submissions"] - ) - - # Filter columns based on selection - if show_info: - # Add columns from selected categories - for category in show_info: - columns_to_show.extend( - [ - col - for col in column_groups[category] - if col in summary_df.columns and col not in columns_to_show - ] - ) - - # Filter the dataframe - filtered_df = summary_df.select(columns_to_show) - else: - # Show all columns if nothing is selected - filtered_df = summary_df - - # Display using Streamlit's native dataframe display - # create column config for enumerator and team conditionally - # Build column configuration dynamically - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - } - - # Add team column if available - if team: - column_config[team] = st.column_config.TextColumn("Team", pinned=True) - - # Add remaining columns - column_config.update( - { - "# submissions": st.column_config.NumberColumn( - "# of Submissions", format="%d", pinned=True - ), - "# unique dates": st.column_config.NumberColumn("# of Days", format="%d"), - "# submissions today": st.column_config.NumberColumn( - "# submitted Today", format="%d" - ), - "# submissions this week": st.column_config.NumberColumn( - "# submitted This Week", format="%d" - ), - "# submissions this month": st.column_config.NumberColumn( - "# submitted This Month", format="%d" - ), - "% Null values": st.column_config.NumberColumn( - "% Null Values", format="%.2f%%" - ), - "% Total Missing": st.column_config.NumberColumn( - "% Total Missing", format="%.2f%%" - ), - "% consent": st.column_config.NumberColumn("% Consent", format="%.2f%%"), - "% completed survey": st.column_config.NumberColumn( - "% Completed", format="%.2f%%" - ), - "min duration": st.column_config.NumberColumn( - "Min Duration (s)", format="%.2f" - ), - "mean duration": st.column_config.NumberColumn( - "Mean Duration (s)", format="%.2f" - ), - "median duration": st.column_config.NumberColumn( - "Median Duration (s)", format="%.2f" - ), - "max duration": st.column_config.NumberColumn( - "Max Duration (s)", format="%.2f" - ), - } - ) - - st.dataframe( - filtered_df, - hide_index=True, - width="stretch", - column_config=column_config, - ) - - -# ============================================================================= -# Display Functions - Productivity -# ============================================================================= - - -def _render_enumerator_productivity( - data: pl.DataFrame, - date: str, - enumerator: str, - team: str | None, - settings_file: str, -) -> None: - """Display enumerator productivity table. - - Shows submission counts by enumerator over time with configurable - time periods (daily, weekly, monthly). - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - settings_file : str - Path to settings file for saving/loading configurations. - """ - if not (enumerator and date): - st.info( - "Enumerator productivity requires a date and enumerator column to be selected. " - "Go to the :material/settings: settings section above to select them." - ) - return - - _render_enumerator_productivity_table(data, date, enumerator, team, settings_file) - - -@st.fragment -def _render_enumerator_productivity_table( - data: pl.DataFrame, - date: str, - enumerator: str, - team: str | None, - settings_file: str, -) -> None: - """Display enumerator productivity table. - Shows submission counts by enumerator over time with configurable - time periods (daily, weekly, monthly). - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - settings_file : str - Path to settings file for saving/loading configurations. - """ - time_period = _render_time_period_selector(settings_file, tab_name=TAB_NAME) - if time_period == "Week": - weekstartday = _render_weekday_selector(settings_file, tab_name=TAB_NAME) - else: - weekstartday = "MON" # Default value, not used for non-weekly periods - - group_by_cols = [enumerator, team] if team else [enumerator] - productivity_df = compute_enumerator_productivity( - data, date, group_by_cols, time_period, weekstartday - ) - - if team: - # Build column configuration dynamically - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - team: st.column_config.TextColumn("Team", pinned=True), - } - else: - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - } - - column_config.update( - { - col: st.column_config.NumberColumn(col, format="%d") - for col in productivity_df.columns - if col not in group_by_cols - } - ) - - st.dataframe( - productivity_df, hide_index=True, width="stretch", column_config=column_config - ) - - -def _render_time_period_selector( - settings_file: str, - tab_name: str = TAB_NAME, -) -> Literal["Day", "Week", "Month"]: - """Render time period selector widget using pills interface. - - Displays a pills widget allowing users to choose the time aggregation period - for productivity analysis (Day, Week, or Month). - - Parameters - ---------- - settings_file : str - Path to settings file for saving/loading configurations. - tab_name : str - Name of the tab for settings storage (default: TAB_NAME). - - Returns - ------- - Literal["Day", "Week", "Month"] - Selected time period. - """ - options_map = { - "Day": ":material/event: Daily", - "Week": ":material/date_range: Weekly", - "Month": ":material/calendar_month: Monthly", - } - - saved_settings = load_check_settings(settings_file, tab_name) or {} - default_time_period = saved_settings.get( - "time_period_enumerator_productivity", "Day" - ) - - with st.container(horizontal_alignment="left"): - time_period = st.pills( - label="Time Period", - options=options_map.keys(), - format_func=lambda x: options_map[x], - key="time_period_enumerator_productivity_key", - default=default_time_period, - help="Select time period for aggregating productivity", - selection_mode="single", - on_change=trigger_save, - kwargs={"state_name": tab_name + "_time_period"}, - ) - save_check_settings(settings_file, tab_name, {"time_period": time_period}) - - return time_period or "Day" - - -def _render_weekday_selector( - settings_file: str, - tab_name: str = TAB_NAME, -) -> str: - """Render weekday selector widget for productivity analysis. - - Displays a selectbox allowing users to choose the first day of the week - for weekly productivity calculations. - - Parameters - ---------- - settings_file : str - Path to settings file for saving/loading configurations. - tab_name : str - Name of the tab for settings storage (default: TAB_NAME). - - Returns - ------- - str - Weekday offset code (e.g., "SUN", "MON") for calculations. - """ - saved_settings = load_check_settings(settings_file, tab_name) or {} - default_weekstartday_sel = saved_settings.get( - "weekstartday_enumerator_productivity", "Monday" - ) - default_weekstartday_sel_index = WEEKDAY_NAMES.index(default_weekstartday_sel) - - cl1, _ = st.columns([1, 3]) - with cl1: - weekstartday_sel = st.selectbox( - label="Select the first day of the week", - options=WEEKDAY_NAMES, - index=default_weekstartday_sel_index, - key="week_start_day_enumerator_productivity_key", - help="Select the first day of the week", - on_change=trigger_save, - kwargs={"state_name": tab_name + "_weekstartday"}, - ) - save_check_settings(settings_file, tab_name, {"weekstartday": weekstartday_sel}) - - return WEEKDAY_OFFSET_MAP[weekstartday_sel] - - -# ============================================================================= -# Display Functions - Statistics -# ============================================================================= - - -def _load_statistics_settings(settings_file: str) -> StatisticsSettings: - """Load and validate statistics settings from file. - - Parameters - ---------- - settings_file : str - Path to settings file. - - Returns - ------- - StatisticsSettings - Validated statistics settings. - """ - saved_settings = load_check_settings(settings_file, TAB_NAME) or {} - try: - return StatisticsSettings(**saved_settings) - except ValueError: - # Return default settings if validation fails - return StatisticsSettings() - - -def _get_numeric_columns( - data: pl.DataFrame, exclude_cols: list[str] | None = None -) -> list[str]: - """Extract numeric column names from DataFrame. - - Parameters - ---------- - data : pl.DataFrame - DataFrame to extract columns from. - exclude_cols : list[str] | None - Columns to exclude from the result. - - Returns - ------- - list[str] - List of numeric column names. - """ - exclude_cols = exclude_cols or [] - return [ - col - for col in data.columns - if data[col].dtype in pl.NUMERIC_DTYPES and col not in exclude_cols - ] - - -def _render_column_selector( - numeric_cols: list[str], - default_cols: list[str] | None, - settings_file: str, -) -> list[str]: - """Render column selection widget. - - Parameters - ---------- - numeric_cols : list[str] - Available numeric columns. - default_cols : list[str] | None - Default selected columns. - settings_file : str - Path to settings file. - - Returns - ------- - list[str] - Selected columns. - """ - selected_cols = st.multiselect( - label="Select columns:", - options=numeric_cols, - default=default_cols, - help="Select columns to include in statistics", - key="selected_columns_enumerator", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_statscols"}, - ) - save_check_settings(settings_file, TAB_NAME, {"statscols": selected_cols}) - return selected_cols - - -def _render_statistics_selector( - default_stats: list[str], - settings_file: str, -) -> list[str]: - """Render statistics selection widget. - - Parameters - ---------- - default_stats : list[str] - Default selected statistics. - settings_file : str - Path to settings file. - - Returns - ------- - list[str] - Selected statistics. - """ - selected_stats = st.multiselect( - "Select statistics:", - options=ALLOWED_STATISTICS, - default=default_stats, - help="Select statistics to calculate", - key="statistics_options_enumerator", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_stats"}, - ) - save_check_settings(settings_file, TAB_NAME, {"stats": selected_stats}) - return selected_stats - - -@st.fragment -def _render_enumerator_statistics_table( - data: pl.DataFrame, - enumerator: str, - team: str | None, - settings_file: str, -) -> None: - """Display enumerator statistics table with team support. - - Shows configurable summary statistics for selected numeric columns - grouped by enumerator (and optionally team). - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - enumerator : str - Enumerator column name. - team : str | None - Team column name (optional). - settings_file : str - Path to settings file for saving/loading configurations. - """ - # Validate inputs - if not enumerator: - st.info( - "Enumerator statistics requires an enumerator column to be selected. " - "Go to the :material/settings: settings section above to select it." - ) - return - - # Load and validate settings using Pydantic - settings = _load_statistics_settings(settings_file) - - # Build exclusion list for numeric columns - exclude_cols = [enumerator, "consent_granted_agg_col", "completed_survey_agg_col"] - if team: - exclude_cols.append(team) - - numeric_cols = _get_numeric_columns(data, exclude_cols=exclude_cols) - - # Render UI in two columns - col1, col2 = st.columns(2) - - with col1: - statscols = _render_column_selector( - numeric_cols, settings.statscols, settings_file - ) - - with col2: - stats = _render_statistics_selector(settings.stats, settings_file) - - # Compute and display statistics - if statscols: - group_by_cols = [enumerator, team] if team else [enumerator] - stats_df = compute_enumerator_statistics( - data=data, - group_by_cols=group_by_cols, - statscols=statscols, - stats=stats, - ) - - # Build column configuration dynamically with pinning - if team: - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - team: st.column_config.TextColumn("Team", pinned=True), - } - else: - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - } - - st.dataframe( - stats_df, hide_index=True, width="stretch", column_config=column_config - ) - else: - st.info( - "No columns selected for statistics calculation.", icon=":material/info:" - ) - - -def _render_enumerator_statistics( - data: pl.DataFrame, - enumerator: str, - team: str | None, - settings_file: str, -) -> None: - """Display enumerator statistics table with team support. - - Shows configurable summary statistics for selected numeric columns - grouped by enumerator (and optionally team). - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - enumerator : str - Enumerator column name. - team : str | None - Team column name (optional). - settings_file : str - Path to settings file for saving/loading configurations. - """ - # Validate inputs - if not enumerator: - st.info( - "Enumerator statistics requires an enumerator column to be selected. " - "Go to the :material/settings: settings section above to select it." - ) - return - - _render_enumerator_statistics_table( - data=data, enumerator=enumerator, team=team, settings_file=settings_file - ) - - -def _load_statistics_overtime_settings( - settings_file: str, -) -> StatisticsOvertimeSettings: - """Load and validate statistics overtime settings from file. - - Parameters - ---------- - settings_file : str - Path to settings file. - - Returns - ------- - StatisticsOvertimeSettings - Validated statistics overtime settings. - """ - saved_settings = load_check_settings(settings_file, TAB_NAME) or {} - try: - return StatisticsOvertimeSettings(**saved_settings) - except ValueError: - # Return default settings if validation fails - return StatisticsOvertimeSettings() - - -def _render_period_selector_overtime( - settings_file: str, - default_period: str = "Week", -) -> str: - """Render time period selection widget. - - Parameters - ---------- - default_period : str - Default selected period. - settings_file : str - Path to settings file. - - Returns - ------- - str - Selected time period. - """ - options_map = { - "Day": ":material/event: Daily", - "Week": ":material/date_range: Weekly", - "Month": ":material/calendar_month: Monthly", - } - period = st.pills( - label="Select Time Period:", - options=options_map.keys(), - format_func=lambda x: options_map[x], - default=default_period, - key="project_enumerator_statistics_overtime_period_pills", - help="Select time period for aggregating statistics", - selection_mode="single", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_period_overtime"}, - ) - save_check_settings(settings_file, TAB_NAME, {"period_overtime": period}) - return period or "Day" - - -def _render_weekday_selector_overtime( - default_weekday: str, - settings_file: str, -) -> str: - """Render weekday selection widget (for weekly period). - - Parameters - ---------- - default_weekday : str - Default selected weekday. - settings_file : str - Path to settings file. - - Returns - ------- - str - Selected weekday offset code (e.g., "SUN", "MON"). - """ - default_weekday_index = WEEKDAY_NAMES.index(default_weekday) - - weekday_sel = st.selectbox( - label="Select the first day of the week", - options=WEEKDAY_NAMES, - index=default_weekday_index, - help="Select the first day of the week", - key="project_week_start_day_enumerator_overtime", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_weekstartday_overtime"}, - ) - save_check_settings(settings_file, TAB_NAME, {"weekstartday": weekday_sel}) - - return WEEKDAY_OFFSET_MAP[weekday_sel] - - -def _render_statistic_selector( - default_stat: str, - settings_file: str, -) -> str: - """Render statistic selection widget. - - Parameters - ---------- - default_stat : str - Default selected statistic. - settings_file : str - Path to settings file. - - Returns - ------- - str - Selected statistic. - """ - default_stat_index = ALLOWED_STATISTICS_OVERTIME.index(default_stat) - - stat = st.selectbox( - label="Select statistic:", - options=ALLOWED_STATISTICS_OVERTIME, - index=default_stat_index, - help="Select statistic to calculate over time", - key="enumerator_statistics_overtime_stat", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_stat_overtime"}, - ) - save_check_settings(settings_file, TAB_NAME, {"stat": stat}) - return stat - - -def _render_column_selector_single( - numeric_cols: list[str], - default_col: str | None, - settings_file: str, -) -> str | None: - """Render single column selection widget. - - Parameters - ---------- - numeric_cols : list[str] - Available numeric columns. - default_col : str | None - Default selected column. - settings_file : str - Path to settings file. - - Returns - ------- - str | None - Selected column. - """ - default_col_index = ( - numeric_cols.index(default_col) - if default_col and default_col in numeric_cols - else None - ) - - statscol = st.selectbox( - label="Select column:", - options=numeric_cols, - index=default_col_index, - help="Select column to include in statistics", - key="enumerator_statistics_overtime_column", - on_change=trigger_save, - kwargs={"state_name": TAB_NAME + "_statscol_overtime"}, - ) - save_check_settings(settings_file, TAB_NAME, {"statscol": statscol}) - return statscol - - -@st.fragment -def _render_enumerator_statistics_overtime_table( - data: pl.DataFrame, - date: str, - enumerator: str, - team: str | None, - settings_file: str, -) -> None: - """Display enumerator statistics over time table with team support. - - Shows how a specific statistic changes over time periods for each - enumerator (and optionally team) with configurable time periods and statistics. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - team : str | None - Team column name (optional). - settings_file : str - Path to settings file for saving/loading configurations. - """ - # Validate inputs - if not (enumerator and date): - return - - # Load and validate settings using Pydantic - settings = _load_statistics_overtime_settings(settings_file) - - # Build exclusion list for numeric columns - exclude_cols = [enumerator, "consent_granted_agg_col", "completed_survey_agg_col"] - if team: - exclude_cols.append(team) - - numeric_cols = _get_numeric_columns(data, exclude_cols=exclude_cols) - - # Render UI in three columns - col1, col2, col3 = st.columns([0.3, 0.2, 0.5]) - - with col1: - statscol = _render_column_selector_single( - numeric_cols, settings.statscol, settings_file - ) - - with col2: - stat = _render_statistic_selector(settings.stat, settings_file) - - with col3: - period = _render_period_selector_overtime( - settings_file, settings.period_overtime - ) - # Conditionally render weekday selector for weekly period - weekstartday = "SAT" # Default - if period == "Week": - weekstartday = _render_weekday_selector_overtime( - settings.weekstartday, settings_file - ) - - # Compute and display statistics - if statscol: - group_by_cols = [enumerator, team] if team else [enumerator] - stats_overtime_df = compute_enumerator_statistics_overtime( - data=data, - date=date, - group_by_cols=group_by_cols, - statscol=statscol, - stat=stat, - period=period, - weekstartday=weekstartday, - ) - - # Build column configuration dynamically with pinning - if team: - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - team: st.column_config.TextColumn("Team", pinned=True), - } - else: - column_config = { - enumerator: st.column_config.TextColumn("Enumerator", pinned=True), - } - - st.dataframe( - stats_overtime_df, - hide_index=True, - width="stretch", - column_config=column_config, - ) - else: - st.info( - "No column selected for statistics calculation.", icon=":material/info:" - ) - - -def _render_enumerator_statistics_overtime( - data: pl.DataFrame, - date: str, - enumerator: str, - team: str | None, - settings_file: str, -) -> None: - """Display enumerator statistics over time table with team support. - - Shows how a specific statistic changes over time periods for each - enumerator (and optionally team) with configurable time periods and statistics. - - Parameters - ---------- - data : pl.DataFrame - DataFrame containing survey data. - date : str - Date column name. - enumerator : str - Enumerator column name. - team : str | None - Team column name (optional). - settings_file : str - Path to settings file for saving/loading configurations. - """ - # Validate inputs - if not (enumerator and date): - st.info( - "Enumerator statistics over time requires a date and enumerator column to be selected. " - "Go to the :material/settings: settings section above to select them." - ) - return - - _render_enumerator_statistics_overtime_table( - data=data, - date=date, - enumerator=enumerator, - team=team, - settings_file=settings_file, - ) - - -# ============================================================================= -# Main Enumerator Report Function -# ============================================================================= - - -def enumerator_report( - project_id: str, - data: pl.DataFrame, - setting_file: str, - config: dict, - survey_columns: ColumnByType, -) -> None: - """Generate a comprehensive enumerator performance report. - - Creates a complete enumerator analysis report including: - - Overview metrics and statistics - - Comprehensive enumerator summary table - - Productivity tracking over time - - Statistical analysis across enumerators - - Time-series analysis of performance - - Parameters - ---------- - project_id : str - Unique project identifier for configuration lookup. - data : pl.DataFrame - Dataset containing survey data to analyze. - settings_file : str - Path to settings file for persisting configurations. - missing_settings_file : str - Path to missing codes configuration file. - page_num : int - Page number for configuration defaults (1-indexed). - """ - categorical_columns = survey_columns.categorical_columns - datetime_columns = survey_columns.datetime_columns - - st.title("Enumerator Report") - - demo_callout( - """ - This tab tracks enumerator performance across your survey dataset. - - It has four sections: - - **Enumerator Overview**: 8 summary metrics at the top. - - **Enumerator Summary**: Tabular breakdown by enumerator with pill-based view switching. - - **Enumerator Productivity**: Submission counts per enumerator over time. - - **Column Statistics by Enumerator**: Per-column statistics broken down by enumerator. - - **Enumerator Statistics Over Time**: Time-series chart of a selected statistic. - - **Start here**: Open the :material/settings: settings panel above to confirm your column - selections and apply consent and outcome settings. - """ - ) - - if data.is_empty(): - st.info( - "No data available for the enumerator report. " - "Please upload data to proceed." - ) - return - - config_settings = EnumeratorSettings(**config) - - enumerator_settings = enumerator_report_settings( - project_id, - setting_file, - data, - config_settings, - categorical_columns, - datetime_columns, - ) - - # get data for enumerator report - data_enum_report = duckdb_get_table( - project_id, - "enumerator_data_with_consent_outcome", - "intermediate", - ) - - if data_enum_report.is_empty(): - data_enum_report = data - - demo_callout( - """ - ##### Enumerator Overview - Eight metrics appear here in two rows of four: - - Row 1: Total enumerators, Total teams, Active enumerators (past 7 days), - % active enumerators. - - Row 2: Fewest submissions, Highest submissions, Average submissions, - Total submissions. - """ - ) - - _render_enumerator_overview_metrics( - data_enum_report, - enumerator_settings.survey_date, - enumerator_settings.enumerator, - enumerator_settings.team, - ) - - st.write("---") - st.subheader("Enumerator Summary") - - demo_callout( - """ - ##### Enumerator Summary - Use the pills above the table to switch between views. Each pill shows a different - set of columns: - - **Submissions**: First/last submission dates, submission counts (total, today, - this week, this month), unique active days. - - **Missing Data**: Percentage of missing values per enumerator. - - **Duration**: Min, max, mean, and median interview duration. - - **Form Version**: Form versions used by each enumerator. - - **Consent & Outcome**: Consent rate and completed survey rate per enumerator. - - You can select multiple pills to see combined columns side by side. - """ - ) - - _render_enumerator_summary_table( - project_id, - data_enum_report, - enumerator_settings.survey_date, - enumerator_settings.enumerator, - enumerator_settings.team, - enumerator_settings.formversion, - enumerator_settings.duration, - ) - - st.write("---") - st.subheader("Enumerator Productivity") - - demo_callout( - """ - ##### Enumerator Productivity - This section shows submission counts per enumerator over time as a table. - Use the **Daily / Weekly / Monthly** pills to change the time period granularity. - """ - ) - - _render_enumerator_productivity( - data_enum_report, - enumerator_settings.survey_date, - enumerator_settings.enumerator, - enumerator_settings.team, - setting_file, - ) - - st.write("---") - st.subheader("Column Statistics by Enumerator") - - demo_callout( - """ - ##### Column Statistics by Enumerator - Use the **column multiselect** to choose one or more numeric columns to analyse, - then use the **statistics multiselect** to choose which statistics to display - (count, mean, median, min, max, std, 25th percentile, 75th percentile). - - ##### Instructions for Demo: - Select **household_size** as the column and choose **count**, **mean**, **min**, - and **max** to check whether enumerators are recording consistent household sizes. - """ - ) - - _render_enumerator_statistics( - data_enum_report, - enumerator_settings.enumerator, - enumerator_settings.team, - setting_file, - ) - - st.write("---") - st.subheader("Enumerator Statistics Over Time") - - demo_callout( - """ - ##### Enumerator Statistics Over Time - This section renders a line chart showing how a selected statistic for a chosen - column changes over time, broken down per enumerator. Use the **column selectbox** - to pick the variable, the **statistic selectbox** to choose the metric - (e.g., mean, count, missing), and the **Daily / Weekly / Monthly** pills to set - the time granularity. - - ##### Instructions for Demo: - Select **household_size** as the column and **mean** as the statistic, then switch - to **Weekly** to see whether average household size varies by enumerator over time. - """ - ) - - _render_enumerator_statistics_overtime( - data_enum_report, - enumerator_settings.survey_date, - enumerator_settings.enumerator, - enumerator_settings.team, - setting_file, - ) - - demo_callout( - "**Next**: :material/arrow_upward: Scroll up and select the **Backcheck Analysis** tab." - ) diff --git a/src/datasure/checks/enumerator/__init__.py b/src/datasure/checks/enumerator/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/datasure/checks/enumerator/compute.py b/src/datasure/checks/enumerator/compute.py new file mode 100644 index 00000000..3429f2f8 --- /dev/null +++ b/src/datasure/checks/enumerator/compute.py @@ -0,0 +1,685 @@ +"""Pure data-computation functions for the enumerator performance module.""" + +from datetime import date as dt_date +from datetime import timedelta + +import polars as pl +import streamlit as st + +from datasure.checks import missing +from datasure.checks.enumerator.models import ( + WEEKDAY_OFFSET_TO_NUMERIC, + EnumeratorOverviewMetrics, +) +from datasure.utils.duckdb_utils import load_missing_codes_from_db + +# ============================================================================= +# Shared Time-Period Helpers +# ============================================================================= + +_LEGACY_PERIOD_ALIASES = {"Day": "Daily", "Week": "Weekly", "Month": "Monthly"} + + +def _add_time_period_column( + df: pl.DataFrame, date: str, period: str, weekstartday: str +) -> pl.DataFrame: + """Add a user-friendly "TIME PERIOD" column for the selected granularity. + + Shared by compute_enumerator_productivity and + compute_enumerator_statistics_overtime, which both bucket submissions + into daily/weekly/monthly periods the same way before aggregating. + + Parameters + ---------- + df : pl.DataFrame + DataFrame containing the date column to bucket. + date : str + Date column name. + period : str + Time period: "Daily", "Weekly", "Monthly", "Day", "Week", or "Month". + weekstartday : str + Start day of the week (e.g., "SUN", "MON") for weekly analysis. + + Returns + ------- + pl.DataFrame + `df` with a "TIME PERIOD" column added (unchanged if `period` is + not recognized). + """ + period_normalized = _LEGACY_PERIOD_ALIASES.get(period, period) + + if period_normalized == "Daily": + # Format as "Jan 1, 2025" + return df.with_columns( + pl.col(date).dt.strftime("%b %d, %Y").alias("TIME PERIOD") + ) + + if period_normalized == "Weekly": + # Calculate week start and end dates for user-friendly display + offset = WEEKDAY_OFFSET_TO_NUMERIC.get(weekstartday, 1) + + # Calculate the week start date (beginning of the week containing this date) + # weekday() returns 0=Monday, 6=Sunday + df = df.with_columns( + [ + # Calculate days since the start of the week + ((pl.col(date).dt.weekday() - offset + 7) % 7).alias( + "_days_since_week_start" + ), + ] + ) + + # Calculate week_start_date by subtracting days_since_week_start + df = df.with_columns( + [ + ( + pl.col(date) - pl.duration(days=pl.col("_days_since_week_start")) + ).alias("_week_start"), + ( + pl.col(date) + - pl.duration(days=pl.col("_days_since_week_start")) + + pl.duration(days=6) + ).alias("_week_end"), + ] + ) + + # Format as "Jan 1, 2025 to Jan 7, 2025" + return df.with_columns( + ( + pl.col("_week_start").dt.strftime("%b %d, %Y") + + " to " + + pl.col("_week_end").dt.strftime("%b %d, %Y") + ).alias("TIME PERIOD") + ) + + if period_normalized == "Monthly": + # Format as "January 2025" + return df.with_columns(pl.col(date).dt.strftime("%B %Y").alias("TIME PERIOD")) + + return df + + +# ============================================================================= +# Overview Computation Functions +# ============================================================================= + + +def compute_enumerator_overview( + data: pl.DataFrame, date: str, enumerator: str, team: str | None +) -> EnumeratorOverviewMetrics: + """Compute enumerator overview metrics. + + Calculates key metrics including total submissions, active enumerators, + team counts, and submission statistics. + + Cached for 5 minutes to improve performance for repeated calls. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + team : str | None + Team column name (optional). + + Returns + ------- + EnumeratorOverviewMetrics + Overview metrics for enumerators including counts and statistics. + """ + if data.is_empty(): + raise ValueError( + "Input data is empty. Cannot compute enumerator overview metrics." + ) + data = data.sort([enumerator, date]) + + all_submissions = data.height + + # Calculate daily submissions + data = data.with_row_index(name="TOKEN KEY") + daily_submissions_sum = data.group_by([date, enumerator]).agg( + pl.col("TOKEN KEY").count().alias("count") + ) + + # Calculate active enumerators (past 7 days) + from datetime import date as dt_date + from datetime import timedelta + + active_date_cut_off = dt_date.today() - timedelta(weeks=1) + + daily_submissions_sum = daily_submissions_sum.with_columns( + (pl.col(date).cast(pl.Date) > active_date_cut_off).alias("active") + ) + + num_active_enumerators = ( + daily_submissions_sum.filter(pl.col("active")) + .select(pl.col(enumerator).n_unique()) + .item() + ) + + num_enumerators = data[enumerator].n_unique() + num_teams = data[team].n_unique() if team else "n/a" + min_submissions = int(daily_submissions_sum["count"].min()) + max_submissions = int(daily_submissions_sum["count"].max()) + avg_submissions = int(daily_submissions_sum["count"].mean()) + + pct_active_enumerators = f"{(num_active_enumerators / num_enumerators) * 100:.0f}%" + + return EnumeratorOverviewMetrics( + all_submissions=all_submissions, + num_active_enumerators=num_active_enumerators, + num_enumerators=num_enumerators, + num_teams=num_teams, + min_submissions=min_submissions, + max_submissions=max_submissions, + avg_submissions=avg_submissions, + pct_active_enumerators=pct_active_enumerators, + ) + + +def compute_enumerator_missing_table( + data: pl.DataFrame, missing_codes_config: pl.DataFrame, group_by_col: list[str] +) -> pl.DataFrame: + """Compute missing data statistics per enumerator. + + Calculates missing data counts and percentages for each enumerator + based on provided missing codes configuration. + + Cached for 5 minutes to improve performance for repeated calls. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + missing_settings_file : str + Path to missing codes configuration file. + enumerator : str + Enumerator column name. + + Returns + ------- + pl.DataFrame + DataFrame with missing data statistics per enumerator. + """ + # define columns to exclude from missing data stats + columns_to_exclude = ["consent_granted_agg_col", "completed_survey_agg_col"] + + data_for_missing = data.select( + [col for col in data.columns if col not in columns_to_exclude] + ) + + # Metadata for missing data calculation + enum_data_missing = data_for_missing.select(group_by_col) + + # If missing_codes_config is empty, calculate only null missingness + if missing_codes_config.is_empty(): + # Calculate overall null missingness per enumerator + columns_to_check = [ + col for col in data_for_missing.columns if col not in group_by_col + ] + + # Count nulls per row and calculate percentage + missing_summary = data_for_missing.with_columns( + [ + pl.sum_horizontal( + [pl.col(col).is_null().cast(pl.Int32) for col in columns_to_check] + ).alias("_null_count"), + pl.lit(len(columns_to_check)).alias("_total_fields"), + ] + ).with_columns( + (pl.col("_null_count") / pl.col("_total_fields") * 100).alias( + "% Null values" + ) + ) + + # Group by enumerator and calculate mean missingness rate + result_df = missing_summary.group_by(group_by_col, maintain_order=True).agg( + pl.col("% Null values").mean() + ) + + return result_df + + # If missing_codes_config is provided, calculate missingness by category + # Get missing code pairs from the config + missing_code_pairs = missing._get_missing_code_pairs(missing_codes_config) + + # Compute missing data with paired encoding + missing_data_encoded = missing._compute_missing_data_paired( + data_for_missing, missing_codes_config + ) + + # Get columns to check (exclude enumerator column) + columns_to_check = [ + col for col in missing_data_encoded.columns if col not in group_by_col + ] + total_fields = len(columns_to_check) + + # Calculate counts for each missing category per row + agg_expressions = [] + + # Add null values count (encoded as 1) + agg_expressions.append( + pl.sum_horizontal( + [(pl.col(col) == 1).cast(pl.Int32) for col in columns_to_check] + ).alias("_null_count") + ) + + # Add count for each special missing code category (starting from 2) + for idx, label in enumerate(missing_code_pairs.keys(), start=2): + agg_expressions.append( + pl.sum_horizontal( + [(pl.col(col) == idx).cast(pl.Int32) for col in columns_to_check] + ).alias(f"_{label}_count") + ) + + # Add total missing count (any value > 0) + agg_expressions.append( + pl.sum_horizontal( + [(pl.col(col) > 0).cast(pl.Int32) for col in columns_to_check] + ).alias("_total_missing_count") + ) + + # Add total fields count + agg_expressions.append(pl.lit(total_fields).alias("_total_fields")) + + # Apply the aggregations + missing_counts = missing_data_encoded.select( + [pl.col(group_by_col)] + agg_expressions + ) + + # Calculate percentages + percentage_expressions = [] + + # Null values percentage + percentage_expressions.append( + (pl.col("_null_count") / pl.col("_total_fields") * 100).alias("% Null values") + ) + + # Special missing code category percentages + for label in missing_code_pairs: + percentage_expressions.append( + (pl.col(f"_{label}_count") / pl.col("_total_fields") * 100).alias( + f"% {label}" + ) + ) + + # Total missing percentage + percentage_expressions.append( + (pl.col("_total_missing_count") / pl.col("_total_fields") * 100).alias( + "% Total Missing" + ) + ) + + missing_with_percentages = missing_counts.with_columns(percentage_expressions) + + # Group by enumerator and calculate mean percentages + final_agg_expressions = [pl.col("% Null values").mean()] + + for label in missing_code_pairs: + final_agg_expressions.append(pl.col(f"% {label}").mean()) + + final_agg_expressions.append(pl.col("% Total Missing").mean()) + + # drop enumerator column from missing_with_percentages + missing_with_percentages = missing_with_percentages.select( + [col for col in missing_with_percentages.columns if col not in group_by_col] + ) + # merge missing_with_percentages with enumerator column + missing_with_percentages = pl.concat( + [enum_data_missing, missing_with_percentages], how="horizontal" + ) + + result_df = missing_with_percentages.group_by( + group_by_col, maintain_order=True + ).agg(final_agg_expressions) + + return result_df + + +def compute_enumerator_summary( + project_id: str, + data: pl.DataFrame, + date: str, + enumerator: str, + team: str | None, + formversion: str | None, + duration: str | None, +) -> pl.DataFrame: + """Compute comprehensive enumerator summary statistics. + + Calculates submission counts, date ranges, duration statistics, + form version tracking, consent rates, outcome rates, and missing data + patterns for each enumerator. + + Cached for 5 minutes to improve performance for repeated calls. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + missing_settings_file : str + Path to missing codes configuration file. + date : str + Date column name. + enumerator : str + Enumerator column name. + formdef_version : str | None + Form version column name (optional). + duration : str | None + Duration column name (optional). + consent : str | None + Consent column name (optional). + consent_vals : list[str] | None + List of values indicating valid consent (optional). + outcome : str | None + Outcome column name (optional). + outcome_vals : list[str] | None + List of values indicating completed surveys (optional). + + Returns + ------- + pl.DataFrame + Comprehensive summary DataFrame with enumerator statistics. + """ + group_by_cols = [enumerator, team] if team else [enumerator] + # Format date column + df = data.with_columns(pl.col(date).dt.strftime("%b %d, %Y").alias(date)) + + # Basic summary aggregations + summary_df = df.group_by(group_by_cols, maintain_order=True).agg( + [ + pl.col(date).min().alias("first submission"), + pl.col(date).max().alias("last submission"), + pl.col(date).count().alias("# submissions"), + pl.col(date).n_unique().alias("# unique dates"), + ] + ) + + # Calculate time-based submissions + today = dt_date.today() + start_of_week = today - timedelta(days=today.weekday()) + start_of_month = today.replace(day=1) + + today_str = today.strftime("%b %d, %Y") + week_str = start_of_week.strftime("%b %d, %Y") + month_str = start_of_month.strftime("%b %d, %Y") + + df = df.with_columns( + [ + (pl.col(date) == today_str).alias("submitted_today"), + (pl.col(date) >= week_str).alias("submitted_this_week"), + (pl.col(date) >= month_str).alias("submitted_this_month"), + ] + ) + + lagged_df = df.group_by(group_by_cols, maintain_order=True).agg( + [ + pl.col("submitted_today").sum().alias("# submissions today"), + pl.col("submitted_this_week").sum().alias("# submissions this week"), + pl.col("submitted_this_month").sum().alias("# submissions this month"), + ] + ) + + summary_df = summary_df.join(lagged_df, on=group_by_cols, how="left") + + # Add missing data statistics + missing_settings_file = load_missing_codes_from_db(project_id) + enumerator_missing_df = compute_enumerator_missing_table( + data, missing_settings_file, group_by_cols + ) + summary_df = summary_df.join(enumerator_missing_df, on=group_by_cols, how="left") + + # Add duration statistics if available + if duration: + duration_df = df.group_by(group_by_cols, maintain_order=True).agg( + [ + pl.col(duration).min().alias("min duration"), + pl.col(duration).mean().alias("mean duration"), + pl.col(duration).median().alias("median duration"), + pl.col(duration).max().alias("max duration"), + ] + ) + summary_df = summary_df.join(duration_df, on=group_by_cols, how="left") + + # Add form version statistics if available + if formversion: + # Get latest form version per date + formdef_outdated = df.group_by(date, maintain_order=True).agg( + pl.col(formversion).max().alias("latest daily form version") + ) + + df = df.join(formdef_outdated, on=date, how="left") + df = df.with_columns( + (pl.col(formversion) != pl.col("latest daily form version")).alias( + "outdated_form_version" + ) + ) + + formdef_outdated_df = df.group_by(group_by_cols, maintain_order=True).agg( + pl.col("outdated_form_version").sum().alias("# of outdated form versions") + ) + + formdef_df = df.group_by(group_by_cols, maintain_order=True).agg( + [ + pl.col(formversion).n_unique().alias("# form versions"), + pl.col(formversion).max().alias("latest form version"), + ] + ) + + latest_enum_formversion = df.group_by(group_by_cols, maintain_order=True).agg( + pl.col(formversion).max().alias("last form version") + ) + + summary_df = summary_df.join(formdef_df, on=group_by_cols, how="left") + summary_df = summary_df.join(formdef_outdated_df, on=group_by_cols, how="left") + summary_df = summary_df.join( + latest_enum_formversion, on=group_by_cols, how="left" + ) + + # Add consent statistics if available + if "consent_granted_agg_col" in df.columns: + consent_df = df.group_by(group_by_cols, maintain_order=True).agg( + pl.col("consent_granted_agg_col").mean().alias("% consent") + ) + summary_df = summary_df.join(consent_df, on=group_by_cols, how="left") + + # Add outcome statistics if available + if "completed_survey_agg_col" in df.columns: + outcome_df = df.group_by(group_by_cols, maintain_order=True).agg( + pl.col("completed_survey_agg_col").mean().alias("% completed survey") + ) + summary_df = summary_df.join(outcome_df, on=group_by_cols, how="left") + + return summary_df + + +# ============================================================================= +# Productivity Computation Functions +# ============================================================================= + + +def compute_enumerator_productivity( + data: pl.DataFrame, + date: str, + group_by_cols: list[str], + period: str, + weekstartday: str, +) -> pl.DataFrame: + """Compute enumerator productivity over time. + + Analyzes submission counts by enumerator across time periods (daily, + weekly, or monthly). + + Cached for 5 minutes to improve performance for repeated calls. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + period : str + Time period: "Daily", "Weekly", "Monthly", "Day", "Week", or "Month". + weekstartday : str + Start day of the week (e.g., "SUN", "MON") for weekly analysis. + + Returns + ------- + pl.DataFrame + Pivoted DataFrame with enumerators as rows and time periods as columns. + """ + prod_df = data.clone() + prod_df = _add_time_period_column(prod_df, date, period, weekstartday) + + # Count submissions per period and enumerator + prod_df = prod_df.with_row_index(name="TOKEN KEY") + prod_res = prod_df.group_by( + ["TIME PERIOD"] + group_by_cols, maintain_order=True + ).agg(pl.col("TOKEN KEY").count().alias("submissions")) + + # Pivot to wide format + prod_res = prod_res.pivot( + index=group_by_cols, + on="TIME PERIOD", + values="submissions", + ).fill_null(0) + + return prod_res + + +# ============================================================================= +# Statistics Computation Functions +# ============================================================================= + + +@st.cache_data(ttl=300) +def compute_enumerator_statistics( + data: pl.DataFrame, + group_by_cols: list[str], + statscols: list[str], + stats: list[str], +) -> pl.DataFrame: + """Compute enumerator statistics across specified columns. + + Calculates summary statistics (mean, median, std, etc.) for numeric + columns grouped by enumerator (and optionally team). + + Cached for 5 minutes to improve performance for repeated calls. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + group_by_cols : list[str] + List of columns to group by (e.g., ["enumerator"] or ["enumerator", "team"]). + statscols : list[str] + List of columns to compute statistics on. + stats : list[str] + List of statistics to compute (e.g., ["mean", "median", "std"]). + + Returns + ------- + pl.DataFrame + DataFrame with enumerators (and teams) and computed statistics. + """ + # Map stat names to Polars expressions + stat_mapping = { + "count": "count", + "min": "min", + "mean": "mean", + "median": "median", + "max": "max", + "std": "std", + "25th percentile": "quantile", + "75th percentile": "quantile", + } + + agg_exprs = [] + for col in statscols: + for stat in stats: + if stat == "25th percentile": + agg_exprs.append(pl.col(col).quantile(0.25).alias(f"{col}_{stat}")) + elif stat == "75th percentile": + agg_exprs.append(pl.col(col).quantile(0.75).alias(f"{col}_{stat}")) + else: + method = stat_mapping.get(stat, stat) + agg_exprs.append(getattr(pl.col(col), method)().alias(f"{col}_{stat}")) + + stats_res = data.group_by(group_by_cols, maintain_order=True).agg(agg_exprs) + + return stats_res + + +def compute_enumerator_statistics_overtime( + data: pl.DataFrame, + date: str, + group_by_cols: list[str], + statscol: str, + stat: str, + period: str, + weekstartday: str, +) -> pl.DataFrame: + """Compute enumerator statistics over time for a specific column. + + Analyzes how a specific statistic changes over time periods for each + enumerator (and optionally team). + + Cached for 5 minutes to improve performance for repeated calls. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + group_by_cols : list[str] + List of columns to group by (e.g., ["enumerator"] or ["enumerator", "team"]). + statscol : str + Column to compute statistics on. + stat : str + Statistic to compute (e.g., "mean", "median", "missing"). + period : str + Time period: "Daily", "Weekly", "Monthly", "Day", "Week", or "Month". + weekstartday : str + Start day of the week for weekly analysis. + + Returns + ------- + pl.DataFrame + Pivoted DataFrame with enumerators (and teams) as rows and time + periods as columns. + """ + stats_overtime_df = data.select([date] + group_by_cols + [statscol]).clone() + stats_overtime_df = _add_time_period_column( + stats_overtime_df, date, period, weekstartday + ) + + # Calculate statistic + if stat == "missing": + stats_overtime_res = stats_overtime_df.group_by( + ["TIME PERIOD"] + group_by_cols, maintain_order=True + ).agg(pl.col(statscol).is_null().mean().alias("_STAT")) + elif stat == "25th percentile": + stats_overtime_res = stats_overtime_df.group_by( + ["TIME PERIOD"] + group_by_cols, maintain_order=True + ).agg(pl.col(statscol).quantile(0.25).alias("_STAT")) + elif stat == "75th percentile": + stats_overtime_res = stats_overtime_df.group_by( + ["TIME PERIOD"] + group_by_cols, maintain_order=True + ).agg(pl.col(statscol).quantile(0.75).alias("_STAT")) + else: + stats_overtime_res = stats_overtime_df.group_by( + ["TIME PERIOD"] + group_by_cols, maintain_order=True + ).agg(getattr(pl.col(statscol), stat)().alias("_STAT")) + + # Pivot to wide format + stats_overtime_res = stats_overtime_res.pivot( + index=group_by_cols, on="TIME PERIOD", values="_STAT" + ) + + return stats_overtime_res diff --git a/src/datasure/checks/enumerator/models.py b/src/datasure/checks/enumerator/models.py new file mode 100644 index 00000000..9614afdc --- /dev/null +++ b/src/datasure/checks/enumerator/models.py @@ -0,0 +1,267 @@ +"""Pydantic models and constants for the enumerator performance module.""" + +from pydantic import BaseModel, Field, field_validator + +TAB_NAME: str = "enumerators" + + +# ============================================================================= +# Pydantic Models for Data Validation +# ============================================================================= + + +class EnumeratorSettings(BaseModel): + """Settings for enumerator report configuration. + + Attributes + ---------- + date : str | None + Column name containing survey submission date. + survey_id : str | None + Column name containing survey ID. + enumerator : str | None + Column name containing enumerator identifier (required). + formdef_version : str | None + Column name containing form version. + duration : str | None + Column name containing survey duration in seconds. + team : str | None + Column name containing team identifier. + consent : str | None + Column name containing consent status. + consent_vals : list[str] | None + List of values indicating valid consent. + outcome : str | None + Column name containing survey outcome status. + outcome_vals : list[str] | None + List of values indicating completed surveys. + """ + + survey_key: str | None = Field(None, description="Survey key column") + survey_id: str | None = Field(..., min_length=1, description="Survey ID column") + survey_date: str | None = Field(None, description="Survey date column") + enumerator: str | None = Field(None, description="Enumerator ID column") + formversion: str | None = Field(None, description="Form version column") + duration: str | None = Field(None, description="Duration column") + duration_unit: str = Field("default='seconds'", description="Duration unit") + team: str | None = Field(None, description="Team identifier column") + + +class ConsentOutcomeSettings(BaseModel): + """Settings for consent and outcome configuration. + + Attributes + ---------- + consent : str | None + Column name containing consent status. + consent_vals : list[str] | None + List of values indicating valid consent. + outcome : str | None + Column name containing survey outcome status. + outcome_vals : list[str] | None + List of values indicating completed surveys. + """ + + consent: str | None = Field(None, description="Consent status column") + consent_vals: list[str] | None = Field(None, description="Valid consent values") + outcome: str | None = Field(None, description="Outcome status column") + outcome_vals: list[str] | None = Field(None, description="Completed survey values") + + +class ProductivitySettings(BaseModel): + """Settings for productivity analysis configuration. + + Attributes + ---------- + view_option : str + Time period for analysis: Daily, Weekly, or Monthly. + weekstartday : str + First day of the week for weekly analysis. + """ + + view_option: str = Field(default="Daily", description="Time period view") + weekstartday: str = Field(default="Monday", description="Week start day") + + @field_validator("view_option") + @classmethod + def validate_view_option(cls, v: str) -> str: + """Validate view option is one of the allowed values.""" + allowed = ["Daily", "Weekly", "Monthly"] + if v not in allowed: + raise ValueError(f"view_option must be one of {allowed}") + return v + + @field_validator("weekstartday") + @classmethod + def validate_weekstartday(cls, v: str) -> str: + """Validate weekstartday is one of the allowed values.""" + allowed = [ + "Monday", + "Tuesday", + "Wednesday", + "Thursday", + "Friday", + "Saturday", + "Sunday", + ] + if v not in allowed: + raise ValueError(f"weekstartday must be one of {allowed}") + return v + + +# Constants for statistics options +ALLOWED_STATISTICS = [ + "count", + "min", + "mean", + "median", + "max", + "std", + "25th percentile", + "75th percentile", +] +ALLOWED_STATISTICS_OVERTIME = ALLOWED_STATISTICS + ["missing"] +ALLOWED_TIME_PERIODS = ["Daily", "Weekly", "Monthly"] +WEEKDAY_NAMES = [ + "Monday", + "Tuesday", + "Wednesday", + "Thursday", + "Friday", + "Saturday", + "Sunday", +] + +# Maps weekday names to offset codes used in computation +WEEKDAY_OFFSET_MAP = { + "Monday": "SUN", + "Tuesday": "MON", + "Wednesday": "TUE", + "Thursday": "WED", + "Friday": "THU", + "Saturday": "FRI", + "Sunday": "SAT", +} + +# Maps offset codes to numeric values for week calculations +WEEKDAY_OFFSET_TO_NUMERIC = { + "SUN": 0, + "MON": 1, + "TUE": 2, + "WED": 3, + "THU": 4, + "FRI": 5, + "SAT": 6, +} + + +class StatisticsSettings(BaseModel): + """Settings for statistics analysis configuration. + + Attributes + ---------- + statscols : list[str] | None + Columns to compute statistics on. + stats : list[str] + Statistics to compute (count, mean, median, etc.). + """ + + statscols: list[str] | None = Field(None, description="Columns for statistics") + stats: list[str] = Field( + default=["count", "mean"], description="Statistics to compute" + ) + + @field_validator("stats") + @classmethod + def validate_stats(cls, v: list[str]) -> list[str]: + """Validate that statistics are from allowed list.""" + for stat in v: + if stat not in ALLOWED_STATISTICS: + raise ValueError( + f"Invalid statistic: {stat}. Must be one of {ALLOWED_STATISTICS}" + ) + return v + + +class StatisticsOvertimeSettings(BaseModel): + """Settings for statistics over time analysis configuration. + + Attributes + ---------- + period : str + Time period for analysis (Daily, Weekly, Monthly). + weekstartday : str + First day of the week for weekly analysis. + stat : str + Statistic to compute over time. + statscol : str | None + Column to compute statistics on. + """ + + period_overtime: str = Field(default="Week", description="Time period for analysis") + weekstartday: str = Field(default="Monday", description="Week start day") + stat: str = Field(default="count", description="Statistic to compute") + statscol: str | None = Field(None, description="Column for statistics") + + @field_validator("period_overtime") + @classmethod + def validate_period(cls, v: str) -> str: + """Validate period is from allowed list.""" + if v not in ALLOWED_TIME_PERIODS: + raise ValueError( + f"Invalid period: {v}. Must be one of {ALLOWED_TIME_PERIODS}" + ) + return v + + @field_validator("weekstartday") + @classmethod + def validate_weekstartday(cls, v: str) -> str: + """Validate weekstartday is from allowed list.""" + if v not in WEEKDAY_NAMES: + raise ValueError( + f"Invalid weekstartday: {v}. Must be one of {WEEKDAY_NAMES}" + ) + return v + + @field_validator("stat") + @classmethod + def validate_stat(cls, v: str) -> str: + """Validate stat is from allowed list.""" + if v not in ALLOWED_STATISTICS_OVERTIME: + raise ValueError( + f"Invalid statistic: {v}. Must be one of {ALLOWED_STATISTICS_OVERTIME}" + ) + return v + + +class EnumeratorOverviewMetrics(BaseModel): + """Metrics for enumerator overview. + + Attributes + ---------- + all_submissions : int + Total number of submissions. + num_active_enumerators : int + Number of enumerators active in past 7 days. + num_enumerators : int + Total number of enumerators. + num_teams : int | str + Number of teams or 'n/a' if not available. + min_submissions : int + Minimum daily submissions. + max_submissions : int + Maximum daily submissions. + avg_submissions : int + Average daily submissions. + pct_active_enumerators : str + Percentage of active enumerators formatted as string. + """ + + all_submissions: int = Field(ge=0) + num_active_enumerators: int = Field(ge=0) + num_enumerators: int = Field(ge=0) + num_teams: int | str + min_submissions: int = Field(ge=0) + max_submissions: int = Field(ge=0) + avg_submissions: int = Field(ge=0) + pct_active_enumerators: str diff --git a/src/datasure/checks/enumerator/report_ui.py b/src/datasure/checks/enumerator/report_ui.py new file mode 100644 index 00000000..7037c83e --- /dev/null +++ b/src/datasure/checks/enumerator/report_ui.py @@ -0,0 +1,1257 @@ +"""Report-rendering UI for the enumerator performance report.""" + +from typing import Literal + +import polars as pl +import streamlit as st + +from datasure.checks.enumerator.compute import ( + compute_enumerator_overview, + compute_enumerator_productivity, + compute_enumerator_statistics, + compute_enumerator_statistics_overtime, + compute_enumerator_summary, +) +from datasure.checks.enumerator.models import ( + ALLOWED_STATISTICS, + ALLOWED_STATISTICS_OVERTIME, + TAB_NAME, + WEEKDAY_NAMES, + WEEKDAY_OFFSET_MAP, + EnumeratorOverviewMetrics, + EnumeratorSettings, + StatisticsOvertimeSettings, + StatisticsSettings, +) +from datasure.checks.enumerator.settings_ui import enumerator_report_settings +from datasure.utils.dataframe_utils import ColumnByType +from datasure.utils.duckdb_utils import duckdb_get_table +from datasure.utils.navigations_utils import demo_callout +from datasure.utils.settings_utils import ( + load_check_settings, + save_check_settings, + trigger_save, +) + +# ============================================================================= +# Display Functions - Overview +# ============================================================================= + + +def _render_enumerator_overview_metrics( + data: pl.DataFrame, date: str, enumerator: str, team: str | None +) -> None: + """Display enumerator overview metrics. + + Shows key metrics including total submissions, active enumerators, + team counts, and submission statistics in a grid layout. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + team : str | None + Team column name (optional). + """ + if not (enumerator and date): + st.info( + "Enumerator overview requires a date and enumerator column to be selected. " + "Go to the :material/settings: settings section above to select them." + ) + return + + metrics: EnumeratorOverviewMetrics = compute_enumerator_overview( + data, date, enumerator, team + ) + + tc1, tc2, tc3, tc4 = st.columns(4, border=True) + num_enumerators_formatted = ( + f"{metrics.num_enumerators:,}" + if isinstance(metrics.num_enumerators, int) + else metrics.num_enumerators + ) + tc1.metric( + r"\# of enumerators", + num_enumerators_formatted, + help="Total unique enumerators in the dataset", + ) + num_teams_formatted = ( + f"{metrics.num_teams:,}" + if isinstance(metrics.num_teams, int) + else metrics.num_teams + ) + tc2.metric( + r"\# of teams", num_teams_formatted, help="Total unique teams in the dataset" + ) + num_active_enumerators_formatted = f"{metrics.num_active_enumerators:,}" + tc3.metric( + r"\# of Active enumerators (past 7 days)", + num_active_enumerators_formatted, + help="Number of enumerators with submissions in the past 7 days", + ) + pct_active_enumerators_formatted = f"{metrics.pct_active_enumerators}" + tc4.metric( + "% of active enumerator (past 7 days)", + pct_active_enumerators_formatted, + help="Percentage of enumerators active in the past 7 days", + ) + + bc1, bc2, bc3, bc4 = st.columns(4, border=True) + min_submissions_formatted = f"{metrics.min_submissions:,}" + bc1.metric( + "Fewest enumerator submissions", + min_submissions_formatted, + help="Minimum number of submissions by any enumerator", + ) + max_submissions_formatted = f"{metrics.max_submissions:,}" + bc2.metric( + "Highest enumerator submissions", + max_submissions_formatted, + help="Maximum number of submissions by any enumerator", + ) + avg_submissions_formatted = f"{metrics.avg_submissions:,}" + bc3.metric( + "Average enumerator submissions", + avg_submissions_formatted, + help="Average number of submissions per enumerator", + ) + all_submissions_formatted = f"{metrics.all_submissions:,}" + bc4.metric( + "Total survey submissions", + all_submissions_formatted, + help="Total number of survey submissions in the dataset", + ) + + +@st.fragment +def _render_enumerator_summary_table( + project_id: str, + data: pl.DataFrame, + date: str, + enumerator: str, + team: str | None, + formversion: str | None, + duration: str | None, +) -> None: + """Display enumerator summary table. + + Shows comprehensive enumerator statistics including submission counts, + duration, missing data, consent rates, and outcome rates with styled + formatting. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + missing_settings_file : str + Path to missing codes configuration file. + date : str + Date column name. + enumerator : str + Enumerator column name. + formdef_version : str | None + Form version column name (optional). + duration : str | None + Duration column name (optional). + consent : str | None + Consent column name (optional). + consent_vals : list[str] | None + Valid consent values (optional). + outcome : str | None + Outcome column name (optional). + outcome_vals : list[str] | None + Completed survey values (optional). + """ + if not (enumerator and date): + st.info( + "Enumerator summary requires a date and enumerator column to be selected. " + "Go to the :material/settings: settings section above to select them." + ) + return + + summary_df = compute_enumerator_summary( + project_id, + data, + date, + enumerator, + team, + formversion, + duration, + ) + + options_map = { + "submissions": ":material/arrow_upload_progress: Submissions", + "missing": ":material/incomplete_circle: Missing Data", + "duration": ":material/timer: Duration", + "formversion": ":material/difference: Form Version", + "consent_outcome": ":material/check_circle: Consent & Outcome", + } + with st.container(horizontal_alignment="left"): + show_info = st.pills( + "Select Summary Information to Display", + options=options_map.keys(), + format_func=lambda x: options_map[x], + key="show_info_enumerator", + help="Select which summary information to display in the table", + selection_mode="multi", + ) + + # Define column groups + column_groups = { + "submissions": [ + "first submission", + "last submission", + "# submissions", + "# unique dates", + "# submissions today", + "# submissions this week", + "# submissions this month", + ], + "missing": [ + col + for col in summary_df.columns + if "%" in col + and ( + "Null" in col + or "Missing" in col + or any( + keyword in col + for keyword in ["Don't Know", "Refuse", "Not Applicable"] + ) + ) + ], + "duration": [ + "min duration", + "mean duration", + "median duration", + "max duration", + ], + "formversion": [ + "# form versions", + "latest form version", + "last form version", + "# of outdated form versions", + ], + "consent_outcome": [ + "% consent", + "% completed survey", + ], + } + + # Always include enumerator and # submissions + columns_to_show = ( + [enumerator, team, "# submissions"] if team else [enumerator, "# submissions"] + ) + + # Filter columns based on selection + if show_info: + # Add columns from selected categories + for category in show_info: + columns_to_show.extend( + [ + col + for col in column_groups[category] + if col in summary_df.columns and col not in columns_to_show + ] + ) + + # Filter the dataframe + filtered_df = summary_df.select(columns_to_show) + else: + # Show all columns if nothing is selected + filtered_df = summary_df + + # Display using Streamlit's native dataframe display + # create column config for enumerator and team conditionally + # Build column configuration dynamically + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + } + + # Add team column if available + if team: + column_config[team] = st.column_config.TextColumn("Team", pinned=True) + + # Add remaining columns + column_config.update( + { + "# submissions": st.column_config.NumberColumn( + "# of Submissions", format="%d", pinned=True + ), + "# unique dates": st.column_config.NumberColumn("# of Days", format="%d"), + "# submissions today": st.column_config.NumberColumn( + "# submitted Today", format="%d" + ), + "# submissions this week": st.column_config.NumberColumn( + "# submitted This Week", format="%d" + ), + "# submissions this month": st.column_config.NumberColumn( + "# submitted This Month", format="%d" + ), + "% Null values": st.column_config.NumberColumn( + "% Null Values", format="%.2f%%" + ), + "% Total Missing": st.column_config.NumberColumn( + "% Total Missing", format="%.2f%%" + ), + "% consent": st.column_config.NumberColumn("% Consent", format="%.2f%%"), + "% completed survey": st.column_config.NumberColumn( + "% Completed", format="%.2f%%" + ), + "min duration": st.column_config.NumberColumn( + "Min Duration (s)", format="%.2f" + ), + "mean duration": st.column_config.NumberColumn( + "Mean Duration (s)", format="%.2f" + ), + "median duration": st.column_config.NumberColumn( + "Median Duration (s)", format="%.2f" + ), + "max duration": st.column_config.NumberColumn( + "Max Duration (s)", format="%.2f" + ), + } + ) + + st.dataframe( + filtered_df, + hide_index=True, + width="stretch", + column_config=column_config, + ) + + +# ============================================================================= +# Display Functions - Productivity +# ============================================================================= + + +def _render_enumerator_productivity( + data: pl.DataFrame, + date: str, + enumerator: str, + team: str | None, + settings_file: str, +) -> None: + """Display enumerator productivity table. + + Shows submission counts by enumerator over time with configurable + time periods (daily, weekly, monthly). + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + settings_file : str + Path to settings file for saving/loading configurations. + """ + if not (enumerator and date): + st.info( + "Enumerator productivity requires a date and enumerator column to be selected. " + "Go to the :material/settings: settings section above to select them." + ) + return + + _render_enumerator_productivity_table(data, date, enumerator, team, settings_file) + + +@st.fragment +def _render_enumerator_productivity_table( + data: pl.DataFrame, + date: str, + enumerator: str, + team: str | None, + settings_file: str, +) -> None: + """Display enumerator productivity table. + Shows submission counts by enumerator over time with configurable + time periods (daily, weekly, monthly). + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + settings_file : str + Path to settings file for saving/loading configurations. + """ + time_period = _render_time_period_selector(settings_file, tab_name=TAB_NAME) + if time_period == "Week": + weekstartday = _render_weekday_selector(settings_file, tab_name=TAB_NAME) + else: + weekstartday = "MON" # Default value, not used for non-weekly periods + + group_by_cols = [enumerator, team] if team else [enumerator] + productivity_df = compute_enumerator_productivity( + data, date, group_by_cols, time_period, weekstartday + ) + + if team: + # Build column configuration dynamically + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + team: st.column_config.TextColumn("Team", pinned=True), + } + else: + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + } + + column_config.update( + { + col: st.column_config.NumberColumn(col, format="%d") + for col in productivity_df.columns + if col not in group_by_cols + } + ) + + st.dataframe( + productivity_df, hide_index=True, width="stretch", column_config=column_config + ) + + +def _render_time_period_selector( + settings_file: str, + tab_name: str = TAB_NAME, +) -> Literal["Day", "Week", "Month"]: + """Render time period selector widget using pills interface. + + Displays a pills widget allowing users to choose the time aggregation period + for productivity analysis (Day, Week, or Month). + + Parameters + ---------- + settings_file : str + Path to settings file for saving/loading configurations. + tab_name : str + Name of the tab for settings storage (default: TAB_NAME). + + Returns + ------- + Literal["Day", "Week", "Month"] + Selected time period. + """ + options_map = { + "Day": ":material/event: Daily", + "Week": ":material/date_range: Weekly", + "Month": ":material/calendar_month: Monthly", + } + + saved_settings = load_check_settings(settings_file, tab_name) or {} + default_time_period = saved_settings.get( + "time_period_enumerator_productivity", "Day" + ) + + with st.container(horizontal_alignment="left"): + time_period = st.pills( + label="Time Period", + options=options_map.keys(), + format_func=lambda x: options_map[x], + key="time_period_enumerator_productivity_key", + default=default_time_period, + help="Select time period for aggregating productivity", + selection_mode="single", + on_change=trigger_save, + kwargs={"state_name": tab_name + "_time_period"}, + ) + save_check_settings(settings_file, tab_name, {"time_period": time_period}) + + return time_period or "Day" + + +def _render_weekday_selector( + settings_file: str, + tab_name: str = TAB_NAME, +) -> str: + """Render weekday selector widget for productivity analysis. + + Displays a selectbox allowing users to choose the first day of the week + for weekly productivity calculations. + + Parameters + ---------- + settings_file : str + Path to settings file for saving/loading configurations. + tab_name : str + Name of the tab for settings storage (default: TAB_NAME). + + Returns + ------- + str + Weekday offset code (e.g., "SUN", "MON") for calculations. + """ + saved_settings = load_check_settings(settings_file, tab_name) or {} + default_weekstartday_sel = saved_settings.get( + "weekstartday_enumerator_productivity", "Monday" + ) + default_weekstartday_sel_index = WEEKDAY_NAMES.index(default_weekstartday_sel) + + cl1, _ = st.columns([1, 3]) + with cl1: + weekstartday_sel = st.selectbox( + label="Select the first day of the week", + options=WEEKDAY_NAMES, + index=default_weekstartday_sel_index, + key="week_start_day_enumerator_productivity_key", + help="Select the first day of the week", + on_change=trigger_save, + kwargs={"state_name": tab_name + "_weekstartday"}, + ) + save_check_settings(settings_file, tab_name, {"weekstartday": weekstartday_sel}) + + return WEEKDAY_OFFSET_MAP[weekstartday_sel] + + +# ============================================================================= +# Display Functions - Statistics +# ============================================================================= + + +def _load_statistics_settings(settings_file: str) -> StatisticsSettings: + """Load and validate statistics settings from file. + + Parameters + ---------- + settings_file : str + Path to settings file. + + Returns + ------- + StatisticsSettings + Validated statistics settings. + """ + saved_settings = load_check_settings(settings_file, TAB_NAME) or {} + try: + return StatisticsSettings(**saved_settings) + except ValueError: + # Return default settings if validation fails + return StatisticsSettings() + + +def _get_numeric_columns( + data: pl.DataFrame, exclude_cols: list[str] | None = None +) -> list[str]: + """Extract numeric column names from DataFrame. + + Parameters + ---------- + data : pl.DataFrame + DataFrame to extract columns from. + exclude_cols : list[str] | None + Columns to exclude from the result. + + Returns + ------- + list[str] + List of numeric column names. + """ + exclude_cols = exclude_cols or [] + return [ + col + for col in data.columns + if data[col].dtype in pl.NUMERIC_DTYPES and col not in exclude_cols + ] + + +def _render_column_selector( + numeric_cols: list[str], + default_cols: list[str] | None, + settings_file: str, +) -> list[str]: + """Render column selection widget. + + Parameters + ---------- + numeric_cols : list[str] + Available numeric columns. + default_cols : list[str] | None + Default selected columns. + settings_file : str + Path to settings file. + + Returns + ------- + list[str] + Selected columns. + """ + selected_cols = st.multiselect( + label="Select columns:", + options=numeric_cols, + default=default_cols, + help="Select columns to include in statistics", + key="selected_columns_enumerator", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_statscols"}, + ) + save_check_settings(settings_file, TAB_NAME, {"statscols": selected_cols}) + return selected_cols + + +def _render_statistics_selector( + default_stats: list[str], + settings_file: str, +) -> list[str]: + """Render statistics selection widget. + + Parameters + ---------- + default_stats : list[str] + Default selected statistics. + settings_file : str + Path to settings file. + + Returns + ------- + list[str] + Selected statistics. + """ + selected_stats = st.multiselect( + "Select statistics:", + options=ALLOWED_STATISTICS, + default=default_stats, + help="Select statistics to calculate", + key="statistics_options_enumerator", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_stats"}, + ) + save_check_settings(settings_file, TAB_NAME, {"stats": selected_stats}) + return selected_stats + + +@st.fragment +def _render_enumerator_statistics_table( + data: pl.DataFrame, + enumerator: str, + team: str | None, + settings_file: str, +) -> None: + """Display enumerator statistics table with team support. + + Shows configurable summary statistics for selected numeric columns + grouped by enumerator (and optionally team). + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + enumerator : str + Enumerator column name. + team : str | None + Team column name (optional). + settings_file : str + Path to settings file for saving/loading configurations. + """ + # Validate inputs + if not enumerator: + st.info( + "Enumerator statistics requires an enumerator column to be selected. " + "Go to the :material/settings: settings section above to select it." + ) + return + + # Load and validate settings using Pydantic + settings = _load_statistics_settings(settings_file) + + # Build exclusion list for numeric columns + exclude_cols = [enumerator, "consent_granted_agg_col", "completed_survey_agg_col"] + if team: + exclude_cols.append(team) + + numeric_cols = _get_numeric_columns(data, exclude_cols=exclude_cols) + + # Render UI in two columns + col1, col2 = st.columns(2) + + with col1: + statscols = _render_column_selector( + numeric_cols, settings.statscols, settings_file + ) + + with col2: + stats = _render_statistics_selector(settings.stats, settings_file) + + # Compute and display statistics + if statscols: + group_by_cols = [enumerator, team] if team else [enumerator] + stats_df = compute_enumerator_statistics( + data=data, + group_by_cols=group_by_cols, + statscols=statscols, + stats=stats, + ) + + # Build column configuration dynamically with pinning + if team: + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + team: st.column_config.TextColumn("Team", pinned=True), + } + else: + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + } + + st.dataframe( + stats_df, hide_index=True, width="stretch", column_config=column_config + ) + else: + st.info( + "No columns selected for statistics calculation.", icon=":material/info:" + ) + + +def _render_enumerator_statistics( + data: pl.DataFrame, + enumerator: str, + team: str | None, + settings_file: str, +) -> None: + """Display enumerator statistics table with team support. + + Shows configurable summary statistics for selected numeric columns + grouped by enumerator (and optionally team). + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + enumerator : str + Enumerator column name. + team : str | None + Team column name (optional). + settings_file : str + Path to settings file for saving/loading configurations. + """ + # Validate inputs + if not enumerator: + st.info( + "Enumerator statistics requires an enumerator column to be selected. " + "Go to the :material/settings: settings section above to select it." + ) + return + + _render_enumerator_statistics_table( + data=data, enumerator=enumerator, team=team, settings_file=settings_file + ) + + +def _load_statistics_overtime_settings( + settings_file: str, +) -> StatisticsOvertimeSettings: + """Load and validate statistics overtime settings from file. + + Parameters + ---------- + settings_file : str + Path to settings file. + + Returns + ------- + StatisticsOvertimeSettings + Validated statistics overtime settings. + """ + saved_settings = load_check_settings(settings_file, TAB_NAME) or {} + try: + return StatisticsOvertimeSettings(**saved_settings) + except ValueError: + # Return default settings if validation fails + return StatisticsOvertimeSettings() + + +def _render_period_selector_overtime( + settings_file: str, + default_period: str = "Week", +) -> str: + """Render time period selection widget. + + Parameters + ---------- + default_period : str + Default selected period. + settings_file : str + Path to settings file. + + Returns + ------- + str + Selected time period. + """ + options_map = { + "Day": ":material/event: Daily", + "Week": ":material/date_range: Weekly", + "Month": ":material/calendar_month: Monthly", + } + period = st.pills( + label="Select Time Period:", + options=options_map.keys(), + format_func=lambda x: options_map[x], + default=default_period, + key="project_enumerator_statistics_overtime_period_pills", + help="Select time period for aggregating statistics", + selection_mode="single", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_period_overtime"}, + ) + save_check_settings(settings_file, TAB_NAME, {"period_overtime": period}) + return period or "Day" + + +def _render_weekday_selector_overtime( + default_weekday: str, + settings_file: str, +) -> str: + """Render weekday selection widget (for weekly period). + + Parameters + ---------- + default_weekday : str + Default selected weekday. + settings_file : str + Path to settings file. + + Returns + ------- + str + Selected weekday offset code (e.g., "SUN", "MON"). + """ + default_weekday_index = WEEKDAY_NAMES.index(default_weekday) + + weekday_sel = st.selectbox( + label="Select the first day of the week", + options=WEEKDAY_NAMES, + index=default_weekday_index, + help="Select the first day of the week", + key="project_week_start_day_enumerator_overtime", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_weekstartday_overtime"}, + ) + save_check_settings(settings_file, TAB_NAME, {"weekstartday": weekday_sel}) + + return WEEKDAY_OFFSET_MAP[weekday_sel] + + +def _render_statistic_selector( + default_stat: str, + settings_file: str, +) -> str: + """Render statistic selection widget. + + Parameters + ---------- + default_stat : str + Default selected statistic. + settings_file : str + Path to settings file. + + Returns + ------- + str + Selected statistic. + """ + default_stat_index = ALLOWED_STATISTICS_OVERTIME.index(default_stat) + + stat = st.selectbox( + label="Select statistic:", + options=ALLOWED_STATISTICS_OVERTIME, + index=default_stat_index, + help="Select statistic to calculate over time", + key="enumerator_statistics_overtime_stat", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_stat_overtime"}, + ) + save_check_settings(settings_file, TAB_NAME, {"stat": stat}) + return stat + + +def _render_column_selector_single( + numeric_cols: list[str], + default_col: str | None, + settings_file: str, +) -> str | None: + """Render single column selection widget. + + Parameters + ---------- + numeric_cols : list[str] + Available numeric columns. + default_col : str | None + Default selected column. + settings_file : str + Path to settings file. + + Returns + ------- + str | None + Selected column. + """ + default_col_index = ( + numeric_cols.index(default_col) + if default_col and default_col in numeric_cols + else None + ) + + statscol = st.selectbox( + label="Select column:", + options=numeric_cols, + index=default_col_index, + help="Select column to include in statistics", + key="enumerator_statistics_overtime_column", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_statscol_overtime"}, + ) + save_check_settings(settings_file, TAB_NAME, {"statscol": statscol}) + return statscol + + +@st.fragment +def _render_enumerator_statistics_overtime_table( + data: pl.DataFrame, + date: str, + enumerator: str, + team: str | None, + settings_file: str, +) -> None: + """Display enumerator statistics over time table with team support. + + Shows how a specific statistic changes over time periods for each + enumerator (and optionally team) with configurable time periods and statistics. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + team : str | None + Team column name (optional). + settings_file : str + Path to settings file for saving/loading configurations. + """ + # Validate inputs + if not (enumerator and date): + return + + # Load and validate settings using Pydantic + settings = _load_statistics_overtime_settings(settings_file) + + # Build exclusion list for numeric columns + exclude_cols = [enumerator, "consent_granted_agg_col", "completed_survey_agg_col"] + if team: + exclude_cols.append(team) + + numeric_cols = _get_numeric_columns(data, exclude_cols=exclude_cols) + + # Render UI in three columns + col1, col2, col3 = st.columns([0.3, 0.2, 0.5]) + + with col1: + statscol = _render_column_selector_single( + numeric_cols, settings.statscol, settings_file + ) + + with col2: + stat = _render_statistic_selector(settings.stat, settings_file) + + with col3: + period = _render_period_selector_overtime( + settings_file, settings.period_overtime + ) + # Conditionally render weekday selector for weekly period + weekstartday = "SAT" # Default + if period == "Week": + weekstartday = _render_weekday_selector_overtime( + settings.weekstartday, settings_file + ) + + # Compute and display statistics + if statscol: + group_by_cols = [enumerator, team] if team else [enumerator] + stats_overtime_df = compute_enumerator_statistics_overtime( + data=data, + date=date, + group_by_cols=group_by_cols, + statscol=statscol, + stat=stat, + period=period, + weekstartday=weekstartday, + ) + + # Build column configuration dynamically with pinning + if team: + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + team: st.column_config.TextColumn("Team", pinned=True), + } + else: + column_config = { + enumerator: st.column_config.TextColumn("Enumerator", pinned=True), + } + + st.dataframe( + stats_overtime_df, + hide_index=True, + width="stretch", + column_config=column_config, + ) + else: + st.info( + "No column selected for statistics calculation.", icon=":material/info:" + ) + + +def _render_enumerator_statistics_overtime( + data: pl.DataFrame, + date: str, + enumerator: str, + team: str | None, + settings_file: str, +) -> None: + """Display enumerator statistics over time table with team support. + + Shows how a specific statistic changes over time periods for each + enumerator (and optionally team) with configurable time periods and statistics. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + date : str + Date column name. + enumerator : str + Enumerator column name. + team : str | None + Team column name (optional). + settings_file : str + Path to settings file for saving/loading configurations. + """ + # Validate inputs + if not (enumerator and date): + st.info( + "Enumerator statistics over time requires a date and enumerator column to be selected. " + "Go to the :material/settings: settings section above to select them." + ) + return + + _render_enumerator_statistics_overtime_table( + data=data, + date=date, + enumerator=enumerator, + team=team, + settings_file=settings_file, + ) + + +# ============================================================================= +# Main Enumerator Report Function +# ============================================================================= + + +def enumerator_report( + project_id: str, + data: pl.DataFrame, + setting_file: str, + config: dict, + survey_columns: ColumnByType, +) -> None: + """Generate a comprehensive enumerator performance report. + + Creates a complete enumerator analysis report including: + - Overview metrics and statistics + - Comprehensive enumerator summary table + - Productivity tracking over time + - Statistical analysis across enumerators + - Time-series analysis of performance + + Parameters + ---------- + project_id : str + Unique project identifier for configuration lookup. + data : pl.DataFrame + Dataset containing survey data to analyze. + settings_file : str + Path to settings file for persisting configurations. + missing_settings_file : str + Path to missing codes configuration file. + page_num : int + Page number for configuration defaults (1-indexed). + """ + categorical_columns = survey_columns.categorical_columns + datetime_columns = survey_columns.datetime_columns + + st.title("Enumerator Report") + + demo_callout( + """ + This tab tracks enumerator performance across your survey dataset. + + It has four sections: + - **Enumerator Overview**: 8 summary metrics at the top. + - **Enumerator Summary**: Tabular breakdown by enumerator with pill-based view switching. + - **Enumerator Productivity**: Submission counts per enumerator over time. + - **Column Statistics by Enumerator**: Per-column statistics broken down by enumerator. + - **Enumerator Statistics Over Time**: Time-series chart of a selected statistic. + + **Start here**: Open the :material/settings: settings panel above to confirm your column + selections and apply consent and outcome settings. + """ + ) + + if data.is_empty(): + st.info( + "No data available for the enumerator report. " + "Please upload data to proceed." + ) + return + + config_settings = EnumeratorSettings(**config) + + enumerator_settings = enumerator_report_settings( + project_id, + setting_file, + data, + config_settings, + categorical_columns, + datetime_columns, + ) + + # get data for enumerator report + data_enum_report = duckdb_get_table( + project_id, + "enumerator_data_with_consent_outcome", + "intermediate", + ) + + if data_enum_report.is_empty(): + data_enum_report = data + + demo_callout( + """ + ##### Enumerator Overview + Eight metrics appear here in two rows of four: + - Row 1: Total enumerators, Total teams, Active enumerators (past 7 days), + % active enumerators. + - Row 2: Fewest submissions, Highest submissions, Average submissions, + Total submissions. + """ + ) + + _render_enumerator_overview_metrics( + data_enum_report, + enumerator_settings.survey_date, + enumerator_settings.enumerator, + enumerator_settings.team, + ) + + st.write("---") + st.subheader("Enumerator Summary") + + demo_callout( + """ + ##### Enumerator Summary + Use the pills above the table to switch between views. Each pill shows a different + set of columns: + - **Submissions**: First/last submission dates, submission counts (total, today, + this week, this month), unique active days. + - **Missing Data**: Percentage of missing values per enumerator. + - **Duration**: Min, max, mean, and median interview duration. + - **Form Version**: Form versions used by each enumerator. + - **Consent & Outcome**: Consent rate and completed survey rate per enumerator. + + You can select multiple pills to see combined columns side by side. + """ + ) + + _render_enumerator_summary_table( + project_id, + data_enum_report, + enumerator_settings.survey_date, + enumerator_settings.enumerator, + enumerator_settings.team, + enumerator_settings.formversion, + enumerator_settings.duration, + ) + + st.write("---") + st.subheader("Enumerator Productivity") + + demo_callout( + """ + ##### Enumerator Productivity + This section shows submission counts per enumerator over time as a table. + Use the **Daily / Weekly / Monthly** pills to change the time period granularity. + """ + ) + + _render_enumerator_productivity( + data_enum_report, + enumerator_settings.survey_date, + enumerator_settings.enumerator, + enumerator_settings.team, + setting_file, + ) + + st.write("---") + st.subheader("Column Statistics by Enumerator") + + demo_callout( + """ + ##### Column Statistics by Enumerator + Use the **column multiselect** to choose one or more numeric columns to analyse, + then use the **statistics multiselect** to choose which statistics to display + (count, mean, median, min, max, std, 25th percentile, 75th percentile). + + ##### Instructions for Demo: + Select **household_size** as the column and choose **count**, **mean**, **min**, + and **max** to check whether enumerators are recording consistent household sizes. + """ + ) + + _render_enumerator_statistics( + data_enum_report, + enumerator_settings.enumerator, + enumerator_settings.team, + setting_file, + ) + + st.write("---") + st.subheader("Enumerator Statistics Over Time") + + demo_callout( + """ + ##### Enumerator Statistics Over Time + This section renders a line chart showing how a selected statistic for a chosen + column changes over time, broken down per enumerator. Use the **column selectbox** + to pick the variable, the **statistic selectbox** to choose the metric + (e.g., mean, count, missing), and the **Daily / Weekly / Monthly** pills to set + the time granularity. + + ##### Instructions for Demo: + Select **household_size** as the column and **mean** as the statistic, then switch + to **Weekly** to see whether average household size varies by enumerator over time. + """ + ) + + _render_enumerator_statistics_overtime( + data_enum_report, + enumerator_settings.survey_date, + enumerator_settings.enumerator, + enumerator_settings.team, + setting_file, + ) + + demo_callout( + "**Next**: :material/arrow_upward: Scroll up and select the **Backcheck Analysis** tab." + ) diff --git a/src/datasure/checks/enumerator/settings_ui.py b/src/datasure/checks/enumerator/settings_ui.py new file mode 100644 index 00000000..62262a08 --- /dev/null +++ b/src/datasure/checks/enumerator/settings_ui.py @@ -0,0 +1,499 @@ +"""Settings UI for the enumerator performance report.""" + +import polars as pl +import streamlit as st + +from datasure.checks.enumerator.models import ( + TAB_NAME, + ConsentOutcomeSettings, + EnumeratorSettings, +) +from datasure.utils.duckdb_utils import duckdb_save_table +from datasure.utils.onboarding_utils import demo_output_onboarding +from datasure.utils.settings_utils import ( + load_check_settings, + save_check_settings, + trigger_save, +) + +# ============================================================================= +# Settings Management Functions +# ============================================================================= + + +def _render_column_select( + label: str, + field_key: str, + options: list[str], + help_text: str, + default_settings: EnumeratorSettings, + settings_file: str, +) -> str | None: + """Render a settings selectbox bound to an EnumeratorSettings field. + + Every single-column picker in the settings UI (survey key, survey ID, + survey date, enumerator, team, duration, form version) shares this same + default-index-lookup + selectbox + save_check_settings shape. + + Parameters + ---------- + label : str + Widget label shown above the selectbox. + field_key : str + Name of the EnumeratorSettings field this selectbox configures; also + used to derive the widget key and the saved settings key. + options : list[str] + Columns to offer as selectbox options. + help_text : str + Tooltip text for the selectbox. + default_settings : EnumeratorSettings + Previously saved settings, used to preselect a default option. + settings_file : str + Path to settings file for saving the selected value. + + Returns + ------- + str | None + The selected column name. + """ + default_value = getattr(default_settings, field_key) + default_index = ( + options.index(default_value) + if default_value and default_value in options + else None + ) + selected = st.selectbox( + label, + options=options, + key=f"{field_key}_enumerator", + help=help_text, + index=default_index, + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + f"_{field_key}"}, + ) + save_check_settings(settings_file, TAB_NAME, {field_key: selected}) + return selected + + +@st.cache_data(ttl=60) +def load_default_enumerator_settings( + settings_file: str, config: EnumeratorSettings +) -> EnumeratorSettings: + """Load and merge saved settings with default configuration. + + Loads previously saved duplicates report settings from the settings file + and merges them with the provided default configuration. Saved settings + take precedence over defaults. + + Cached for 60 seconds to reduce file I/O operations. + + Parameters + ---------- + settings_file : str + Path to the settings file containing saved configurations. + config : DuplicatesSettings + Default configuration to use as fallback for missing settings. + + Returns + ------- + DuplicatesSettings + Merged settings combining saved and default configurations. + """ + saved_settings = load_check_settings(settings_file, TAB_NAME) + + default_settings: dict = dict(config) + default_settings.update(saved_settings) + + return EnumeratorSettings(**default_settings) + + +@demo_output_onboarding(TAB_NAME) +def enumerator_report_settings( + project_id: str, + settings_file: str, + data: pl.DataFrame, + config: EnumeratorSettings, + categorical_columns: list[str], + datetime_columns: list[str], +) -> EnumeratorSettings: + """Create and render the settings UI for duplicates report configuration. + + This function creates a comprehensive Streamlit UI for configuring + duplicates report settings. It includes: + - Survey identifiers (key and ID columns) + - Survey date column selection + - Enumerator ID column + - Filtering conditions for targeted duplicate detection + + Settings are automatically saved to the settings file when changed + and loaded from previous sessions if available. + + Parameters + ---------- + project_id : str + Unique project identifier for database operations. + settings_file : str + Path to settings file for saving/loading configurations. + data : pl.DataFrame + Dataset to analyze for duplicates. + config : DuplicatesSettings + Default configuration used as fallback values. + categorical_columns : list[str] + Available categorical columns for selection (survey key, ID, enumerator). + datetime_columns : list[str] + Available datetime columns for date selection. + + Returns + ------- + DuplicatesSettings + User-configured settings from the UI. + """ + with st.expander("settings", icon=":material/settings:"): + st.markdown("## Configure settings for enumerator report") + st.write("---") + + default_settings = load_default_enumerator_settings(settings_file, config) + + # Survey Identifiers + with st.container(border=True): + st.subheader("Survey Identifiers") + si1, si2, _ = st.columns(3) + + with si1: + survey_key = _render_column_select( + "Survey Key", + "survey_key", + categorical_columns, + "Select the column that contains the survey key", + default_settings, + settings_file, + ) + + with si2: + survey_id = _render_column_select( + "Survey ID", + "survey_id", + categorical_columns, + "Select the column that contains the survey ID", + default_settings, + settings_file, + ) + + with st.container(border=True): + st.subheader("Survey Date") + + sd1, _, _ = st.columns(3) + + with sd1: + survey_date = _render_column_select( + "Survey Date", + "survey_date", + datetime_columns, + "Select the column that contains the survey date", + default_settings, + settings_file, + ) + + with st.container(border=True): + st.subheader("Enumerator") + ec1, ec2, _ = st.columns(3) + with ec1: + enumerator = _render_column_select( + "Enumerator ID", + "enumerator", + categorical_columns, + "Select the column that contains the enumerator ID", + default_settings, + settings_file, + ) + + with ec2: + team = _render_column_select( + "Team", + "team", + categorical_columns, + "Select the column that contains the team identifier", + default_settings, + settings_file, + ) + + with st.container(border=True): + st.subheader("Survey Duration") + dc1, dc2, _ = st.columns(3) + with dc1: + duration = _render_column_select( + "Duration Column", + "duration", + categorical_columns, + "Select the column that contains the survey duration in seconds", + default_settings, + settings_file, + ) + + with dc2: + default_duration_unit = default_settings.duration_unit + default_duration_unit_index = ( + ["seconds", "minutes", "hours"].index(default_duration_unit) + if default_duration_unit in ["seconds", "minutes", "hours"] + else 0 + ) + duration_unit = st.selectbox( + "Duration Unit", + options=["seconds", "minutes", "hours"], + key="duration_unit_enumerator", + help="Select the unit for survey duration", + index=default_duration_unit_index, + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + "_duration_unit"}, + ) + save_check_settings( + settings_file, TAB_NAME, {"duration_unit": duration_unit} + ) + + with st.container(border=True): + st.subheader("Form Version") + fv1, _ = st.columns([1, 2]) + with fv1: + formversion = _render_column_select( + "Form Version Column", + "formversion", + categorical_columns, + "Select the column that contains the form version", + default_settings, + settings_file, + ) + + with st.container(border=True): + st.subheader("Consent and Outcome Settings") + st.info( + "Configure consent and outcome columns along with their valid values." + ) + + _render_consent_outcome_settings( + project_id, data, categorical_columns, settings_file + ) + if st.session_state.get("st_apply_consent_outcome_enumerator"): + st.success("Consent and outcome settings applied successfully.") + st.session_state["st_apply_consent_outcome_enumerator"] = False + + return EnumeratorSettings( + survey_key=survey_key, + survey_id=survey_id, + survey_date=survey_date, + enumerator=enumerator, + team=team, + formversion=formversion, + duration=duration, + duration_unit=duration_unit, + ) + + +def _render_category_settings( + column_label: str, + field_key: str, + column_help: str, + values_label: str, + values_help: str, + categorical_columns: list, + default_settings: dict, + data: pl.DataFrame, + settings_file: str, +) -> tuple[str | None, list[str]]: + """Render a column selector + valid-values multiselect pair. + + The consent and outcome settings sections share this same + column-then-values-multiselect layout, differing only in labels, + help text, and which EnumeratorSettings field they populate. + + Parameters + ---------- + column_label : str + Widget label for the column selectbox. + field_key : str + Base settings key (e.g. "consent" or "outcome"); the values + multiselect is saved under f"{field_key}_vals". + column_help : str + Tooltip text for the column selectbox. + values_label : str + Widget label for the values multiselect. + values_help : str + Tooltip text for the values multiselect. + categorical_columns : list + Columns to offer as selectbox options. + default_settings : dict + Previously saved settings, used to preselect defaults. + data : pl.DataFrame + DataFrame containing survey data, used to derive value options. + settings_file : str + Path to settings file for saving/loading configurations. + + Returns + ------- + tuple[str | None, list[str]] + The selected column name and selected valid values. + """ + col1, col2 = st.columns([0.3, 0.7]) + with col1: + default_col = default_settings.get(field_key) + default_index = ( + categorical_columns.index(default_col) + if default_col and default_col in categorical_columns + else 0 + ) + selected_col = st.selectbox( + column_label, + options=categorical_columns, + help=column_help, + key=f"{field_key}_enumerator", + index=default_index, + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + f"_{field_key}"}, + ) + save_check_settings(settings_file, TAB_NAME, {field_key: selected_col}) + + with col2: + vals_key = f"{field_key}_vals" + default_vals = default_settings.get(vals_key, []) + val_options = data[selected_col].unique().to_list() + selected_vals = st.multiselect( + values_label, + options=val_options, + default=default_vals, + help=values_help, + key=f"{vals_key}_enumerator", + on_change=trigger_save, + kwargs={"state_name": TAB_NAME + f"_{vals_key}"}, + ) + save_check_settings(settings_file, TAB_NAME, {vals_key: selected_vals}) + + return selected_col, selected_vals + + +@st.fragment +def _render_consent_outcome_settings( + project_id: str, data: pl.DataFrame, categorical_columns: list, settings_file: str +) -> ConsentOutcomeSettings: + """Render consent and outcome settings UI. + + Parameters + ---------- + data : pl.DataFrame + DataFrame containing survey data. + settings_file : str + Path to settings file for saving/loading configurations. + + Returns + ------- + ConsentOutcomeSettings + User-configured consent and outcome settings. + """ + default_settings = load_check_settings(settings_file, TAB_NAME) + + with st.container(border=True): + st.subheader("Consent Settings") + consent_col, consent_vals = _render_category_settings( + "Consent Column", + "consent", + "Select the column that contains consent status", + "Valid Consent Values", + "Select values that indicate valid consent", + categorical_columns, + default_settings, + data, + settings_file, + ) + + with st.container(border=True): + st.subheader("Outcome Settings") + outcome_col, outcome_vals = _render_category_settings( + "Outcome Column", + "outcome", + "Select the column that contains survey outcome status", + "Completed Survey Values", + "Select values that indicate completed surveys", + categorical_columns, + default_settings, + data, + settings_file, + ) + + config_dict = { + "consent": consent_col, + "consent_vals": consent_vals, + "outcome": outcome_col, + "outcome_vals": outcome_vals, + } + config = ConsentOutcomeSettings(**config_dict) + + if st.button( + "Apply Consent and Outcome Settings", + key="apply_consent_outcome_enumerator", + type="primary", + width="stretch", + ): + _create_enum_data_on_settings(project_id, data, config) + _trigger_success_message("st_apply_consent_outcome_enumerator") + st.rerun() + + +def _trigger_success_message(button_key: str) -> None: + """Trigger a success message after button click. + + Parameters + ---------- + button_key : str + Unique key of the button to associate the success message with. + """ + st.session_state[f"{button_key}"] = True + + +def _create_enum_data_on_settings( + project_id: str, + data: pl.DataFrame, + config: ConsentOutcomeSettings, +) -> None: + """Create enumerator data based on consent and outcome settings. + + Parameters + ---------- + project_id : str + Unique project identifier for database operations. + data : pl.DataFrame + DataFrame containing survey data. + conditions : ConsentOutcomeSettings + Consent and outcome configuration settings. + """ + # If consent and consent values are provided, create a dummy column indicating + # valid consent else set to 1 + + if config.consent and config.consent_vals: + enum_data = data.with_columns( + pl.col(config.consent) + .is_in(config.consent_vals) + .cast(pl.Int32) + .alias("consent_granted_agg_col") + ) + else: + enum_data = data.with_columns( + pl.lit(1).cast(pl.Int32).alias("consent_granted_agg_col") + ) + + if config.outcome and config.outcome_vals: + enum_data = enum_data.with_columns( + pl.col(config.outcome) + .is_in(config.outcome_vals) + .cast(pl.Int32) + .alias("completed_survey_agg_col") + ) + else: + enum_data = enum_data.with_columns( + pl.lit(1).cast(pl.Int32).alias("completed_survey_agg_col") + ) + + # save to database + duckdb_save_table( + project_id, + enum_data, + "enumerator_data_with_consent_outcome", + "intermediate", + ) diff --git a/src/datasure/views/output_view_template.py b/src/datasure/views/output_view_template.py index c1118c74..7f474f91 100644 --- a/src/datasure/views/output_view_template.py +++ b/src/datasure/views/output_view_template.py @@ -17,7 +17,7 @@ from datasure.checks.backchecks import backchecks_report from datasure.checks.descriptive import descriptive_report from datasure.checks.duplicates import duplicates_report -from datasure.checks.enumerator import enumerator_report +from datasure.checks.enumerator.report_ui import enumerator_report from datasure.checks.gpschecks import gpschecks_report from datasure.checks.missing import missing_report from datasure.checks.outliers import outliers_report diff --git a/tests/checks/enumerator/conftest.py b/tests/checks/enumerator/conftest.py new file mode 100644 index 00000000..abf3eecf --- /dev/null +++ b/tests/checks/enumerator/conftest.py @@ -0,0 +1,133 @@ +"""Shared fixtures for the split enumerator module test suite.""" + +import json +from datetime import date, timedelta +from unittest.mock import MagicMock + +import polars as pl +import pytest + +from datasure.checks.enumerator.models import EnumeratorSettings + +# ============================================ +# MOCK STREAMLIT HELPERS +# ============================================ + + +def make_mock_st(): + """Create a mock Streamlit module for testing UI render functions.""" + + def make_col(): + col = MagicMock() + col.number_input.return_value = 0.0 + col.selectbox.return_value = None + col.text_input.return_value = "" + col.multiselect.return_value = [] + return col + + def _col_factory(n_or_spec, **kwargs): + if isinstance(n_or_spec, int): + n = n_or_spec + elif isinstance(n_or_spec, list | tuple): + n = len(n_or_spec) + else: + n = 2 + return tuple(make_col() for _ in range(n)) + + mock_st = MagicMock() + mock_st.fragment = lambda func: func + mock_st.dialog = lambda *args, **kwargs: lambda func: func + mock_st.columns.side_effect = _col_factory + mock_st.selectbox.return_value = None + mock_st.multiselect.return_value = [] + mock_st.pills.return_value = None + mock_st.button.return_value = False + mock_st.toggle.return_value = False + mock_st.number_input.return_value = 0 + mock_st.text_input.return_value = "" + mock_st.session_state = {} + return mock_st + + +# ============================================ +# FIXTURES FOR ENUMERATOR-SPECIFIC DATA +# ============================================ + + +@pytest.fixture(autouse=True) +def mock_database_functions(monkeypatch): + """Override the autouse fixture from conftest. + + Disables database mocking for these tests. + """ + pass + + +@pytest.fixture +def sample_enumerator_data(): + """Create sample enumerator data as Polars DataFrame.""" + today = date.today() + return pl.DataFrame( + { + "survey_id": ["S001", "S002", "S003", "S004", "S005", "S006"], + "submission_date": [ + today - timedelta(days=1), + today - timedelta(days=2), + today - timedelta(days=3), + today - timedelta(days=1), + today - timedelta(days=8), + today, + ], + "enumerator": ["E1", "E1", "E2", "E2", "E3", "E1"], + "team": ["T1", "T1", "T1", "T2", "T2", "T1"], + "duration": [3600, 4200, 3800, 4000, 3900, 3700], + "formversion": ["v1", "v1", "v2", "v2", "v1", "v2"], + "age": [25, 30, 35, 28, 32, 27], + "income": [50000, 60000, 55000, 52000, 58000, 51000], + "consent_granted_agg_col": [1, 1, 1, 0, 1, 1], + "completed_survey_agg_col": [1, 1, 1, 1, 0, 1], + } + ) + + +@pytest.fixture +def sample_enumerator_settings(): + """Create sample EnumeratorSettings for testing.""" + return EnumeratorSettings( + survey_key="survey_id", + survey_id="survey_id", + survey_date="submission_date", + enumerator="enumerator", + formversion="formversion", + duration="duration", + duration_unit="seconds", + team="team", + ) + + +@pytest.fixture +def sample_missing_codes_config(): + """Create sample missing codes configuration.""" + return pl.DataFrame( + { + "label": ["Refused", "Don't know"], + "codes": ["-99", "-88"], + } + ) + + +@pytest.fixture +def enumerator_settings_file(tmp_path): + """Create a temporary enumerator settings file.""" + settings = { + "enumerators": { + "survey_key": "survey_id", + "survey_id": "survey_id", + "survey_date": "submission_date", + "enumerator": "enumerator", + "team": "team", + } + } + file_path = tmp_path / "enumerator_settings.json" + file_path.write_text(json.dumps(settings)) + return str(file_path) diff --git a/tests/checks/enumerator/test_compute.py b/tests/checks/enumerator/test_compute.py new file mode 100644 index 00000000..947e6bd4 --- /dev/null +++ b/tests/checks/enumerator/test_compute.py @@ -0,0 +1,818 @@ +"""Tests for datasure.checks.enumerator.compute.""" + +from datetime import date, timedelta +from unittest.mock import patch + +import polars as pl +import pytest +from polars.exceptions import ColumnNotFoundError + +from datasure.checks.enumerator.compute import ( + compute_enumerator_missing_table, + compute_enumerator_overview, + compute_enumerator_productivity, + compute_enumerator_statistics, + compute_enumerator_statistics_overtime, + compute_enumerator_summary, +) + +# ============================================ +# COMPUTE_ENUMERATOR_OVERVIEW TESTS +# ============================================ + + +def test_compute_enumerator_overview_basic(sample_enumerator_data): + """Test basic enumerator overview computation.""" + result = compute_enumerator_overview( + sample_enumerator_data, "submission_date", "enumerator", "team" + ) + + assert result.all_submissions == 6 + assert result.num_enumerators == 3 + assert result.num_teams == 2 + assert result.num_active_enumerators >= 0 + assert result.min_submissions > 0 + assert result.max_submissions > 0 + assert result.avg_submissions > 0 + + +def test_compute_enumerator_overview_without_team(sample_enumerator_data): + """Test enumerator overview without team column.""" + result = compute_enumerator_overview( + sample_enumerator_data, "submission_date", "enumerator", None + ) + + assert result.all_submissions == 6 + assert result.num_enumerators == 3 + assert result.num_teams == "n/a" + + +def test_compute_enumerator_overview_empty_data(): + """Test enumerator overview with empty data.""" + empty_data = pl.DataFrame( + schema={"submission_date": pl.Date, "enumerator": pl.Utf8} + ) + + with pytest.raises(ValueError, match="Input data is empty"): + compute_enumerator_overview(empty_data, "submission_date", "enumerator", None) + + +def test_compute_enumerator_overview_active_enumerators(): + """Test active enumerators calculation.""" + today = date.today() + data = pl.DataFrame( + { + "submission_date": [ + today - timedelta(days=1), + today - timedelta(days=10), + today, + ], + "enumerator": ["E1", "E2", "E1"], + "team": ["T1", "T1", "T1"], + } + ) + + result = compute_enumerator_overview(data, "submission_date", "enumerator", "team") + + # Only E1 should be active (has submissions in past 7 days) + assert result.num_active_enumerators == 1 + assert result.num_enumerators == 2 + + +# ============================================ +# COMPUTE_ENUMERATOR_MISSING_TABLE TESTS +# ============================================ + + +def test_compute_enumerator_missing_table_empty_config(sample_enumerator_data): + """Test missing table with empty missing codes config.""" + empty_config = pl.DataFrame() + + result = compute_enumerator_missing_table( + sample_enumerator_data, empty_config, ["enumerator"] + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "% Null values" in result.columns + + +def test_compute_enumerator_missing_table_with_config( + sample_enumerator_data, sample_missing_codes_config +): + """Test missing table with missing codes config.""" + result = compute_enumerator_missing_table( + sample_enumerator_data, sample_missing_codes_config, ["enumerator"] + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_missing_table_with_team(sample_enumerator_data): + """Test missing table with team grouping.""" + empty_config = pl.DataFrame() + + result = compute_enumerator_missing_table( + sample_enumerator_data, empty_config, ["enumerator", "team"] + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "team" in result.columns + + +# ============================================ +# COMPUTE_ENUMERATOR_SUMMARY TESTS +# ============================================ + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_basic(mock_load_missing, sample_enumerator_data): + """Test basic enumerator summary computation.""" + mock_load_missing.return_value = pl.DataFrame() + + result = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "formversion", + "duration", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "team" in result.columns + assert "# submissions" in result.columns + assert "first submission" in result.columns + assert "last submission" in result.columns + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_without_team( + mock_load_missing, sample_enumerator_data +): + """Test enumerator summary without team.""" + mock_load_missing.return_value = pl.DataFrame() + + result = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "formversion", + "duration", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "team" not in result.columns + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_without_duration( + mock_load_missing, sample_enumerator_data +): + """Test enumerator summary without duration.""" + mock_load_missing.return_value = pl.DataFrame() + + result = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "formversion", + None, + ) + + assert not result.is_empty() + assert "min duration" not in result.columns + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_without_formversion( + mock_load_missing, sample_enumerator_data +): + """Test enumerator summary without formversion.""" + mock_load_missing.return_value = pl.DataFrame() + + result = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + None, + "duration", + ) + + assert not result.is_empty() + assert "# form versions" not in result.columns + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_with_consent( + mock_load_missing, sample_enumerator_data +): + """Test enumerator summary with consent column.""" + mock_load_missing.return_value = pl.DataFrame() + + result = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "formversion", + "duration", + ) + + assert not result.is_empty() + assert "% consent" in result.columns + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_with_outcome( + mock_load_missing, sample_enumerator_data +): + """Test enumerator summary with outcome column.""" + mock_load_missing.return_value = pl.DataFrame() + + result = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "formversion", + "duration", + ) + + assert not result.is_empty() + assert "% completed survey" in result.columns + + +# ============================================ +# COMPUTE_ENUMERATOR_PRODUCTIVITY TESTS +# ============================================ + + +def test_compute_enumerator_productivity_daily(sample_enumerator_data): + """Test productivity computation with daily period.""" + result = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Daily", "SUN" + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_productivity_weekly(sample_enumerator_data): + """Test productivity computation with weekly period.""" + result = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Weekly", "MON" + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_productivity_monthly(sample_enumerator_data): + """Test productivity computation with monthly period.""" + result = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Monthly", "SUN" + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_productivity_legacy_period_names(sample_enumerator_data): + """Test productivity with legacy period names.""" + # Test "Day" -> "Daily" + result = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Day", "SUN" + ) + assert not result.is_empty() + + # Test "Week" -> "Weekly" + result = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Week", "MON" + ) + assert not result.is_empty() + + # Test "Month" -> "Monthly" + result = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Month", "SUN" + ) + assert not result.is_empty() + + +def test_compute_enumerator_productivity_with_team(sample_enumerator_data): + """Test productivity with team grouping.""" + result = compute_enumerator_productivity( + sample_enumerator_data, + "submission_date", + ["enumerator", "team"], + "Daily", + "SUN", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "team" in result.columns + + +def test_compute_enumerator_productivity_different_weekstarts(sample_enumerator_data): + """Test productivity with different week start days.""" + for weekstart in ["SUN", "MON", "TUE", "WED", "THU", "FRI", "SAT"]: + result = compute_enumerator_productivity( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "Weekly", + weekstart, + ) + assert not result.is_empty() + + +# ============================================ +# COMPUTE_ENUMERATOR_STATISTICS TESTS +# ============================================ + + +def test_compute_enumerator_statistics_basic(sample_enumerator_data): + """Test basic statistics computation.""" + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = compute_enumerator_statistics( + sample_enumerator_data, + ["enumerator"], + ["age", "income"], + ["count", "mean"], + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "age_count" in result.columns + assert "age_mean" in result.columns + assert "income_count" in result.columns + assert "income_mean" in result.columns + + +def test_compute_enumerator_statistics_all_stats(sample_enumerator_data): + """Test statistics with all stat types.""" + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = compute_enumerator_statistics( + sample_enumerator_data, + ["enumerator"], + ["age"], + [ + "count", + "min", + "mean", + "median", + "max", + "std", + "25th percentile", + "75th percentile", + ], + ) + + assert not result.is_empty() + assert "age_count" in result.columns + assert "age_min" in result.columns + assert "age_mean" in result.columns + assert "age_median" in result.columns + assert "age_max" in result.columns + assert "age_std" in result.columns + assert "age_25th percentile" in result.columns + assert "age_75th percentile" in result.columns + + +def test_compute_enumerator_statistics_with_team(sample_enumerator_data): + """Test statistics with team grouping.""" + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = compute_enumerator_statistics( + sample_enumerator_data, + ["enumerator", "team"], + ["age"], + ["mean"], + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "team" in result.columns + + +def test_compute_enumerator_statistics_multiple_columns(sample_enumerator_data): + """Test statistics with multiple columns.""" + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = compute_enumerator_statistics( + sample_enumerator_data, + ["enumerator"], + ["age", "income", "duration"], + ["mean", "median"], + ) + + assert not result.is_empty() + assert len([col for col in result.columns if "_mean" in col]) == 3 + assert len([col for col in result.columns if "_median" in col]) == 3 + + +# ============================================ +# COMPUTE_ENUMERATOR_STATISTICS_OVERTIME TESTS +# ============================================ + + +def test_compute_enumerator_statistics_overtime_daily(sample_enumerator_data): + """Test statistics overtime with daily period.""" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Daily", + "SUN", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_statistics_overtime_weekly(sample_enumerator_data): + """Test statistics overtime with weekly period.""" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Weekly", + "MON", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_statistics_overtime_monthly(sample_enumerator_data): + """Test statistics overtime with monthly period.""" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Monthly", + "SUN", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + + +def test_compute_enumerator_statistics_overtime_missing_stat(sample_enumerator_data): + """Test statistics overtime with missing statistic.""" + # Add some null values + data_with_nulls = sample_enumerator_data.clone() + data_with_nulls = data_with_nulls.with_columns( + pl.when(pl.col("enumerator") == "E1") + .then(None) + .otherwise(pl.col("age")) + .alias("age") + ) + + result = compute_enumerator_statistics_overtime( + data_with_nulls, + "submission_date", + ["enumerator"], + "age", + "missing", + "Daily", + "SUN", + ) + + assert not result.is_empty() + + +def test_compute_enumerator_statistics_overtime_percentile_stats( + sample_enumerator_data, +): + """Test statistics overtime with percentile statistics.""" + result_25th = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "25th percentile", + "Daily", + "SUN", + ) + assert not result_25th.is_empty() + + result_75th = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "75th percentile", + "Daily", + "SUN", + ) + assert not result_75th.is_empty() + + +def test_compute_enumerator_statistics_overtime_legacy_periods(sample_enumerator_data): + """Test statistics overtime with legacy period names.""" + # Test "Day" -> "Daily" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Day", + "SUN", + ) + assert not result.is_empty() + + # Test "Week" -> "Weekly" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Week", + "MON", + ) + assert not result.is_empty() + + # Test "Month" -> "Monthly" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Month", + "SUN", + ) + assert not result.is_empty() + + +def test_compute_enumerator_statistics_overtime_with_team(sample_enumerator_data): + """Test statistics overtime with team grouping.""" + result = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator", "team"], + "age", + "mean", + "Daily", + "SUN", + ) + + assert not result.is_empty() + assert "enumerator" in result.columns + assert "team" in result.columns + + +# ============================================ +# EDGE CASES AND ERROR HANDLING TESTS +# ============================================ + + +def test_edge_case_single_enumerator(): + """Test handling of single enumerator.""" + data = pl.DataFrame( + { + "submission_date": [date.today(), date.today() - timedelta(days=1)], + "enumerator": ["E1", "E1"], + "age": [25, 30], + } + ) + + result = compute_enumerator_overview(data, "submission_date", "enumerator", None) + assert result.num_enumerators == 1 + assert result.all_submissions == 2 + + +def test_edge_case_all_null_values(): + """Test handling of all null values in column.""" + data = pl.DataFrame( + { + "submission_date": [date.today(), date.today()], + "enumerator": ["E1", "E2"], + "age": [None, None], + } + ) + + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = compute_enumerator_statistics( + data, + ["enumerator"], + ["age"], + ["mean"], + ) + + assert not result.is_empty() + # Mean of nulls should be null + assert result["age_mean"].is_null().any() + + +def test_edge_case_single_submission(): + """Test handling of single submission.""" + data = pl.DataFrame( + { + "submission_date": [date.today()], + "enumerator": ["E1"], + "team": ["T1"], + "age": [25], + } + ) + + result = compute_enumerator_overview(data, "submission_date", "enumerator", "team") + assert result.all_submissions == 1 + assert result.num_enumerators == 1 + + +def test_edge_case_date_ranges(): + """Test handling of various date ranges.""" + today = date.today() + data = pl.DataFrame( + { + "submission_date": [ + today, + today - timedelta(days=365), + today - timedelta(days=1000), + ], + "enumerator": ["E1", "E2", "E3"], + } + ) + + result = compute_enumerator_productivity( + data, "submission_date", ["enumerator"], "Monthly", "SUN" + ) + assert not result.is_empty() + + +# ============================================ +# INTEGRATION TESTS +# ============================================ + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_full_enumerator_workflow(mock_load_missing, sample_enumerator_data): + """Test complete enumerator workflow from overview to statistics.""" + mock_load_missing.return_value = pl.DataFrame() + + # Step 1: Compute overview + overview = compute_enumerator_overview( + sample_enumerator_data, "submission_date", "enumerator", "team" + ) + assert overview.num_enumerators == 3 + + # Step 2: Compute summary + summary = compute_enumerator_summary( + "test_project", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "formversion", + "duration", + ) + assert len(summary) > 0 + + # Step 3: Compute productivity + productivity = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], "Daily", "SUN" + ) + assert not productivity.is_empty() + + # Step 4: Compute statistics + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + statistics = compute_enumerator_statistics( + sample_enumerator_data, + ["enumerator"], + ["age", "income"], + ["mean", "median"], + ) + assert not statistics.is_empty() + + # Step 5: Compute statistics overtime + stats_overtime = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "Daily", + "SUN", + ) + assert not stats_overtime.is_empty() + + +def test_enumerator_workflow_with_missing_data(): + """Test enumerator workflow with missing data.""" + data = pl.DataFrame( + { + "submission_date": [date.today(), date.today() - timedelta(days=1)], + "enumerator": ["E1", "E2"], + "age": [25, None], + "income": [50000, 60000], + } + ) + + # Overview should work with missing data + overview = compute_enumerator_overview(data, "submission_date", "enumerator", None) + assert overview.all_submissions == 2 + + # Statistics should handle missing values + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + statistics = compute_enumerator_statistics( + data, + ["enumerator"], + ["age", "income"], + ["count", "mean"], + ) + assert not statistics.is_empty() + + +def test_enumerator_workflow_different_time_periods(sample_enumerator_data): + """Test enumerator workflow with different time periods.""" + periods = ["Daily", "Weekly", "Monthly"] + + for period in periods: + # Test productivity + productivity = compute_enumerator_productivity( + sample_enumerator_data, "submission_date", ["enumerator"], period, "MON" + ) + assert not productivity.is_empty() + + # Test statistics overtime + stats_overtime = compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + period, + "MON", + ) + assert not stats_overtime.is_empty() + + +# ============================================ +# BRANCH MISS TESTS (COMPUTE FUNCTIONS) +# ============================================ + + +@patch("datasure.checks.enumerator.compute.load_missing_codes_from_db") +def test_compute_enumerator_summary_without_consent_outcome_cols(mock_load_missing): + """compute_enumerator_summary skips consent/outcome when cols are absent.""" + mock_load_missing.return_value = pl.DataFrame() + data = pl.DataFrame( + { + "submission_date": [date.today()], + "enumerator": ["E1"], + } + ) + result = compute_enumerator_summary( + "proj", data, "submission_date", "enumerator", None, None, None + ) + assert not result.is_empty() + assert "% consent" not in result.columns + assert "% completed survey" not in result.columns + + +def test_compute_enumerator_productivity_unknown_period(sample_enumerator_data): + """compute_enumerator_productivity falls through period branches on unknown.""" + with pytest.raises(ColumnNotFoundError): + compute_enumerator_productivity( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "UnknownPeriod", + "SUN", + ) + + +def test_compute_enumerator_statistics_overtime_unknown_period(sample_enumerator_data): + """compute_enumerator_statistics_overtime falls through period branches.""" + with pytest.raises(ColumnNotFoundError): + compute_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + ["enumerator"], + "age", + "mean", + "UnknownPeriod", + "SUN", + ) diff --git a/tests/checks/enumerator/test_models.py b/tests/checks/enumerator/test_models.py new file mode 100644 index 00000000..47b5ebc8 --- /dev/null +++ b/tests/checks/enumerator/test_models.py @@ -0,0 +1,257 @@ +"""Tests for datasure.checks.enumerator.models.""" + +import pytest +from pydantic import ValidationError + +from datasure.checks.enumerator.models import ( + ALLOWED_STATISTICS, + ALLOWED_STATISTICS_OVERTIME, + ALLOWED_TIME_PERIODS, + TAB_NAME, + WEEKDAY_NAMES, + WEEKDAY_OFFSET_MAP, + WEEKDAY_OFFSET_TO_NUMERIC, + ConsentOutcomeSettings, + EnumeratorOverviewMetrics, + EnumeratorSettings, + ProductivitySettings, + StatisticsOvertimeSettings, + StatisticsSettings, +) + +# ============================================ +# CONSTANTS TESTS +# ============================================ + + +def test_constants(): + """Test that all constants are defined correctly.""" + assert TAB_NAME == "enumerators" + assert len(ALLOWED_STATISTICS) == 8 + assert "count" in ALLOWED_STATISTICS + assert "mean" in ALLOWED_STATISTICS + assert "median" in ALLOWED_STATISTICS + assert len(ALLOWED_STATISTICS_OVERTIME) == 9 + assert "missing" in ALLOWED_STATISTICS_OVERTIME + assert len(ALLOWED_TIME_PERIODS) == 3 + assert "Daily" in ALLOWED_TIME_PERIODS + assert "Weekly" in ALLOWED_TIME_PERIODS + assert "Monthly" in ALLOWED_TIME_PERIODS + + +def test_weekday_constants(): + """Test weekday mapping constants.""" + assert len(WEEKDAY_NAMES) == 7 + assert "Monday" in WEEKDAY_NAMES + assert len(WEEKDAY_OFFSET_MAP) == 7 + assert WEEKDAY_OFFSET_MAP["Monday"] == "SUN" + assert len(WEEKDAY_OFFSET_TO_NUMERIC) == 7 + assert WEEKDAY_OFFSET_TO_NUMERIC["SUN"] == 0 + + +# ============================================ +# PYDANTIC MODELS TESTS +# ============================================ + + +def test_enumerator_settings_model_valid(): + """Test EnumeratorSettings model with valid data.""" + settings = EnumeratorSettings( + survey_key="survey_id", + survey_id="survey_id", + survey_date="submission_date", + enumerator="enumerator", + formversion="formversion", + duration="duration", + duration_unit="seconds", + team="team", + ) + assert settings.survey_key == "survey_id" + assert settings.survey_id == "survey_id" + assert settings.enumerator == "enumerator" + assert settings.team == "team" + + +def test_enumerator_settings_model_required_survey_id(): + """Test EnumeratorSettings model requires survey_id.""" + with pytest.raises(ValidationError): + EnumeratorSettings(survey_id="") + + +def test_enumerator_settings_model_optional_fields(): + """Test EnumeratorSettings model with optional fields.""" + settings = EnumeratorSettings( + survey_key="survey_id", + survey_id="survey_id", + ) + assert settings.survey_date is None + assert settings.enumerator is None + assert settings.team is None + + +def test_consent_outcome_settings_model(): + """Test ConsentOutcomeSettings model.""" + settings = ConsentOutcomeSettings( + consent="consent_col", + consent_vals=["yes", "agreed"], + outcome="outcome_col", + outcome_vals=["completed"], + ) + assert settings.consent == "consent_col" + assert settings.consent_vals == ["yes", "agreed"] + assert settings.outcome == "outcome_col" + assert settings.outcome_vals == ["completed"] + + +def test_consent_outcome_settings_defaults(): + """Test ConsentOutcomeSettings model with defaults.""" + settings = ConsentOutcomeSettings() + assert settings.consent is None + assert settings.consent_vals is None + assert settings.outcome is None + assert settings.outcome_vals is None + + +def test_productivity_settings_model(): + """Test ProductivitySettings model.""" + settings = ProductivitySettings( + view_option="Weekly", + weekstartday="Monday", + ) + assert settings.view_option == "Weekly" + assert settings.weekstartday == "Monday" + + +def test_productivity_settings_validation_view_option(): + """Test ProductivitySettings validation for view_option.""" + with pytest.raises(ValidationError, match="view_option must be one of"): + ProductivitySettings(view_option="Invalid") + + +def test_productivity_settings_validation_weekstartday(): + """Test ProductivitySettings validation for weekstartday.""" + with pytest.raises(ValidationError, match="weekstartday must be one of"): + ProductivitySettings(weekstartday="InvalidDay") + + +def test_productivity_settings_defaults(): + """Test ProductivitySettings model with defaults.""" + settings = ProductivitySettings() + assert settings.view_option == "Daily" + assert settings.weekstartday == "Monday" + + +def test_statistics_settings_model(): + """Test StatisticsSettings model.""" + settings = StatisticsSettings( + statscols=["age", "income"], + stats=["count", "mean", "median"], + ) + assert settings.statscols == ["age", "income"] + assert settings.stats == ["count", "mean", "median"] + + +def test_statistics_settings_validation(): + """Test StatisticsSettings validation for stats.""" + with pytest.raises(ValidationError, match="Invalid statistic"): + StatisticsSettings(stats=["invalid_stat"]) + + +def test_statistics_settings_defaults(): + """Test StatisticsSettings model with defaults.""" + settings = StatisticsSettings() + assert settings.statscols is None + assert settings.stats == ["count", "mean"] + + +def test_statistics_overtime_settings_model(): + """Test StatisticsOvertimeSettings model.""" + settings = StatisticsOvertimeSettings( + period_overtime="Weekly", + weekstartday="Tuesday", + stat="median", + statscol="age", + ) + assert settings.period_overtime == "Weekly" + assert settings.weekstartday == "Tuesday" + assert settings.stat == "median" + assert settings.statscol == "age" + + +def test_statistics_overtime_settings_validation_period(): + """Test StatisticsOvertimeSettings validation for period.""" + with pytest.raises(ValidationError, match="Invalid period"): + StatisticsOvertimeSettings(period_overtime="InvalidPeriod") + + +def test_statistics_overtime_settings_validation_weekstartday(): + """Test StatisticsOvertimeSettings validation for weekstartday.""" + with pytest.raises(ValidationError, match="Invalid weekstartday"): + StatisticsOvertimeSettings(weekstartday="InvalidDay") + + +def test_statistics_overtime_settings_validation_stat(): + """Test StatisticsOvertimeSettings validation for stat.""" + with pytest.raises(ValidationError, match="Invalid statistic"): + StatisticsOvertimeSettings(stat="invalid_stat") + + +def test_statistics_overtime_settings_defaults(): + """Test StatisticsOvertimeSettings model with defaults.""" + settings = StatisticsOvertimeSettings() + assert settings.period_overtime == "Week" + assert settings.weekstartday == "Monday" + assert settings.stat == "count" + assert settings.statscol is None + + +def test_enumerator_overview_metrics_model(): + """Test EnumeratorOverviewMetrics model.""" + metrics = EnumeratorOverviewMetrics( + all_submissions=100, + num_active_enumerators=5, + num_enumerators=10, + num_teams=3, + min_submissions=5, + max_submissions=25, + avg_submissions=15, + pct_active_enumerators="50%", + ) + assert metrics.all_submissions == 100 + assert metrics.num_active_enumerators == 5 + assert metrics.num_enumerators == 10 + assert metrics.num_teams == 3 + assert metrics.min_submissions == 5 + assert metrics.max_submissions == 25 + assert metrics.avg_submissions == 15 + assert metrics.pct_active_enumerators == "50%" + + +def test_enumerator_overview_metrics_validation(): + """Test EnumeratorOverviewMetrics validation for non-negative integers.""" + with pytest.raises(ValidationError): + EnumeratorOverviewMetrics( + all_submissions=-1, + num_active_enumerators=5, + num_enumerators=10, + num_teams=3, + min_submissions=5, + max_submissions=25, + avg_submissions=15, + pct_active_enumerators="50%", + ) + + +def test_enumerator_overview_metrics_num_teams_string(): + """Test EnumeratorOverviewMetrics allows string for num_teams.""" + metrics = EnumeratorOverviewMetrics( + all_submissions=100, + num_active_enumerators=5, + num_enumerators=10, + num_teams="n/a", + min_submissions=5, + max_submissions=25, + avg_submissions=15, + pct_active_enumerators="50%", + ) + assert metrics.num_teams == "n/a" diff --git a/tests/checks/enumerator/test_report_ui.py b/tests/checks/enumerator/test_report_ui.py new file mode 100644 index 00000000..ca4abda3 --- /dev/null +++ b/tests/checks/enumerator/test_report_ui.py @@ -0,0 +1,559 @@ +"""Tests for datasure.checks.enumerator.report_ui.""" + +import importlib +import sys +from unittest.mock import MagicMock, patch + +import polars as pl +import pytest + +from datasure.checks.enumerator.report_ui import ( + _get_numeric_columns, + _load_statistics_overtime_settings, + _load_statistics_settings, + _render_column_selector, + _render_column_selector_single, + _render_enumerator_overview_metrics, + _render_enumerator_productivity, + _render_enumerator_statistics, + _render_enumerator_statistics_overtime, + _render_period_selector_overtime, + _render_statistic_selector, + _render_statistics_selector, + _render_time_period_selector, + _render_weekday_selector, + _render_weekday_selector_overtime, +) +from datasure.models.schemas import ColumnByType +from tests.checks.enumerator.conftest import make_mock_st + +# ============================================ +# PATCHED_ENUM FIXTURE (patch st in report_ui for non-fragment UI functions) +# ============================================ + + +@pytest.fixture +def patched_enum(): + """Patch st in report_ui module for non-fragment UI function tests.""" + mock_st = make_mock_st() + with ( + patch("datasure.checks.enumerator.report_ui.st", mock_st), + patch("datasure.checks.enumerator.report_ui.save_check_settings"), + patch( + "datasure.checks.enumerator.report_ui.load_check_settings", return_value={} + ), + patch("datasure.checks.enumerator.report_ui.trigger_save"), + patch( + "datasure.checks.enumerator.report_ui.duckdb_get_table", + return_value=pl.DataFrame(), + ), + patch("datasure.checks.enumerator.report_ui.demo_callout"), + patch("datasure.utils.onboarding_utils.is_demo_project", return_value=False), + ): + yield mock_st + + +# ============================================ +# ENUM_BC FIXTURE (reload compute/settings_ui/report_ui with mocked streamlit) +# ============================================ + + +@pytest.fixture +def enum_bc(): + """Reload enumerator submodules with mocked Streamlit to strip fragments.""" + mock_st = make_mock_st() + original_st = sys.modules.get("streamlit") + sys.modules["streamlit"] = mock_st + + import datasure.checks.enumerator.compute as compute_module + import datasure.checks.enumerator.report_ui as report_ui_module + import datasure.checks.enumerator.settings_ui as settings_ui_module + + try: + # Reload in dependency order so decorators pick up the mocked st and + # cross-module references (report_ui imports from settings_ui) stay wired. + importlib.reload(compute_module) + importlib.reload(settings_ui_module) + importlib.reload(report_ui_module) + + compute_module.load_missing_codes_from_db = MagicMock( + return_value=pl.DataFrame() + ) + + settings_ui_module.load_check_settings = MagicMock(return_value={}) + settings_ui_module.save_check_settings = MagicMock() + settings_ui_module.trigger_save = MagicMock() + settings_ui_module.duckdb_save_table = MagicMock() + + report_ui_module.load_check_settings = MagicMock(return_value={}) + report_ui_module.save_check_settings = MagicMock() + report_ui_module.trigger_save = MagicMock() + report_ui_module.duckdb_get_table = MagicMock(return_value=pl.DataFrame()) + report_ui_module.demo_callout = MagicMock() + + with patch( + "datasure.utils.onboarding_utils.is_demo_project", return_value=False + ): + yield report_ui_module + finally: + if original_st is not None: + sys.modules["streamlit"] = original_st + else: + sys.modules.pop("streamlit", None) + importlib.reload(compute_module) + importlib.reload(settings_ui_module) + importlib.reload(report_ui_module) + + +# ============================================ +# HELPER FUNCTIONS TESTS +# ============================================ + + +def test_get_numeric_columns(): + """Test _get_numeric_columns helper function.""" + data = pl.DataFrame( + { + "age": [25, 30, 35], + "income": [50000, 60000, 55000], + "name": ["Alice", "Bob", "Charlie"], + "is_active": [True, False, True], + } + ) + + result = _get_numeric_columns(data) + assert "age" in result + assert "income" in result + assert "name" not in result + assert "is_active" not in result + + +def test_get_numeric_columns_with_exclude(): + """Test _get_numeric_columns with exclude list.""" + data = pl.DataFrame( + { + "age": [25, 30, 35], + "income": [50000, 60000, 55000], + "duration": [3600, 4200, 3800], + } + ) + + result = _get_numeric_columns(data, exclude_cols=["duration"]) + assert "age" in result + assert "income" in result + assert "duration" not in result + + +def test_get_numeric_columns_empty_dataframe(): + """Test _get_numeric_columns with empty DataFrame.""" + data = pl.DataFrame() + result = _get_numeric_columns(data) + assert result == [] + + +def test_get_numeric_columns_no_numeric(): + """Test _get_numeric_columns with no numeric columns.""" + data = pl.DataFrame( + { + "name": ["Alice", "Bob"], + "city": ["NYC", "LA"], + } + ) + result = _get_numeric_columns(data) + assert result == [] + + +# ============================================ +# PATCHED_ENUM UI FUNCTION TESTS +# ============================================ + + +def test_render_enumerator_overview_no_date_enum(patched_enum): + """_render_enumerator_overview_metrics shows info when date/enum is None.""" + _render_enumerator_overview_metrics(pl.DataFrame(), None, None, None) + patched_enum.info.assert_called() + + +def test_render_enumerator_overview_with_data(patched_enum, sample_enumerator_data): + """_render_enumerator_overview_metrics renders metrics with valid data.""" + _render_enumerator_overview_metrics( + sample_enumerator_data, "submission_date", "enumerator", "team" + ) + patched_enum.columns.assert_called() + + +def test_render_enumerator_overview_no_team(patched_enum, sample_enumerator_data): + """_render_enumerator_overview_metrics renders without team column.""" + _render_enumerator_overview_metrics( + sample_enumerator_data, "submission_date", "enumerator", None + ) + patched_enum.columns.assert_called() + + +def test_render_time_period_selector_default(patched_enum): + """_render_time_period_selector returns Day when pills returns None.""" + result = _render_time_period_selector("settings.json") + assert result == "Day" + patched_enum.pills.assert_called() + + +def test_render_time_period_selector_week(patched_enum): + """_render_time_period_selector returns Week when pills returns Week.""" + patched_enum.pills.return_value = "Week" + result = _render_time_period_selector("settings.json") + assert result == "Week" + + +def test_render_weekday_selector(patched_enum): + """_render_weekday_selector returns offset code for the selected weekday.""" + patched_enum.selectbox.return_value = "Monday" + result = _render_weekday_selector("settings.json") + assert result == "SUN" + + +def test_load_statistics_settings_default(patched_enum): + """_load_statistics_settings returns default StatisticsSettings.""" + result = _load_statistics_settings("settings.json") + assert result.stats == ["count", "mean"] + + +def test_render_column_selector(patched_enum): + """_render_column_selector returns list from multiselect.""" + result = _render_column_selector(["age", "income"], None, "settings.json") + assert isinstance(result, list) + patched_enum.multiselect.assert_called() + + +def test_render_statistics_selector(patched_enum): + """_render_statistics_selector returns list from multiselect.""" + result = _render_statistics_selector(["count", "mean"], "settings.json") + assert isinstance(result, list) + patched_enum.multiselect.assert_called() + + +def test_render_enumerator_statistics_no_enum(patched_enum): + """_render_enumerator_statistics shows info when enumerator is None.""" + _render_enumerator_statistics(pl.DataFrame(), None, None, "settings.json") + patched_enum.info.assert_called() + + +def test_load_statistics_overtime_settings_default(patched_enum): + """_load_statistics_overtime_settings returns default settings.""" + result = _load_statistics_overtime_settings("settings.json") + assert result.stat == "count" + + +def test_render_period_selector_overtime_default(patched_enum): + """_render_period_selector_overtime returns Day when pills returns None.""" + result = _render_period_selector_overtime("settings.json") + assert result == "Day" + patched_enum.pills.assert_called() + + +def test_render_period_selector_overtime_week(patched_enum): + """_render_period_selector_overtime returns Week when pills returns Week.""" + patched_enum.pills.return_value = "Week" + result = _render_period_selector_overtime("settings.json", default_period="Week") + assert result == "Week" + + +def test_render_weekday_selector_overtime(patched_enum): + """_render_weekday_selector_overtime returns offset code.""" + patched_enum.selectbox.return_value = "Tuesday" + result = _render_weekday_selector_overtime("Monday", "settings.json") + assert result == "MON" + + +def test_render_statistic_selector(patched_enum): + """_render_statistic_selector returns the selectbox value.""" + patched_enum.selectbox.return_value = "mean" + result = _render_statistic_selector("count", "settings.json") + assert result == "mean" + + +def test_render_column_selector_single_default_none(patched_enum): + """_render_column_selector_single returns None when selectbox returns None.""" + result = _render_column_selector_single(["age", "income"], None, "settings.json") + assert result is None + patched_enum.selectbox.assert_called() + + +def test_render_column_selector_single_with_value(patched_enum): + """_render_column_selector_single returns the selectbox value.""" + patched_enum.selectbox.return_value = "age" + result = _render_column_selector_single(["age", "income"], "age", "settings.json") + assert result == "age" + + +def test_render_enumerator_productivity_no_enum(patched_enum): + """_render_enumerator_productivity shows info when enum/date is None.""" + _render_enumerator_productivity(pl.DataFrame(), None, None, None, "settings.json") + patched_enum.info.assert_called() + + +def test_render_enumerator_statistics_overtime_no_enum(patched_enum): + """_render_enumerator_statistics_overtime shows info when enum/date is None.""" + _render_enumerator_statistics_overtime( + pl.DataFrame(), None, None, None, "settings.json" + ) + patched_enum.info.assert_called() + + +def test_load_statistics_settings_fallback(patched_enum): + """_load_statistics_settings returns defaults when saved settings invalid.""" + with patch( + "datasure.checks.enumerator.report_ui.load_check_settings", + return_value={"stats": ["not_a_real_stat"]}, + ): + result = _load_statistics_settings("settings.json") + assert result.stats == ["count", "mean"] + + +def test_load_statistics_overtime_settings_fallback(patched_enum): + """_load_statistics_overtime_settings returns defaults when settings invalid.""" + with patch( + "datasure.checks.enumerator.report_ui.load_check_settings", + return_value={"period_overtime": "bad_period"}, + ): + result = _load_statistics_overtime_settings("settings.json") + assert result.stat == "count" + + +# ============================================ +# ENUM_BC UI FRAGMENT TESTS +# ============================================ + + +def test_render_enumerator_summary_table_no_date_enum(enum_bc): + """_render_enumerator_summary_table shows info when date/enum is None.""" + enum_bc._render_enumerator_summary_table( + "proj", pl.DataFrame(), None, None, None, None, None + ) + enum_bc.st.info.assert_called() + + +def test_render_enumerator_summary_table_with_data(enum_bc, sample_enumerator_data): + """_render_enumerator_summary_table renders table with valid data.""" + enum_bc._render_enumerator_summary_table( + "proj", + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "formversion", + "duration", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_summary_table_with_show_info( + enum_bc, sample_enumerator_data +): + """_render_enumerator_summary_table filters columns when show_info selected.""" + enum_bc.st.pills.return_value = ["submissions"] + enum_bc._render_enumerator_summary_table( + "proj", + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "formversion", + "duration", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_productivity_table_day(enum_bc, sample_enumerator_data): + """_render_enumerator_productivity_table renders with Day period.""" + enum_bc._render_enumerator_productivity_table( + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "settings.json", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_productivity_table_week(enum_bc, sample_enumerator_data): + """_render_enumerator_productivity_table renders weekday selector for Week.""" + enum_bc.st.pills.return_value = "Week" + enum_bc.st.selectbox.return_value = "Monday" + enum_bc._render_enumerator_productivity_table( + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "settings.json", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_statistics_table_no_enum(enum_bc): + """_render_enumerator_statistics_table shows info when enum is None.""" + enum_bc._render_enumerator_statistics_table(pl.DataFrame(), None, None, "s.json") + enum_bc.st.info.assert_called() + + +def test_render_enumerator_statistics_table_no_cols(enum_bc, sample_enumerator_data): + """_render_enumerator_statistics_table shows info when no cols selected.""" + enum_bc._render_enumerator_statistics_table( + sample_enumerator_data, "enumerator", None, "settings.json" + ) + enum_bc.st.info.assert_called() + + +def test_render_enumerator_statistics_table_with_cols(enum_bc, sample_enumerator_data): + """_render_enumerator_statistics_table renders table when cols selected.""" + enum_bc.st.multiselect.side_effect = [["age"], ["count", "mean"]] + enum_bc._render_enumerator_statistics_table( + sample_enumerator_data, "enumerator", "team", "settings.json" + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_statistics_with_enum(enum_bc, sample_enumerator_data): + """_render_enumerator_statistics calls the fragment table function.""" + enum_bc._render_enumerator_statistics( + sample_enumerator_data, "enumerator", None, "settings.json" + ) + enum_bc.st.info.assert_called() + + +def test_render_enumerator_statistics_overtime_table_no_statscol( + enum_bc, sample_enumerator_data +): + """_render_enumerator_statistics_overtime_table shows info when statscol None.""" + enum_bc._render_enumerator_statistics_overtime_table( + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "settings.json", + ) + enum_bc.st.info.assert_called() + + +def test_render_enumerator_statistics_overtime_table_with_col( + enum_bc, sample_enumerator_data +): + """_render_enumerator_statistics_overtime_table renders with valid col.""" + enum_bc.st.selectbox.side_effect = ["age", "count"] + enum_bc._render_enumerator_statistics_overtime_table( + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "settings.json", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_statistics_overtime_table_week_period( + enum_bc, sample_enumerator_data +): + """_render_enumerator_statistics_overtime_table handles Week period.""" + enum_bc.st.pills.return_value = "Week" + enum_bc.st.selectbox.side_effect = ["age", "count", "Monday"] + enum_bc._render_enumerator_statistics_overtime_table( + sample_enumerator_data, + "submission_date", + "enumerator", + "team", + "settings.json", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_productivity_with_valid_params( + enum_bc, sample_enumerator_data +): + """_render_enumerator_productivity calls the table fragment when params valid.""" + enum_bc._render_enumerator_productivity( + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "settings.json", + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_statistics_overtime_with_valid_params( + enum_bc, sample_enumerator_data +): + """_render_enumerator_statistics_overtime calls the table fragment.""" + enum_bc._render_enumerator_statistics_overtime( + sample_enumerator_data, + "submission_date", + "enumerator", + None, + "settings.json", + ) + enum_bc.st.info.assert_called() + + +def test_enumerator_report_empty_data(enum_bc): + """enumerator_report shows info and returns early when data is empty.""" + survey_cols = ColumnByType( + categorical_columns=["survey_id", "enumerator"], + datetime_columns=["submission_date"], + ) + enum_bc.enumerator_report( + "proj_id", + pl.DataFrame(), + "settings.json", + {"survey_id": "survey_id"}, + survey_cols, + ) + enum_bc.st.info.assert_called() + + +def test_enumerator_report_with_data(enum_bc, sample_enumerator_data): + """enumerator_report renders full report with valid data.""" + survey_cols = ColumnByType( + categorical_columns=list(sample_enumerator_data.columns), + datetime_columns=["submission_date"], + ) + # `_render_consent_outcome_settings` lives in settings_ui (called internally by + # enumerator_report_settings) rather than report_ui, so it must be patched on + # the module that actually owns it for the stub to take effect. + from datasure.checks.enumerator import settings_ui as settings_ui_module + + settings_ui_module._render_consent_outcome_settings = MagicMock() + + def _selectbox_side_effect(label, options=None, **kwargs): + if options and "seconds" in options: + return "seconds" + return None + + enum_bc.st.selectbox.side_effect = _selectbox_side_effect + enum_bc.enumerator_report( + "proj_id", + sample_enumerator_data, + "settings.json", + {"survey_id": "survey_id"}, + survey_cols, + ) + enum_bc.st.title.assert_called() + + +def test_render_enumerator_statistics_table_with_cols_no_team( + enum_bc, sample_enumerator_data +): + """_render_enumerator_statistics_table uses else column config when no team.""" + enum_bc.st.multiselect.side_effect = [["age"], ["count"]] + enum_bc._render_enumerator_statistics_table( + sample_enumerator_data, "enumerator", None, "settings.json" + ) + enum_bc.st.dataframe.assert_called() + + +def test_render_enumerator_statistics_overtime_table_early_return( + enum_bc, sample_enumerator_data +): + """_render_enumerator_statistics_overtime_table returns early when enum None.""" + enum_bc._render_enumerator_statistics_overtime_table( + sample_enumerator_data, "submission_date", None, None, "settings.json" + ) + enum_bc.st.dataframe.assert_not_called() diff --git a/tests/checks/enumerator/test_settings_ui.py b/tests/checks/enumerator/test_settings_ui.py new file mode 100644 index 00000000..77365608 --- /dev/null +++ b/tests/checks/enumerator/test_settings_ui.py @@ -0,0 +1,259 @@ +"""Tests for datasure.checks.enumerator.settings_ui.""" + +import importlib +import sys +from unittest.mock import MagicMock, patch + +import polars as pl +import pytest + +from datasure.checks.enumerator.models import ConsentOutcomeSettings, EnumeratorSettings +from datasure.checks.enumerator.settings_ui import ( + _create_enum_data_on_settings, + _trigger_success_message, + load_default_enumerator_settings, +) +from tests.checks.enumerator.conftest import make_mock_st + +# ============================================ +# ENUM_BC FIXTURE (reload settings_ui with mocked streamlit) +# ============================================ + + +@pytest.fixture +def enum_bc(): + """Reload settings_ui module with mocked Streamlit to strip fragment decorators.""" + mock_st = make_mock_st() + original_st = sys.modules.get("streamlit") + sys.modules["streamlit"] = mock_st + import datasure.checks.enumerator.settings_ui as settings_ui_module + + try: + importlib.reload(settings_ui_module) + settings_ui_module.load_check_settings = MagicMock(return_value={}) + settings_ui_module.save_check_settings = MagicMock() + settings_ui_module.trigger_save = MagicMock() + settings_ui_module.duckdb_save_table = MagicMock() + with patch( + "datasure.utils.onboarding_utils.is_demo_project", return_value=False + ): + yield settings_ui_module + finally: + if original_st is not None: + sys.modules["streamlit"] = original_st + else: + sys.modules.pop("streamlit", None) + importlib.reload(settings_ui_module) + + +# ============================================ +# SETTINGS TESTS +# ============================================ + + +def test_load_default_enumerator_settings_valid(enumerator_settings_file): + """Test loading enumerator settings from valid file.""" + config = EnumeratorSettings( + survey_key="default_key", + survey_id="default_id", + ) + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = load_default_enumerator_settings(enumerator_settings_file, config) + + # Saved settings should override defaults + assert result.survey_key == "survey_id" + assert result.enumerator == "enumerator" + + +def test_load_default_enumerator_settings_missing_file(): + """Test loading enumerator settings when file doesn't exist.""" + config = EnumeratorSettings( + survey_key="default_key", + survey_id="default_id", + enumerator="default_enum", + ) + with patch("streamlit.cache_data", lambda ttl: lambda f: f): + result = load_default_enumerator_settings("nonexistent.json", config) + + # Should return default config when file doesn't exist + assert result.survey_key == "default_key" + assert result.enumerator == "default_enum" + + +@patch("datasure.checks.enumerator.settings_ui.st") +def test_trigger_success_message(mock_st): + """Test _trigger_success_message function.""" + mock_st.session_state = {} + + _trigger_success_message("test_button") + assert mock_st.session_state["test_button"] is True + + +def test_create_enum_data_on_settings_with_consent_and_outcome(): + """Test _create_enum_data_on_settings with consent and outcome values.""" + data = pl.DataFrame( + { + "survey_id": ["S001", "S002", "S003"], + "consent": ["yes", "no", "yes"], + "outcome": ["completed", "incomplete", "completed"], + } + ) + + config = ConsentOutcomeSettings( + consent="consent", + consent_vals=["yes"], + outcome="outcome", + outcome_vals=["completed"], + ) + + with patch("datasure.checks.enumerator.settings_ui.duckdb_save_table") as mock_save: + _create_enum_data_on_settings("test_project", data, config) + + # Verify function was called + mock_save.assert_called_once() + saved_data = mock_save.call_args[0][1] + + # Check that consent and outcome columns were created + assert "consent_granted_agg_col" in saved_data.columns + assert "completed_survey_agg_col" in saved_data.columns + assert saved_data["consent_granted_agg_col"].to_list() == [1, 0, 1] + assert saved_data["completed_survey_agg_col"].to_list() == [1, 0, 1] + + +def test_create_enum_data_on_settings_without_consent(): + """Test _create_enum_data_on_settings without consent values.""" + data = pl.DataFrame( + { + "survey_id": ["S001", "S002"], + "outcome": ["completed", "completed"], + } + ) + + config = ConsentOutcomeSettings( + consent=None, + consent_vals=None, + outcome="outcome", + outcome_vals=["completed"], + ) + + with patch("datasure.checks.enumerator.settings_ui.duckdb_save_table") as mock_save: + _create_enum_data_on_settings("test_project", data, config) + + saved_data = mock_save.call_args[0][1] + # Consent should default to 1 + assert saved_data["consent_granted_agg_col"].to_list() == [1, 1] + + +def test_create_enum_data_on_settings_without_outcome(): + """Test _create_enum_data_on_settings without outcome values.""" + data = pl.DataFrame( + { + "survey_id": ["S001", "S002"], + "consent": ["yes", "yes"], + } + ) + + config = ConsentOutcomeSettings( + consent="consent", + consent_vals=["yes"], + outcome=None, + outcome_vals=None, + ) + + with patch("datasure.checks.enumerator.settings_ui.duckdb_save_table") as mock_save: + _create_enum_data_on_settings("test_project", data, config) + + saved_data = mock_save.call_args[0][1] + # Outcome should default to 1 + assert saved_data["completed_survey_agg_col"].to_list() == [1, 1] + + +# ============================================ +# INTEGRATION TESTS +# ============================================ + + +@patch("datasure.checks.enumerator.settings_ui.duckdb_save_table") +def test_consent_outcome_integration(mock_save, sample_enumerator_data): + """Test consent and outcome settings integration.""" + # Create consent and outcome settings + config = ConsentOutcomeSettings( + consent="consent", + consent_vals=["yes"], + outcome="outcome", + outcome_vals=["completed"], + ) + + # Add consent and outcome columns + data = sample_enumerator_data.with_columns( + [ + pl.lit("yes").alias("consent"), + pl.lit("completed").alias("outcome"), + ] + ) + + # Create enum data with settings + _create_enum_data_on_settings("test_project", data, config) + + # Verify the data was saved + assert mock_save.called + + +# ============================================ +# ENUM_BC UI FRAGMENT TESTS +# ============================================ + + +def test_enumerator_report_settings_basic(enum_bc, sample_enumerator_data): + """enumerator_report_settings returns EnumeratorSettings from UI.""" + enum_bc.st.selectbox.return_value = "survey_id" + categorical_cols = list(sample_enumerator_data.columns) + datetime_cols = ["submission_date"] + config = EnumeratorSettings(survey_id="survey_id") + result = enum_bc.enumerator_report_settings( + "proj_id", + "settings.json", + sample_enumerator_data, + config, + categorical_cols, + datetime_cols, + ) + assert result is not None + + +def test_render_consent_outcome_settings(enum_bc, sample_enumerator_data): + """_render_consent_outcome_settings renders consent/outcome selectors.""" + enum_bc.st.selectbox.return_value = "survey_id" + categorical_cols = list(sample_enumerator_data.columns) + enum_bc._render_consent_outcome_settings( + "proj_id", sample_enumerator_data, categorical_cols, "settings.json" + ) + enum_bc.st.button.assert_called() + + +def test_render_consent_outcome_settings_button_click(enum_bc, sample_enumerator_data): + """_render_consent_outcome_settings calls create when button clicked.""" + enum_bc.st.selectbox.return_value = "survey_id" + enum_bc.st.button.return_value = True + categorical_cols = list(sample_enumerator_data.columns) + enum_bc._render_consent_outcome_settings( + "proj_id", sample_enumerator_data, categorical_cols, "settings.json" + ) + enum_bc.duckdb_save_table.assert_called() + + +def test_enumerator_report_settings_success_flag(enum_bc, sample_enumerator_data): + """enumerator_report_settings shows success message when consent flag set.""" + enum_bc.st.selectbox.return_value = "survey_id" + enum_bc.st.session_state["st_apply_consent_outcome_enumerator"] = True + categorical_cols = list(sample_enumerator_data.columns) + config = EnumeratorSettings(survey_id="survey_id") + enum_bc.enumerator_report_settings( + "proj_id", + "settings.json", + sample_enumerator_data, + config, + categorical_cols, + ["submission_date"], + ) + enum_bc.st.success.assert_called() diff --git a/tests/checks/test_enumerator.py b/tests/checks/test_enumerator.py deleted file mode 100644 index 380892c8..00000000 --- a/tests/checks/test_enumerator.py +++ /dev/null @@ -1,1921 +0,0 @@ -"""Tests for enumerator module. - -This module tests the refactored enumerator performance analysis system using -Polars DataFrames and Pydantic models for validation and configuration. -""" - -import importlib -import json -import sys -from datetime import date, timedelta -from unittest.mock import MagicMock, patch - -import polars as pl -import pytest -from polars.exceptions import ColumnNotFoundError -from pydantic import ValidationError - -from datasure.checks.enumerator import ( - ALLOWED_STATISTICS, - ALLOWED_STATISTICS_OVERTIME, - ALLOWED_TIME_PERIODS, - TAB_NAME, - WEEKDAY_NAMES, - WEEKDAY_OFFSET_MAP, - WEEKDAY_OFFSET_TO_NUMERIC, - ConsentOutcomeSettings, - EnumeratorOverviewMetrics, - EnumeratorSettings, - ProductivitySettings, - StatisticsOvertimeSettings, - StatisticsSettings, - _create_enum_data_on_settings, - _get_numeric_columns, - _load_statistics_overtime_settings, - _load_statistics_settings, - _render_column_selector, - _render_column_selector_single, - _render_enumerator_overview_metrics, - _render_enumerator_productivity, - _render_enumerator_statistics, - _render_enumerator_statistics_overtime, - _render_period_selector_overtime, - _render_statistic_selector, - _render_statistics_selector, - _render_time_period_selector, - _render_weekday_selector, - _render_weekday_selector_overtime, - _trigger_success_message, - compute_enumerator_missing_table, - compute_enumerator_overview, - compute_enumerator_productivity, - compute_enumerator_statistics, - compute_enumerator_statistics_overtime, - compute_enumerator_summary, - load_default_enumerator_settings, -) -from datasure.models.schemas import ColumnByType - -# ============================================ -# MOCK STREAMLIT HELPERS AND UI FIXTURES -# ============================================ - - -def _make_mock_st(): - """Create a mock Streamlit module for testing UI render functions.""" - - def make_col(): - col = MagicMock() - col.number_input.return_value = 0.0 - col.selectbox.return_value = None - col.text_input.return_value = "" - col.multiselect.return_value = [] - return col - - def _col_factory(n_or_spec, **kwargs): - if isinstance(n_or_spec, int): - n = n_or_spec - elif isinstance(n_or_spec, list | tuple): - n = len(n_or_spec) - else: - n = 2 - return tuple(make_col() for _ in range(n)) - - mock_st = MagicMock() - mock_st.fragment = lambda func: func - mock_st.dialog = lambda *args, **kwargs: lambda func: func - mock_st.columns.side_effect = _col_factory - mock_st.selectbox.return_value = None - mock_st.multiselect.return_value = [] - mock_st.pills.return_value = None - mock_st.button.return_value = False - mock_st.toggle.return_value = False - mock_st.number_input.return_value = 0 - mock_st.text_input.return_value = "" - mock_st.session_state = {} - return mock_st - - -@pytest.fixture -def patched_enum(): - """Patch st in enumerator module for non-fragment UI function tests.""" - mock_st = _make_mock_st() - with ( - patch("datasure.checks.enumerator.st", mock_st), - patch("datasure.checks.enumerator.save_check_settings"), - patch("datasure.checks.enumerator.load_check_settings", return_value={}), - patch("datasure.checks.enumerator.trigger_save"), - patch( - "datasure.checks.enumerator.duckdb_get_table", - return_value=pl.DataFrame(), - ), - patch("datasure.checks.enumerator.duckdb_save_table"), - patch("datasure.checks.enumerator.demo_callout"), - patch("datasure.utils.onboarding_utils.is_demo_project", return_value=False), - ): - yield mock_st - - -@pytest.fixture -def enum_bc(): - """Reload enumerator module with mocked Streamlit to strip fragment decorators.""" - mock_st = _make_mock_st() - original_st = sys.modules.get("streamlit") - sys.modules["streamlit"] = mock_st - import datasure.checks.enumerator as enum_module - - try: - importlib.reload(enum_module) - enum_module.load_check_settings = MagicMock(return_value={}) - enum_module.save_check_settings = MagicMock() - enum_module.trigger_save = MagicMock() - enum_module.duckdb_get_table = MagicMock(return_value=pl.DataFrame()) - enum_module.duckdb_save_table = MagicMock() - enum_module.demo_callout = MagicMock() - enum_module.load_missing_codes_from_db = MagicMock(return_value=pl.DataFrame()) - with patch( - "datasure.utils.onboarding_utils.is_demo_project", return_value=False - ): - yield enum_module - finally: - if original_st is not None: - sys.modules["streamlit"] = original_st - else: - sys.modules.pop("streamlit", None) - importlib.reload(enum_module) - - -# ============================================ -# FIXTURES FOR ENUMERATOR-SPECIFIC DATA -# ============================================ - - -@pytest.fixture(autouse=True) -def mock_database_functions(monkeypatch): - """Override the autouse fixture from conftest. - - Disables database mocking for these tests. - """ - pass - - -@pytest.fixture -def sample_enumerator_data(): - """Create sample enumerator data as Polars DataFrame.""" - today = date.today() - return pl.DataFrame( - { - "survey_id": ["S001", "S002", "S003", "S004", "S005", "S006"], - "submission_date": [ - today - timedelta(days=1), - today - timedelta(days=2), - today - timedelta(days=3), - today - timedelta(days=1), - today - timedelta(days=8), - today, - ], - "enumerator": ["E1", "E1", "E2", "E2", "E3", "E1"], - "team": ["T1", "T1", "T1", "T2", "T2", "T1"], - "duration": [3600, 4200, 3800, 4000, 3900, 3700], - "formversion": ["v1", "v1", "v2", "v2", "v1", "v2"], - "age": [25, 30, 35, 28, 32, 27], - "income": [50000, 60000, 55000, 52000, 58000, 51000], - "consent_granted_agg_col": [1, 1, 1, 0, 1, 1], - "completed_survey_agg_col": [1, 1, 1, 1, 0, 1], - } - ) - - -@pytest.fixture -def sample_enumerator_settings(): - """Create sample EnumeratorSettings for testing.""" - return EnumeratorSettings( - survey_key="survey_id", - survey_id="survey_id", - survey_date="submission_date", - enumerator="enumerator", - formversion="formversion", - duration="duration", - duration_unit="seconds", - team="team", - ) - - -@pytest.fixture -def sample_missing_codes_config(): - """Create sample missing codes configuration.""" - return pl.DataFrame( - { - "label": ["Refused", "Don't know"], - "codes": ["-99", "-88"], - } - ) - - -@pytest.fixture -def enumerator_settings_file(tmp_path): - """Create a temporary enumerator settings file.""" - settings = { - "enumerators": { - "survey_key": "survey_id", - "survey_id": "survey_id", - "survey_date": "submission_date", - "enumerator": "enumerator", - "team": "team", - } - } - file_path = tmp_path / "enumerator_settings.json" - file_path.write_text(json.dumps(settings)) - return str(file_path) - - -# ============================================ -# CONSTANTS TESTS -# ============================================ - - -def test_constants(): - """Test that all constants are defined correctly.""" - assert TAB_NAME == "enumerators" - assert len(ALLOWED_STATISTICS) == 8 - assert "count" in ALLOWED_STATISTICS - assert "mean" in ALLOWED_STATISTICS - assert "median" in ALLOWED_STATISTICS - assert len(ALLOWED_STATISTICS_OVERTIME) == 9 - assert "missing" in ALLOWED_STATISTICS_OVERTIME - assert len(ALLOWED_TIME_PERIODS) == 3 - assert "Daily" in ALLOWED_TIME_PERIODS - assert "Weekly" in ALLOWED_TIME_PERIODS - assert "Monthly" in ALLOWED_TIME_PERIODS - - -def test_weekday_constants(): - """Test weekday mapping constants.""" - assert len(WEEKDAY_NAMES) == 7 - assert "Monday" in WEEKDAY_NAMES - assert len(WEEKDAY_OFFSET_MAP) == 7 - assert WEEKDAY_OFFSET_MAP["Monday"] == "SUN" - assert len(WEEKDAY_OFFSET_TO_NUMERIC) == 7 - assert WEEKDAY_OFFSET_TO_NUMERIC["SUN"] == 0 - - -# ============================================ -# PYDANTIC MODELS TESTS -# ============================================ - - -def test_enumerator_settings_model_valid(): - """Test EnumeratorSettings model with valid data.""" - settings = EnumeratorSettings( - survey_key="survey_id", - survey_id="survey_id", - survey_date="submission_date", - enumerator="enumerator", - formversion="formversion", - duration="duration", - duration_unit="seconds", - team="team", - ) - assert settings.survey_key == "survey_id" - assert settings.survey_id == "survey_id" - assert settings.enumerator == "enumerator" - assert settings.team == "team" - - -def test_enumerator_settings_model_required_survey_id(): - """Test EnumeratorSettings model requires survey_id.""" - with pytest.raises(ValidationError): - EnumeratorSettings(survey_id="") - - -def test_enumerator_settings_model_optional_fields(): - """Test EnumeratorSettings model with optional fields.""" - settings = EnumeratorSettings( - survey_key="survey_id", - survey_id="survey_id", - ) - assert settings.survey_date is None - assert settings.enumerator is None - assert settings.team is None - - -def test_consent_outcome_settings_model(): - """Test ConsentOutcomeSettings model.""" - settings = ConsentOutcomeSettings( - consent="consent_col", - consent_vals=["yes", "agreed"], - outcome="outcome_col", - outcome_vals=["completed"], - ) - assert settings.consent == "consent_col" - assert settings.consent_vals == ["yes", "agreed"] - assert settings.outcome == "outcome_col" - assert settings.outcome_vals == ["completed"] - - -def test_consent_outcome_settings_defaults(): - """Test ConsentOutcomeSettings model with defaults.""" - settings = ConsentOutcomeSettings() - assert settings.consent is None - assert settings.consent_vals is None - assert settings.outcome is None - assert settings.outcome_vals is None - - -def test_productivity_settings_model(): - """Test ProductivitySettings model.""" - settings = ProductivitySettings( - view_option="Weekly", - weekstartday="Monday", - ) - assert settings.view_option == "Weekly" - assert settings.weekstartday == "Monday" - - -def test_productivity_settings_validation_view_option(): - """Test ProductivitySettings validation for view_option.""" - with pytest.raises(ValidationError, match="view_option must be one of"): - ProductivitySettings(view_option="Invalid") - - -def test_productivity_settings_validation_weekstartday(): - """Test ProductivitySettings validation for weekstartday.""" - with pytest.raises(ValidationError, match="weekstartday must be one of"): - ProductivitySettings(weekstartday="InvalidDay") - - -def test_productivity_settings_defaults(): - """Test ProductivitySettings model with defaults.""" - settings = ProductivitySettings() - assert settings.view_option == "Daily" - assert settings.weekstartday == "Monday" - - -def test_statistics_settings_model(): - """Test StatisticsSettings model.""" - settings = StatisticsSettings( - statscols=["age", "income"], - stats=["count", "mean", "median"], - ) - assert settings.statscols == ["age", "income"] - assert settings.stats == ["count", "mean", "median"] - - -def test_statistics_settings_validation(): - """Test StatisticsSettings validation for stats.""" - with pytest.raises(ValidationError, match="Invalid statistic"): - StatisticsSettings(stats=["invalid_stat"]) - - -def test_statistics_settings_defaults(): - """Test StatisticsSettings model with defaults.""" - settings = StatisticsSettings() - assert settings.statscols is None - assert settings.stats == ["count", "mean"] - - -def test_statistics_overtime_settings_model(): - """Test StatisticsOvertimeSettings model.""" - settings = StatisticsOvertimeSettings( - period_overtime="Weekly", - weekstartday="Tuesday", - stat="median", - statscol="age", - ) - assert settings.period_overtime == "Weekly" - assert settings.weekstartday == "Tuesday" - assert settings.stat == "median" - assert settings.statscol == "age" - - -def test_statistics_overtime_settings_validation_period(): - """Test StatisticsOvertimeSettings validation for period.""" - with pytest.raises(ValidationError, match="Invalid period"): - StatisticsOvertimeSettings(period_overtime="InvalidPeriod") - - -def test_statistics_overtime_settings_validation_weekstartday(): - """Test StatisticsOvertimeSettings validation for weekstartday.""" - with pytest.raises(ValidationError, match="Invalid weekstartday"): - StatisticsOvertimeSettings(weekstartday="InvalidDay") - - -def test_statistics_overtime_settings_validation_stat(): - """Test StatisticsOvertimeSettings validation for stat.""" - with pytest.raises(ValidationError, match="Invalid statistic"): - StatisticsOvertimeSettings(stat="invalid_stat") - - -def test_statistics_overtime_settings_defaults(): - """Test StatisticsOvertimeSettings model with defaults.""" - settings = StatisticsOvertimeSettings() - assert settings.period_overtime == "Week" - assert settings.weekstartday == "Monday" - assert settings.stat == "count" - assert settings.statscol is None - - -def test_enumerator_overview_metrics_model(): - """Test EnumeratorOverviewMetrics model.""" - metrics = EnumeratorOverviewMetrics( - all_submissions=100, - num_active_enumerators=5, - num_enumerators=10, - num_teams=3, - min_submissions=5, - max_submissions=25, - avg_submissions=15, - pct_active_enumerators="50%", - ) - assert metrics.all_submissions == 100 - assert metrics.num_active_enumerators == 5 - assert metrics.num_enumerators == 10 - assert metrics.num_teams == 3 - assert metrics.min_submissions == 5 - assert metrics.max_submissions == 25 - assert metrics.avg_submissions == 15 - assert metrics.pct_active_enumerators == "50%" - - -def test_enumerator_overview_metrics_validation(): - """Test EnumeratorOverviewMetrics validation for non-negative integers.""" - with pytest.raises(ValidationError): - EnumeratorOverviewMetrics( - all_submissions=-1, - num_active_enumerators=5, - num_enumerators=10, - num_teams=3, - min_submissions=5, - max_submissions=25, - avg_submissions=15, - pct_active_enumerators="50%", - ) - - -def test_enumerator_overview_metrics_num_teams_string(): - """Test EnumeratorOverviewMetrics allows string for num_teams.""" - metrics = EnumeratorOverviewMetrics( - all_submissions=100, - num_active_enumerators=5, - num_enumerators=10, - num_teams="n/a", - min_submissions=5, - max_submissions=25, - avg_submissions=15, - pct_active_enumerators="50%", - ) - assert metrics.num_teams == "n/a" - - -# ============================================ -# SETTINGS TESTS -# ============================================ - - -def test_load_default_enumerator_settings_valid(enumerator_settings_file): - """Test loading enumerator settings from valid file.""" - config = EnumeratorSettings( - survey_key="default_key", - survey_id="default_id", - ) - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = load_default_enumerator_settings(enumerator_settings_file, config) - - # Saved settings should override defaults - assert result.survey_key == "survey_id" - assert result.enumerator == "enumerator" - - -def test_load_default_enumerator_settings_missing_file(): - """Test loading enumerator settings when file doesn't exist.""" - config = EnumeratorSettings( - survey_key="default_key", - survey_id="default_id", - enumerator="default_enum", - ) - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = load_default_enumerator_settings("nonexistent.json", config) - - # Should return default config when file doesn't exist - assert result.survey_key == "default_key" - assert result.enumerator == "default_enum" - - -@patch("datasure.checks.enumerator.st") -def test_trigger_success_message(mock_st): - """Test _trigger_success_message function.""" - mock_st.session_state = {} - - _trigger_success_message("test_button") - assert mock_st.session_state["test_button"] is True - - -def test_create_enum_data_on_settings_with_consent_and_outcome(): - """Test _create_enum_data_on_settings with consent and outcome values.""" - data = pl.DataFrame( - { - "survey_id": ["S001", "S002", "S003"], - "consent": ["yes", "no", "yes"], - "outcome": ["completed", "incomplete", "completed"], - } - ) - - config = ConsentOutcomeSettings( - consent="consent", - consent_vals=["yes"], - outcome="outcome", - outcome_vals=["completed"], - ) - - with patch("datasure.checks.enumerator.duckdb_save_table") as mock_save: - _create_enum_data_on_settings("test_project", data, config) - - # Verify function was called - mock_save.assert_called_once() - saved_data = mock_save.call_args[0][1] - - # Check that consent and outcome columns were created - assert "consent_granted_agg_col" in saved_data.columns - assert "completed_survey_agg_col" in saved_data.columns - assert saved_data["consent_granted_agg_col"].to_list() == [1, 0, 1] - assert saved_data["completed_survey_agg_col"].to_list() == [1, 0, 1] - - -def test_create_enum_data_on_settings_without_consent(): - """Test _create_enum_data_on_settings without consent values.""" - data = pl.DataFrame( - { - "survey_id": ["S001", "S002"], - "outcome": ["completed", "completed"], - } - ) - - config = ConsentOutcomeSettings( - consent=None, - consent_vals=None, - outcome="outcome", - outcome_vals=["completed"], - ) - - with patch("datasure.checks.enumerator.duckdb_save_table") as mock_save: - _create_enum_data_on_settings("test_project", data, config) - - saved_data = mock_save.call_args[0][1] - # Consent should default to 1 - assert saved_data["consent_granted_agg_col"].to_list() == [1, 1] - - -def test_create_enum_data_on_settings_without_outcome(): - """Test _create_enum_data_on_settings without outcome values.""" - data = pl.DataFrame( - { - "survey_id": ["S001", "S002"], - "consent": ["yes", "yes"], - } - ) - - config = ConsentOutcomeSettings( - consent="consent", - consent_vals=["yes"], - outcome=None, - outcome_vals=None, - ) - - with patch("datasure.checks.enumerator.duckdb_save_table") as mock_save: - _create_enum_data_on_settings("test_project", data, config) - - saved_data = mock_save.call_args[0][1] - # Outcome should default to 1 - assert saved_data["completed_survey_agg_col"].to_list() == [1, 1] - - -# ============================================ -# COMPUTE_ENUMERATOR_OVERVIEW TESTS -# ============================================ - - -def test_compute_enumerator_overview_basic(sample_enumerator_data): - """Test basic enumerator overview computation.""" - result = compute_enumerator_overview( - sample_enumerator_data, "submission_date", "enumerator", "team" - ) - - assert result.all_submissions == 6 - assert result.num_enumerators == 3 - assert result.num_teams == 2 - assert result.num_active_enumerators >= 0 - assert result.min_submissions > 0 - assert result.max_submissions > 0 - assert result.avg_submissions > 0 - - -def test_compute_enumerator_overview_without_team(sample_enumerator_data): - """Test enumerator overview without team column.""" - result = compute_enumerator_overview( - sample_enumerator_data, "submission_date", "enumerator", None - ) - - assert result.all_submissions == 6 - assert result.num_enumerators == 3 - assert result.num_teams == "n/a" - - -def test_compute_enumerator_overview_empty_data(): - """Test enumerator overview with empty data.""" - empty_data = pl.DataFrame( - schema={"submission_date": pl.Date, "enumerator": pl.Utf8} - ) - - with pytest.raises(ValueError, match="Input data is empty"): - compute_enumerator_overview(empty_data, "submission_date", "enumerator", None) - - -def test_compute_enumerator_overview_active_enumerators(): - """Test active enumerators calculation.""" - today = date.today() - data = pl.DataFrame( - { - "submission_date": [ - today - timedelta(days=1), - today - timedelta(days=10), - today, - ], - "enumerator": ["E1", "E2", "E1"], - "team": ["T1", "T1", "T1"], - } - ) - - result = compute_enumerator_overview(data, "submission_date", "enumerator", "team") - - # Only E1 should be active (has submissions in past 7 days) - assert result.num_active_enumerators == 1 - assert result.num_enumerators == 2 - - -# ============================================ -# COMPUTE_ENUMERATOR_MISSING_TABLE TESTS -# ============================================ - - -def test_compute_enumerator_missing_table_empty_config(sample_enumerator_data): - """Test missing table with empty missing codes config.""" - empty_config = pl.DataFrame() - - result = compute_enumerator_missing_table( - sample_enumerator_data, empty_config, ["enumerator"] - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "% Null values" in result.columns - - -def test_compute_enumerator_missing_table_with_config( - sample_enumerator_data, sample_missing_codes_config -): - """Test missing table with missing codes config.""" - result = compute_enumerator_missing_table( - sample_enumerator_data, sample_missing_codes_config, ["enumerator"] - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_missing_table_with_team(sample_enumerator_data): - """Test missing table with team grouping.""" - empty_config = pl.DataFrame() - - result = compute_enumerator_missing_table( - sample_enumerator_data, empty_config, ["enumerator", "team"] - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "team" in result.columns - - -# ============================================ -# COMPUTE_ENUMERATOR_SUMMARY TESTS -# ============================================ - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_basic(mock_load_missing, sample_enumerator_data): - """Test basic enumerator summary computation.""" - mock_load_missing.return_value = pl.DataFrame() - - result = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "formversion", - "duration", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "team" in result.columns - assert "# submissions" in result.columns - assert "first submission" in result.columns - assert "last submission" in result.columns - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_without_team( - mock_load_missing, sample_enumerator_data -): - """Test enumerator summary without team.""" - mock_load_missing.return_value = pl.DataFrame() - - result = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "formversion", - "duration", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "team" not in result.columns - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_without_duration( - mock_load_missing, sample_enumerator_data -): - """Test enumerator summary without duration.""" - mock_load_missing.return_value = pl.DataFrame() - - result = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "formversion", - None, - ) - - assert not result.is_empty() - assert "min duration" not in result.columns - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_without_formversion( - mock_load_missing, sample_enumerator_data -): - """Test enumerator summary without formversion.""" - mock_load_missing.return_value = pl.DataFrame() - - result = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - None, - "duration", - ) - - assert not result.is_empty() - assert "# form versions" not in result.columns - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_with_consent( - mock_load_missing, sample_enumerator_data -): - """Test enumerator summary with consent column.""" - mock_load_missing.return_value = pl.DataFrame() - - result = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "formversion", - "duration", - ) - - assert not result.is_empty() - assert "% consent" in result.columns - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_with_outcome( - mock_load_missing, sample_enumerator_data -): - """Test enumerator summary with outcome column.""" - mock_load_missing.return_value = pl.DataFrame() - - result = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "formversion", - "duration", - ) - - assert not result.is_empty() - assert "% completed survey" in result.columns - - -# ============================================ -# COMPUTE_ENUMERATOR_PRODUCTIVITY TESTS -# ============================================ - - -def test_compute_enumerator_productivity_daily(sample_enumerator_data): - """Test productivity computation with daily period.""" - result = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Daily", "SUN" - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_productivity_weekly(sample_enumerator_data): - """Test productivity computation with weekly period.""" - result = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Weekly", "MON" - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_productivity_monthly(sample_enumerator_data): - """Test productivity computation with monthly period.""" - result = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Monthly", "SUN" - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_productivity_legacy_period_names(sample_enumerator_data): - """Test productivity with legacy period names.""" - # Test "Day" -> "Daily" - result = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Day", "SUN" - ) - assert not result.is_empty() - - # Test "Week" -> "Weekly" - result = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Week", "MON" - ) - assert not result.is_empty() - - # Test "Month" -> "Monthly" - result = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Month", "SUN" - ) - assert not result.is_empty() - - -def test_compute_enumerator_productivity_with_team(sample_enumerator_data): - """Test productivity with team grouping.""" - result = compute_enumerator_productivity( - sample_enumerator_data, - "submission_date", - ["enumerator", "team"], - "Daily", - "SUN", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "team" in result.columns - - -def test_compute_enumerator_productivity_different_weekstarts(sample_enumerator_data): - """Test productivity with different week start days.""" - for weekstart in ["SUN", "MON", "TUE", "WED", "THU", "FRI", "SAT"]: - result = compute_enumerator_productivity( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "Weekly", - weekstart, - ) - assert not result.is_empty() - - -# ============================================ -# COMPUTE_ENUMERATOR_STATISTICS TESTS -# ============================================ - - -def test_compute_enumerator_statistics_basic(sample_enumerator_data): - """Test basic statistics computation.""" - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = compute_enumerator_statistics( - sample_enumerator_data, - ["enumerator"], - ["age", "income"], - ["count", "mean"], - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "age_count" in result.columns - assert "age_mean" in result.columns - assert "income_count" in result.columns - assert "income_mean" in result.columns - - -def test_compute_enumerator_statistics_all_stats(sample_enumerator_data): - """Test statistics with all stat types.""" - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = compute_enumerator_statistics( - sample_enumerator_data, - ["enumerator"], - ["age"], - [ - "count", - "min", - "mean", - "median", - "max", - "std", - "25th percentile", - "75th percentile", - ], - ) - - assert not result.is_empty() - assert "age_count" in result.columns - assert "age_min" in result.columns - assert "age_mean" in result.columns - assert "age_median" in result.columns - assert "age_max" in result.columns - assert "age_std" in result.columns - assert "age_25th percentile" in result.columns - assert "age_75th percentile" in result.columns - - -def test_compute_enumerator_statistics_with_team(sample_enumerator_data): - """Test statistics with team grouping.""" - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = compute_enumerator_statistics( - sample_enumerator_data, - ["enumerator", "team"], - ["age"], - ["mean"], - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "team" in result.columns - - -def test_compute_enumerator_statistics_multiple_columns(sample_enumerator_data): - """Test statistics with multiple columns.""" - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = compute_enumerator_statistics( - sample_enumerator_data, - ["enumerator"], - ["age", "income", "duration"], - ["mean", "median"], - ) - - assert not result.is_empty() - assert len([col for col in result.columns if "_mean" in col]) == 3 - assert len([col for col in result.columns if "_median" in col]) == 3 - - -# ============================================ -# COMPUTE_ENUMERATOR_STATISTICS_OVERTIME TESTS -# ============================================ - - -def test_compute_enumerator_statistics_overtime_daily(sample_enumerator_data): - """Test statistics overtime with daily period.""" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Daily", - "SUN", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_statistics_overtime_weekly(sample_enumerator_data): - """Test statistics overtime with weekly period.""" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Weekly", - "MON", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_statistics_overtime_monthly(sample_enumerator_data): - """Test statistics overtime with monthly period.""" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Monthly", - "SUN", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - - -def test_compute_enumerator_statistics_overtime_missing_stat(sample_enumerator_data): - """Test statistics overtime with missing statistic.""" - # Add some null values - data_with_nulls = sample_enumerator_data.clone() - data_with_nulls = data_with_nulls.with_columns( - pl.when(pl.col("enumerator") == "E1") - .then(None) - .otherwise(pl.col("age")) - .alias("age") - ) - - result = compute_enumerator_statistics_overtime( - data_with_nulls, - "submission_date", - ["enumerator"], - "age", - "missing", - "Daily", - "SUN", - ) - - assert not result.is_empty() - - -def test_compute_enumerator_statistics_overtime_percentile_stats( - sample_enumerator_data, -): - """Test statistics overtime with percentile statistics.""" - result_25th = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "25th percentile", - "Daily", - "SUN", - ) - assert not result_25th.is_empty() - - result_75th = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "75th percentile", - "Daily", - "SUN", - ) - assert not result_75th.is_empty() - - -def test_compute_enumerator_statistics_overtime_legacy_periods(sample_enumerator_data): - """Test statistics overtime with legacy period names.""" - # Test "Day" -> "Daily" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Day", - "SUN", - ) - assert not result.is_empty() - - # Test "Week" -> "Weekly" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Week", - "MON", - ) - assert not result.is_empty() - - # Test "Month" -> "Monthly" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Month", - "SUN", - ) - assert not result.is_empty() - - -def test_compute_enumerator_statistics_overtime_with_team(sample_enumerator_data): - """Test statistics overtime with team grouping.""" - result = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator", "team"], - "age", - "mean", - "Daily", - "SUN", - ) - - assert not result.is_empty() - assert "enumerator" in result.columns - assert "team" in result.columns - - -# ============================================ -# HELPER FUNCTIONS TESTS -# ============================================ - - -def test_get_numeric_columns(): - """Test _get_numeric_columns helper function.""" - data = pl.DataFrame( - { - "age": [25, 30, 35], - "income": [50000, 60000, 55000], - "name": ["Alice", "Bob", "Charlie"], - "is_active": [True, False, True], - } - ) - - result = _get_numeric_columns(data) - assert "age" in result - assert "income" in result - assert "name" not in result - assert "is_active" not in result - - -def test_get_numeric_columns_with_exclude(): - """Test _get_numeric_columns with exclude list.""" - data = pl.DataFrame( - { - "age": [25, 30, 35], - "income": [50000, 60000, 55000], - "duration": [3600, 4200, 3800], - } - ) - - result = _get_numeric_columns(data, exclude_cols=["duration"]) - assert "age" in result - assert "income" in result - assert "duration" not in result - - -def test_get_numeric_columns_empty_dataframe(): - """Test _get_numeric_columns with empty DataFrame.""" - data = pl.DataFrame() - result = _get_numeric_columns(data) - assert result == [] - - -def test_get_numeric_columns_no_numeric(): - """Test _get_numeric_columns with no numeric columns.""" - data = pl.DataFrame( - { - "name": ["Alice", "Bob"], - "city": ["NYC", "LA"], - } - ) - result = _get_numeric_columns(data) - assert result == [] - - -# ============================================ -# EDGE CASES AND ERROR HANDLING TESTS -# ============================================ - - -def test_edge_case_single_enumerator(): - """Test handling of single enumerator.""" - data = pl.DataFrame( - { - "submission_date": [date.today(), date.today() - timedelta(days=1)], - "enumerator": ["E1", "E1"], - "age": [25, 30], - } - ) - - result = compute_enumerator_overview(data, "submission_date", "enumerator", None) - assert result.num_enumerators == 1 - assert result.all_submissions == 2 - - -def test_edge_case_all_null_values(): - """Test handling of all null values in column.""" - data = pl.DataFrame( - { - "submission_date": [date.today(), date.today()], - "enumerator": ["E1", "E2"], - "age": [None, None], - } - ) - - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - result = compute_enumerator_statistics( - data, - ["enumerator"], - ["age"], - ["mean"], - ) - - assert not result.is_empty() - # Mean of nulls should be null - assert result["age_mean"].is_null().any() - - -def test_edge_case_single_submission(): - """Test handling of single submission.""" - data = pl.DataFrame( - { - "submission_date": [date.today()], - "enumerator": ["E1"], - "team": ["T1"], - "age": [25], - } - ) - - result = compute_enumerator_overview(data, "submission_date", "enumerator", "team") - assert result.all_submissions == 1 - assert result.num_enumerators == 1 - - -def test_edge_case_date_ranges(): - """Test handling of various date ranges.""" - today = date.today() - data = pl.DataFrame( - { - "submission_date": [ - today, - today - timedelta(days=365), - today - timedelta(days=1000), - ], - "enumerator": ["E1", "E2", "E3"], - } - ) - - result = compute_enumerator_productivity( - data, "submission_date", ["enumerator"], "Monthly", "SUN" - ) - assert not result.is_empty() - - -# ============================================ -# INTEGRATION TESTS -# ============================================ - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_full_enumerator_workflow(mock_load_missing, sample_enumerator_data): - """Test complete enumerator workflow from overview to statistics.""" - mock_load_missing.return_value = pl.DataFrame() - - # Step 1: Compute overview - overview = compute_enumerator_overview( - sample_enumerator_data, "submission_date", "enumerator", "team" - ) - assert overview.num_enumerators == 3 - - # Step 2: Compute summary - summary = compute_enumerator_summary( - "test_project", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "formversion", - "duration", - ) - assert len(summary) > 0 - - # Step 3: Compute productivity - productivity = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], "Daily", "SUN" - ) - assert not productivity.is_empty() - - # Step 4: Compute statistics - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - statistics = compute_enumerator_statistics( - sample_enumerator_data, - ["enumerator"], - ["age", "income"], - ["mean", "median"], - ) - assert not statistics.is_empty() - - # Step 5: Compute statistics overtime - stats_overtime = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "Daily", - "SUN", - ) - assert not stats_overtime.is_empty() - - -def test_enumerator_workflow_with_missing_data(): - """Test enumerator workflow with missing data.""" - data = pl.DataFrame( - { - "submission_date": [date.today(), date.today() - timedelta(days=1)], - "enumerator": ["E1", "E2"], - "age": [25, None], - "income": [50000, 60000], - } - ) - - # Overview should work with missing data - overview = compute_enumerator_overview(data, "submission_date", "enumerator", None) - assert overview.all_submissions == 2 - - # Statistics should handle missing values - with patch("streamlit.cache_data", lambda ttl: lambda f: f): - statistics = compute_enumerator_statistics( - data, - ["enumerator"], - ["age", "income"], - ["count", "mean"], - ) - assert not statistics.is_empty() - - -def test_enumerator_workflow_different_time_periods(sample_enumerator_data): - """Test enumerator workflow with different time periods.""" - periods = ["Daily", "Weekly", "Monthly"] - - for period in periods: - # Test productivity - productivity = compute_enumerator_productivity( - sample_enumerator_data, "submission_date", ["enumerator"], period, "MON" - ) - assert not productivity.is_empty() - - # Test statistics overtime - stats_overtime = compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - period, - "MON", - ) - assert not stats_overtime.is_empty() - - -@patch("datasure.checks.enumerator.duckdb_save_table") -def test_consent_outcome_integration(mock_save, sample_enumerator_data): - """Test consent and outcome settings integration.""" - # Create consent and outcome settings - config = ConsentOutcomeSettings( - consent="consent", - consent_vals=["yes"], - outcome="outcome", - outcome_vals=["completed"], - ) - - # Add consent and outcome columns - data = sample_enumerator_data.with_columns( - [ - pl.lit("yes").alias("consent"), - pl.lit("completed").alias("outcome"), - ] - ) - - # Create enum data with settings - _create_enum_data_on_settings("test_project", data, config) - - # Verify the data was saved - assert mock_save.called - - -# ============================================ -# BRANCH MISS TESTS (COMPUTE FUNCTIONS) -# ============================================ - - -@patch("datasure.checks.enumerator.load_missing_codes_from_db") -def test_compute_enumerator_summary_without_consent_outcome_cols(mock_load_missing): - """compute_enumerator_summary skips consent/outcome when cols are absent.""" - mock_load_missing.return_value = pl.DataFrame() - data = pl.DataFrame( - { - "submission_date": [date.today()], - "enumerator": ["E1"], - } - ) - result = compute_enumerator_summary( - "proj", data, "submission_date", "enumerator", None, None, None - ) - assert not result.is_empty() - assert "% consent" not in result.columns - assert "% completed survey" not in result.columns - - -def test_compute_enumerator_productivity_unknown_period(sample_enumerator_data): - """compute_enumerator_productivity falls through period branches on unknown.""" - with pytest.raises(ColumnNotFoundError): - compute_enumerator_productivity( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "UnknownPeriod", - "SUN", - ) - - -def test_compute_enumerator_statistics_overtime_unknown_period(sample_enumerator_data): - """compute_enumerator_statistics_overtime falls through period branches.""" - with pytest.raises(ColumnNotFoundError): - compute_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - ["enumerator"], - "age", - "mean", - "UnknownPeriod", - "SUN", - ) - - -# ============================================ -# PATCHED_ENUM UI FUNCTION TESTS -# ============================================ - - -def test_render_enumerator_overview_no_date_enum(patched_enum): - """_render_enumerator_overview_metrics shows info when date/enum is None.""" - _render_enumerator_overview_metrics(pl.DataFrame(), None, None, None) - patched_enum.info.assert_called() - - -def test_render_enumerator_overview_with_data(patched_enum, sample_enumerator_data): - """_render_enumerator_overview_metrics renders metrics with valid data.""" - _render_enumerator_overview_metrics( - sample_enumerator_data, "submission_date", "enumerator", "team" - ) - patched_enum.columns.assert_called() - - -def test_render_enumerator_overview_no_team(patched_enum, sample_enumerator_data): - """_render_enumerator_overview_metrics renders without team column.""" - _render_enumerator_overview_metrics( - sample_enumerator_data, "submission_date", "enumerator", None - ) - patched_enum.columns.assert_called() - - -def test_render_time_period_selector_default(patched_enum): - """_render_time_period_selector returns Day when pills returns None.""" - result = _render_time_period_selector("settings.json") - assert result == "Day" - patched_enum.pills.assert_called() - - -def test_render_time_period_selector_week(patched_enum): - """_render_time_period_selector returns Week when pills returns Week.""" - patched_enum.pills.return_value = "Week" - result = _render_time_period_selector("settings.json") - assert result == "Week" - - -def test_render_weekday_selector(patched_enum): - """_render_weekday_selector returns offset code for the selected weekday.""" - patched_enum.selectbox.return_value = "Monday" - result = _render_weekday_selector("settings.json") - assert result == "SUN" - - -def test_load_statistics_settings_default(patched_enum): - """_load_statistics_settings returns default StatisticsSettings.""" - result = _load_statistics_settings("settings.json") - assert result.stats == ["count", "mean"] - - -def test_render_column_selector(patched_enum): - """_render_column_selector returns list from multiselect.""" - result = _render_column_selector(["age", "income"], None, "settings.json") - assert isinstance(result, list) - patched_enum.multiselect.assert_called() - - -def test_render_statistics_selector(patched_enum): - """_render_statistics_selector returns list from multiselect.""" - result = _render_statistics_selector(["count", "mean"], "settings.json") - assert isinstance(result, list) - patched_enum.multiselect.assert_called() - - -def test_render_enumerator_statistics_no_enum(patched_enum): - """_render_enumerator_statistics shows info when enumerator is None.""" - _render_enumerator_statistics(pl.DataFrame(), None, None, "settings.json") - patched_enum.info.assert_called() - - -def test_load_statistics_overtime_settings_default(patched_enum): - """_load_statistics_overtime_settings returns default settings.""" - result = _load_statistics_overtime_settings("settings.json") - assert result.stat == "count" - - -def test_render_period_selector_overtime_default(patched_enum): - """_render_period_selector_overtime returns Day when pills returns None.""" - result = _render_period_selector_overtime("settings.json") - assert result == "Day" - patched_enum.pills.assert_called() - - -def test_render_period_selector_overtime_week(patched_enum): - """_render_period_selector_overtime returns Week when pills returns Week.""" - patched_enum.pills.return_value = "Week" - result = _render_period_selector_overtime("settings.json", default_period="Week") - assert result == "Week" - - -def test_render_weekday_selector_overtime(patched_enum): - """_render_weekday_selector_overtime returns offset code.""" - patched_enum.selectbox.return_value = "Tuesday" - result = _render_weekday_selector_overtime("Monday", "settings.json") - assert result == "MON" - - -def test_render_statistic_selector(patched_enum): - """_render_statistic_selector returns the selectbox value.""" - patched_enum.selectbox.return_value = "mean" - result = _render_statistic_selector("count", "settings.json") - assert result == "mean" - - -def test_render_column_selector_single_default_none(patched_enum): - """_render_column_selector_single returns None when selectbox returns None.""" - result = _render_column_selector_single(["age", "income"], None, "settings.json") - assert result is None - patched_enum.selectbox.assert_called() - - -def test_render_column_selector_single_with_value(patched_enum): - """_render_column_selector_single returns the selectbox value.""" - patched_enum.selectbox.return_value = "age" - result = _render_column_selector_single(["age", "income"], "age", "settings.json") - assert result == "age" - - -def test_render_enumerator_productivity_no_enum(patched_enum): - """_render_enumerator_productivity shows info when enum/date is None.""" - _render_enumerator_productivity(pl.DataFrame(), None, None, None, "settings.json") - patched_enum.info.assert_called() - - -def test_render_enumerator_statistics_overtime_no_enum(patched_enum): - """_render_enumerator_statistics_overtime shows info when enum/date is None.""" - _render_enumerator_statistics_overtime( - pl.DataFrame(), None, None, None, "settings.json" - ) - patched_enum.info.assert_called() - - -# ============================================ -# ENUM_BC UI FRAGMENT TESTS -# ============================================ - - -def test_enumerator_report_settings_basic(enum_bc, sample_enumerator_data): - """enumerator_report_settings returns EnumeratorSettings from UI.""" - enum_bc.st.selectbox.return_value = "survey_id" - categorical_cols = list(sample_enumerator_data.columns) - datetime_cols = ["submission_date"] - config = EnumeratorSettings(survey_id="survey_id") - result = enum_bc.enumerator_report_settings( - "proj_id", - "settings.json", - sample_enumerator_data, - config, - categorical_cols, - datetime_cols, - ) - assert result is not None - - -def test_render_consent_outcome_settings(enum_bc, sample_enumerator_data): - """_render_consent_outcome_settings renders consent/outcome selectors.""" - enum_bc.st.selectbox.return_value = "survey_id" - categorical_cols = list(sample_enumerator_data.columns) - enum_bc._render_consent_outcome_settings( - "proj_id", sample_enumerator_data, categorical_cols, "settings.json" - ) - enum_bc.st.button.assert_called() - - -def test_render_consent_outcome_settings_button_click(enum_bc, sample_enumerator_data): - """_render_consent_outcome_settings calls create when button clicked.""" - enum_bc.st.selectbox.return_value = "survey_id" - enum_bc.st.button.return_value = True - categorical_cols = list(sample_enumerator_data.columns) - enum_bc._render_consent_outcome_settings( - "proj_id", sample_enumerator_data, categorical_cols, "settings.json" - ) - enum_bc.duckdb_save_table.assert_called() - - -def test_render_enumerator_summary_table_no_date_enum(enum_bc): - """_render_enumerator_summary_table shows info when date/enum is None.""" - enum_bc._render_enumerator_summary_table( - "proj", pl.DataFrame(), None, None, None, None, None - ) - enum_bc.st.info.assert_called() - - -def test_render_enumerator_summary_table_with_data(enum_bc, sample_enumerator_data): - """_render_enumerator_summary_table renders table with valid data.""" - enum_bc._render_enumerator_summary_table( - "proj", - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "formversion", - "duration", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_summary_table_with_show_info( - enum_bc, sample_enumerator_data -): - """_render_enumerator_summary_table filters columns when show_info selected.""" - enum_bc.st.pills.return_value = ["submissions"] - enum_bc._render_enumerator_summary_table( - "proj", - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "formversion", - "duration", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_productivity_table_day(enum_bc, sample_enumerator_data): - """_render_enumerator_productivity_table renders with Day period.""" - enum_bc._render_enumerator_productivity_table( - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "settings.json", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_productivity_table_week(enum_bc, sample_enumerator_data): - """_render_enumerator_productivity_table renders weekday selector for Week.""" - enum_bc.st.pills.return_value = "Week" - enum_bc.st.selectbox.return_value = "Monday" - enum_bc._render_enumerator_productivity_table( - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "settings.json", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_statistics_table_no_enum(enum_bc): - """_render_enumerator_statistics_table shows info when enum is None.""" - enum_bc._render_enumerator_statistics_table(pl.DataFrame(), None, None, "s.json") - enum_bc.st.info.assert_called() - - -def test_render_enumerator_statistics_table_no_cols(enum_bc, sample_enumerator_data): - """_render_enumerator_statistics_table shows info when no cols selected.""" - enum_bc._render_enumerator_statistics_table( - sample_enumerator_data, "enumerator", None, "settings.json" - ) - enum_bc.st.info.assert_called() - - -def test_render_enumerator_statistics_table_with_cols(enum_bc, sample_enumerator_data): - """_render_enumerator_statistics_table renders table when cols selected.""" - enum_bc.st.multiselect.side_effect = [["age"], ["count", "mean"]] - enum_bc._render_enumerator_statistics_table( - sample_enumerator_data, "enumerator", "team", "settings.json" - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_statistics_with_enum(enum_bc, sample_enumerator_data): - """_render_enumerator_statistics calls the fragment table function.""" - enum_bc._render_enumerator_statistics( - sample_enumerator_data, "enumerator", None, "settings.json" - ) - enum_bc.st.info.assert_called() - - -def test_render_enumerator_statistics_overtime_table_no_statscol( - enum_bc, sample_enumerator_data -): - """_render_enumerator_statistics_overtime_table shows info when statscol None.""" - enum_bc._render_enumerator_statistics_overtime_table( - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "settings.json", - ) - enum_bc.st.info.assert_called() - - -def test_render_enumerator_statistics_overtime_table_with_col( - enum_bc, sample_enumerator_data -): - """_render_enumerator_statistics_overtime_table renders with valid col.""" - enum_bc.st.selectbox.side_effect = ["age", "count"] - enum_bc._render_enumerator_statistics_overtime_table( - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "settings.json", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_statistics_overtime_table_week_period( - enum_bc, sample_enumerator_data -): - """_render_enumerator_statistics_overtime_table handles Week period.""" - enum_bc.st.pills.return_value = "Week" - enum_bc.st.selectbox.side_effect = ["age", "count", "Monday"] - enum_bc._render_enumerator_statistics_overtime_table( - sample_enumerator_data, - "submission_date", - "enumerator", - "team", - "settings.json", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_productivity_with_valid_params( - enum_bc, sample_enumerator_data -): - """_render_enumerator_productivity calls the table fragment when params valid.""" - enum_bc._render_enumerator_productivity( - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "settings.json", - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_statistics_overtime_with_valid_params( - enum_bc, sample_enumerator_data -): - """_render_enumerator_statistics_overtime calls the table fragment.""" - enum_bc._render_enumerator_statistics_overtime( - sample_enumerator_data, - "submission_date", - "enumerator", - None, - "settings.json", - ) - enum_bc.st.info.assert_called() - - -def test_enumerator_report_empty_data(enum_bc): - """enumerator_report shows info and returns early when data is empty.""" - survey_cols = ColumnByType( - categorical_columns=["survey_id", "enumerator"], - datetime_columns=["submission_date"], - ) - enum_bc.enumerator_report( - "proj_id", - pl.DataFrame(), - "settings.json", - {"survey_id": "survey_id"}, - survey_cols, - ) - enum_bc.st.info.assert_called() - - -def test_enumerator_report_with_data(enum_bc, sample_enumerator_data): - """enumerator_report renders full report with valid data.""" - survey_cols = ColumnByType( - categorical_columns=list(sample_enumerator_data.columns), - datetime_columns=["submission_date"], - ) - enum_bc._render_consent_outcome_settings = MagicMock() - - def _selectbox_side_effect(label, options=None, **kwargs): - if options and "seconds" in options: - return "seconds" - return None - - enum_bc.st.selectbox.side_effect = _selectbox_side_effect - enum_bc.enumerator_report( - "proj_id", - sample_enumerator_data, - "settings.json", - {"survey_id": "survey_id"}, - survey_cols, - ) - enum_bc.st.title.assert_called() - - -def test_load_statistics_settings_fallback(patched_enum): - """_load_statistics_settings returns defaults when saved settings invalid.""" - with patch( - "datasure.checks.enumerator.load_check_settings", - return_value={"stats": ["not_a_real_stat"]}, - ): - result = _load_statistics_settings("settings.json") - assert result.stats == ["count", "mean"] - - -def test_load_statistics_overtime_settings_fallback(patched_enum): - """_load_statistics_overtime_settings returns defaults when settings invalid.""" - with patch( - "datasure.checks.enumerator.load_check_settings", - return_value={"period_overtime": "bad_period"}, - ): - result = _load_statistics_overtime_settings("settings.json") - assert result.stat == "count" - - -def test_render_enumerator_statistics_table_with_cols_no_team( - enum_bc, sample_enumerator_data -): - """_render_enumerator_statistics_table uses else column config when no team.""" - enum_bc.st.multiselect.side_effect = [["age"], ["count"]] - enum_bc._render_enumerator_statistics_table( - sample_enumerator_data, "enumerator", None, "settings.json" - ) - enum_bc.st.dataframe.assert_called() - - -def test_render_enumerator_statistics_overtime_table_early_return( - enum_bc, sample_enumerator_data -): - """_render_enumerator_statistics_overtime_table returns early when enum None.""" - enum_bc._render_enumerator_statistics_overtime_table( - sample_enumerator_data, "submission_date", None, None, "settings.json" - ) - enum_bc.st.dataframe.assert_not_called() - - -def test_enumerator_report_settings_success_flag(enum_bc, sample_enumerator_data): - """enumerator_report_settings shows success message when consent flag set.""" - enum_bc.st.selectbox.return_value = "survey_id" - enum_bc.st.session_state["st_apply_consent_outcome_enumerator"] = True - categorical_cols = list(sample_enumerator_data.columns) - config = EnumeratorSettings(survey_id="survey_id") - enum_bc.enumerator_report_settings( - "proj_id", - "settings.json", - sample_enumerator_data, - config, - categorical_cols, - ["submission_date"], - ) - enum_bc.st.success.assert_called()