diff --git a/CHANGELOG.md b/CHANGELOG.md index 5b95fac0..f6c8c19a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,43 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **Correction log**: New `source` column records which page made each entry; + existing logs load with `source = corrections_page`. A new `check_type` + column goes with the new `accept` action + (`CORRECTION_LOG_SCHEMA`, `ensure_log_columns` in the new Streamlit-free + `src/datasure/processing/correction_log.py`, shared with the replication + package, which now exports legacy logs with these columns) — #296 +- **Accept action**: `CorrectionProcessor.accept_value` records that a flagged + value (outliers, constraints, backchecks, duplicates, GPS) was reviewed and + is correct, with a required reason. An acceptance is rejected if the data + no longer holds the value being accepted. `get_active_acceptances` returns the + acceptances whose recorded value still matches the data (for GPS, both + latitude and longitude). Replay and the generated `4_corrections.do` skip + `accept` rows; `correction_log.csv` keeps them, and the README's correction + counts exclude them. Values compare by value, not text: missing matches + missing (None or NaN), numbers compare numerically, and every row with the + KEY must match — #296 +- **Atomic apply**: `CorrectionProcessor.apply_corrections` applies a list of + `CorrectionEntry` objects all or nothing; if the log save fails, the + corrected data is restored — #296 +- **Shared correction form**: `src/datasure/utils/correction_form.py` + (`render_correction_form`, `render_correction_inputs`, + `apply_correction_entries`) renders the action, new-value and reason inputs + for a prefilled KEY/column/current value, with namespaced widget keys. The + Correct Data page now uses it — #296 + +### Fixed + +- **Corrections cache**: `CorrectionProcessor`'s cached reads were keyed only on + `alias`, so two projects sharing an alias shared cached corrected data and + logs. The processor is now hashed by `project_id` — #296 +- **Apply button**: A new value of `0` no longer disables Apply. An empty + string still does; use "remove value" to blank a cell — #296 +- **Correction log schema**: Removing the last correction entry now leaves an + empty log with the full schema, including status columns — #296 + ## [1.1.0] - 2026-09-21 ### Added diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 1f343e58..ca4fabd0 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -35,13 +35,15 @@ src/datasure/ │ └── local.py # Local file import (csv/xlsx/xls/json/dta/parquet) ├── processing/ │ ├── prep.py # Data preparation operations (Polars) -│ └── corrections.py # Data correction application +│ ├── corrections.py # Corrections and accept entries (CorrectionProcessor) +│ └── correction_log.py # Correction log schema/backfill (no Streamlit) ├── replication/ # Stata/Python replication package export ├── models/ │ ├── schemas.py # Pydantic models │ └── enums.py # Prep action/method enums ├── utils/ # Shared utilities (DuckDB, cache, config, charts, │ # credentials, SurveyCTO API, UI helpers, ...) +│ └── correction_form.py # Shared correction form used across pages └── views/ # Streamlit pages (top-level page scripts) ├── start_view.py # Project selection/creation ├── import_view.py # Credentials + data import diff --git a/docs/USER_GUIDE.md b/docs/USER_GUIDE.md index caa2fc8a..cbb2c6e0 100644 --- a/docs/USER_GUIDE.md +++ b/docs/USER_GUIDE.md @@ -460,6 +460,18 @@ All corrections are tracked with: - Action type - Reason for correction - Timestamp +- Status of the last reapply, with the reason if it failed +- Source: the page that made the entry (`corrections_page` for this page) + +The log can also contain **accept** entries. An accept entry records that a +flagged value was reviewed and is correct. It never changes the data, and the +Check type column shows which check it applies to. It stays in effect only +while the value is unchanged. Accept entries are kept in `correction_log.csv` +in the replication package but are not part of the corrections script. You can +remove an accept entry with "Remove correction step" like any other entry. + +To blank a cell, use "remove value": "modify value" needs a non-empty new +value (`0` is valid). #### Verifying Corrections diff --git a/src/datasure/processing/correction_log.py b/src/datasure/processing/correction_log.py new file mode 100644 index 00000000..4389c183 --- /dev/null +++ b/src/datasure/processing/correction_log.py @@ -0,0 +1,63 @@ +"""Schema and vocabulary of the correction log (`corr_log_{alias}`). + +Kept free of Streamlit so the replication package can read logs with the +same schema and backfill rules as `CorrectionProcessor`. +""" + +import polars as pl + +CORRECTIONS_PAGE_SOURCE = "corrections_page" + +# Actions that change the data when a correction is applied or replayed. +MODIFY_VALUE_ACTION = "modify value" +REMOVE_VALUE_ACTION = "remove value" +REMOVE_ROW_ACTION = "remove row" +CORRECTION_ACTIONS = (MODIFY_VALUE_ACTION, REMOVE_VALUE_ACTION, REMOVE_ROW_ACTION) + +# An "accept" entry records that a flagged value was reviewed and is correct. +# It never changes the data; check pages use it to stop flagging the value. +ACCEPT_ACTION = "accept" +ACCEPT_CHECK_TYPES = ("outliers", "constraints", "backchecks", "duplicates", "gps") + +# Full schema of a persisted correction log (`corr_log_{alias}`), in column order. +CORRECTION_LOG_SCHEMA: dict[str, pl.DataType] = { + "date": pl.Datetime("us"), + "KEY": pl.String, + "ID": pl.String, + "action": pl.String, + "column": pl.String, + "current_value": pl.String, + "new_value": pl.String, + "reason": pl.String, + "status": pl.String, + "status_reason": pl.String, + "source": pl.String, + "check_type": pl.String, +} + +# Values given to columns that were added to the log after some logs were +# already persisted. Every legacy entry came from the Corrections page and +# was applied successfully when it was logged. +_LOG_BACKFILL_DEFAULTS: dict[str, str | None] = { + "status": "Successful", + "status_reason": None, + "source": CORRECTIONS_PAGE_SOURCE, + "check_type": None, +} + + +def ensure_log_columns(df: pl.DataFrame) -> pl.DataFrame: + """Backfill columns missing from logs persisted before those columns existed.""" + if df.width == 0: + return df + for column, default in _LOG_BACKFILL_DEFAULTS.items(): + if column not in df.columns: + df = df.with_columns( + pl.lit(default, dtype=CORRECTION_LOG_SCHEMA[column]).alias(column) + ) + return df + + +def empty_correction_log() -> pl.DataFrame: + """Return a correction log with no entries and the full schema.""" + return pl.DataFrame(schema=CORRECTION_LOG_SCHEMA) diff --git a/src/datasure/processing/corrections.py b/src/datasure/processing/corrections.py index a3b38a46..fb52c0fa 100644 --- a/src/datasure/processing/corrections.py +++ b/src/datasure/processing/corrections.py @@ -1,9 +1,24 @@ +import json +import math +from collections.abc import Sequence +from dataclasses import dataclass from datetime import datetime from typing import Any import polars as pl import streamlit as st +from datasure.processing.correction_log import ( + ACCEPT_ACTION, + ACCEPT_CHECK_TYPES, + CORRECTION_LOG_SCHEMA, + CORRECTIONS_PAGE_SOURCE, + MODIFY_VALUE_ACTION, + REMOVE_ROW_ACTION, + REMOVE_VALUE_ACTION, + empty_correction_log, + ensure_log_columns, +) from datasure.utils.duckdb_utils import duckdb_get_table, duckdb_save_table from datasure.utils.reapply_utils import ReapplyFailure @@ -26,24 +41,218 @@ def _describe_correction_row(row: dict[str, Any]) -> str: column = row["column"] new_value = row["new_value"] - if action == "modify value": + if action == MODIFY_VALUE_ACTION: return f"Modify {column} for key {key_value} to '{new_value}'" - if action == "remove value": + if action == REMOVE_VALUE_ACTION: return f"Remove {column} value for key {key_value}" - if action == "remove row": + if action == REMOVE_ROW_ACTION: return f"Remove entire row for key {key_value}" + if action == ACCEPT_ACTION: + target = column if column is not None else "coordinates" + return f"Accept {row['check_type']} flag on {target} for key {key_value}" return f"{action} for key {key_value}" -def _ensure_status_columns(df: pl.DataFrame) -> pl.DataFrame: - """Backfill status/status_reason columns for logs persisted before they existed.""" - if df.is_empty(): - return df - if "status" not in df.columns: - df = df.with_columns(pl.lit("Successful").alias("status")) - if "status_reason" not in df.columns: - df = df.with_columns(pl.lit(None, dtype=pl.String).alias("status_reason")) - return df +def _is_missing(value: Any) -> bool: + """Whether a value is missing: None, or NaN as pandas reports nulls.""" + return value is None or (isinstance(value, float) and math.isnan(value)) + + +def _encode_scalar(value: Any) -> str | None: + return None if _is_missing(value) else str(value) + + +def _encode_log_value(value: Any) -> str | None: + """Encode a value for the log's string-typed value columns. + + Missing values (None or NaN) are stored as null. GPS acceptances cover a + latitude/longitude pair, passed as a mapping of column name to value and + stored as JSON. Any other value is stored as a string. + """ + if isinstance(value, dict): + return json.dumps( + {col: _encode_scalar(v) for col, v in value.items()}, sort_keys=True + ) + return _encode_scalar(value) + + +def _values_match(actual: Any, recorded: str | None) -> bool: + """Whether a data value equals a value recorded in the log as a string. + + Missing matches missing. Numbers compare numerically, so an integer cell + holding 25 matches "25.0", which is how pandas reports an integer column + that has nulls. + """ + if _is_missing(actual) or recorded is None: + return _is_missing(actual) and recorded is None + if str(actual) == recorded: + return True + if isinstance(actual, int | float) and not isinstance(actual, bool): + try: + return float(actual) == float(recorded) + except ValueError: + return False + return False + + +def _accepted_values(row: dict[str, Any]) -> dict[str, str | None]: + """Return the column -> value snapshot an accept row was recorded against.""" + if row["column"] is not None: + return {row["column"]: row["current_value"]} + return json.loads(row["current_value"]) if row["current_value"] else {} + + +def _acceptance_mismatch( + data: pl.DataFrame, key_col: str, key_value: Any, accepted: dict[str, str | None] +) -> str | None: + """Explain why the data doesn't hold the accepted values, or None if it does. + + `accepted` maps each column to its value as encoded in the log. The log + stores KEY as text, so the key column is compared as text. If the KEY + appears on several rows, every row must hold the accepted values. + """ + if key_col not in data.columns: + return f"Key column '{key_col}' not found in data" + records = data.filter(pl.col(key_col).cast(pl.String) == str(key_value)) + if records.is_empty(): + return f"Key value '{key_value}' not found in data" + for column, value in accepted.items(): + if column not in records.columns: + return f"Column '{column}' not found in data" + if not all(_values_match(v, value) for v in records[column].to_list()): + return ( + f"The value of '{column}' for key '{key_value}' has changed since " + "it was flagged. Refresh the page and review it again." + ) + return None + + +def _acceptance_is_active( + data: pl.DataFrame, key_col: str, row: dict[str, Any] +) -> bool: + """Whether the data still holds the values an acceptance recorded.""" + return ( + _acceptance_mismatch(data, key_col, row["KEY"], _accepted_values(row)) is None + ) + + +def _check_acceptance_against_data( + data: pl.DataFrame, + key_col: str, + key_value: Any, + column: str | None, + current_value: Any, +) -> None: + """Raise ValueError unless the data holds the value being accepted. + + The value is encoded exactly as the log will store it, so an acceptance + that passes this check is active as soon as it is logged. + """ + accepted = _accepted_values( + {"column": column, "current_value": _encode_log_value(current_value)} + ) + mismatch = _acceptance_mismatch(data, key_col, key_value, accepted) + if mismatch: + raise ValueError(mismatch) + + +def _validate_acceptance( + check_type: str | None, column: str | None, current_value: Any +) -> None: + """Raise ValueError if an acceptance's check type, column and value don't fit.""" + if check_type not in ACCEPT_CHECK_TYPES: + raise ValueError( + f"Unknown check type '{check_type}'; expected one of " + f"{', '.join(ACCEPT_CHECK_TYPES)}" + ) + if check_type == "gps": + if ( + column is not None + or not isinstance(current_value, dict) + or len(current_value) != 2 + ): + raise ValueError( + "GPS acceptances take no column and a mapping of the " + "latitude and longitude columns to their values" + ) + elif not column: + raise ValueError(f"A column is required to accept a {check_type} value") + + +def _build_log_row( + key_value: str, + current_id: Any | None, + action: str, + column: str | None, + current_value: Any | None, + new_value: Any | None, + reason: str, + source: str, + check_type: str | None, +) -> dict[str, Any]: + """Build one correction-log row. + + A freshly logged entry has just been applied successfully (the apply + step raises before logging otherwise). + """ + return { + "date": datetime.now(), + "KEY": str(key_value), + "ID": str(current_id) if current_id is not None else None, + "action": str(action), + "column": str(column) if column is not None else None, + "current_value": _encode_log_value(current_value), + "new_value": _encode_log_value(new_value), + "reason": str(reason), + "status": "Successful", + "status_reason": None, + "source": str(source), + "check_type": check_type, + } + + +@dataclass(frozen=True) +class CorrectionEntry: + """One correction or acceptance, applied as part of `apply_corrections`. + + Attributes + ---------- + key_value : str + The KEY of the record + action : str + "modify value", "remove value", "remove row" or "accept" + reason : str + Why the entry is made. Required. + column : str | None + The column affected. None for "remove row" and GPS acceptances. + current_value : Any + The value before the change, or the value being accepted. For GPS + acceptances, a mapping of the latitude and longitude columns to their + values. + new_value : Any + The new value, for "modify value" + survey_id_value : Any + The Survey ID value for this KEY, recorded in the log's ID column + check_type : str | None + For "accept", the check whose flag is accepted + """ + + key_value: str + action: str + reason: str + column: str | None = None + current_value: Any = None + new_value: Any = None + survey_id_value: Any = None + check_type: str | None = None + + +# The cached methods below hash `self` by its project so that two projects +# sharing an alias never share cached data. The key is the class's qualified +# name because the class isn't defined yet when the decorators run. +_PROCESSOR_HASH_FUNCS = { + "datasure.processing.corrections.CorrectionProcessor": lambda p: p.project_id +} class CorrectionProcessor: @@ -59,8 +268,8 @@ def __init__(self, project_id: str) -> None: """ self.project_id = project_id - @st.cache_data(ttl=60, show_spinner=False) - def get_corrected_data(_self, alias: str) -> pl.DataFrame: + @st.cache_data(ttl=60, show_spinner=False, hash_funcs=_PROCESSOR_HASH_FUNCS) + def get_corrected_data(self, alias: str) -> pl.DataFrame: """Get corrected data for a given alias. If no corrected data exists, initializes from prepped data. @@ -76,7 +285,7 @@ def get_corrected_data(_self, alias: str) -> pl.DataFrame: The corrected data """ corrected_data = duckdb_get_table( - project_id=_self.project_id, + project_id=self.project_id, alias=alias, db_name="corrected", ) @@ -84,12 +293,12 @@ def get_corrected_data(_self, alias: str) -> pl.DataFrame: if corrected_data.is_empty(): # Initialize from prepped data prepped_data = duckdb_get_table( - project_id=_self.project_id, + project_id=self.project_id, alias=alias, db_name="prep", ) if not prepped_data.is_empty(): - _self.save_corrected_data(alias, prepped_data) + self.save_corrected_data(alias, prepped_data) return prepped_data return corrected_data @@ -114,8 +323,8 @@ def save_corrected_data(self, alias: str, data: pl.DataFrame) -> None: self.get_corrected_data.clear() self.get_data_summary.clear() - @st.cache_data(ttl=30, show_spinner=False) - def get_correction_log(_self, alias: str) -> pl.DataFrame: + @st.cache_data(ttl=30, show_spinner=False, hash_funcs=_PROCESSOR_HASH_FUNCS) + def get_correction_log(self, alias: str) -> pl.DataFrame: """Get correction log for a given alias. Parameters @@ -128,10 +337,12 @@ def get_correction_log(_self, alias: str) -> pl.DataFrame: pl.DataFrame The correction log """ - return duckdb_get_table( - project_id=_self.project_id, - alias=f"corr_log_{alias}", - db_name="logs", + return ensure_log_columns( + duckdb_get_table( + project_id=self.project_id, + alias=f"corr_log_{alias}", + db_name="logs", + ) ) def add_correction_entry( @@ -144,6 +355,8 @@ def add_correction_entry( current_value: Any | None, new_value: Any | None, reason: str, + source: str = CORRECTIONS_PAGE_SOURCE, + check_type: str | None = None, ) -> None: """Add a new correction entry to the log. @@ -166,51 +379,51 @@ def add_correction_entry( The new value reason : str The reason for correction + source : str + The page that produced the entry, e.g. "corrections_page" or a + check page such as "outliers" + check_type : str | None + For "accept" entries, the check whose flag was accepted """ - current_log = self.get_correction_log(alias) - - # Create new entry DataFrame with proper schema. A freshly added - # correction has just been applied successfully (apply_correction - # would have raised before reaching this point otherwise). - new_entry_data = { - "date": [datetime.now()], - "KEY": [str(key_value)], - "ID": [str(current_id) if current_id is not None else None], - "action": [str(action)], - "column": [str(column) if column is not None else None], - "current_value": [ - str(current_value) if current_value is not None else None + self._append_log_rows( + alias, + [ + _build_log_row( + key_value=key_value, + current_id=current_id, + action=action, + column=column, + current_value=current_value, + new_value=new_value, + reason=reason, + source=source, + check_type=check_type, + ) ], - "new_value": [str(new_value) if new_value is not None else None], - "reason": [str(reason)], - "status": ["Successful"], - "status_reason": [None], - } - new_entry_df = pl.DataFrame(new_entry_data).with_columns( - pl.col("status_reason").cast(pl.String) ) + def _append_log_rows(self, alias: str, rows: list[dict[str, Any]]) -> None: + """Append rows to the correction log in a single save. + + Parameters + ---------- + alias : str + The data alias/table name + rows : list[dict[str, Any]] + Log rows built by `_build_log_row` + """ + current_log = self.get_correction_log(alias) + new_rows = pl.DataFrame(rows, schema=CORRECTION_LOG_SCHEMA) + if current_log.is_empty(): - # If no existing log, use the new entry schema - updated_log = new_entry_df + updated_log = new_rows else: - # Ensure schema compatibility before concatenating - # Cast columns to match the new entry schema - aligned_current_log = _ensure_status_columns(current_log).with_columns( - [ - pl.col("date").cast(pl.Datetime("us")), - pl.col("KEY").cast(pl.String), - pl.col("ID").cast(pl.String), - pl.col("action").cast(pl.String), - pl.col("column").cast(pl.String), - pl.col("current_value").cast(pl.String), - pl.col("new_value").cast(pl.String), - pl.col("reason").cast(pl.String), - pl.col("status").cast(pl.String), - pl.col("status_reason").cast(pl.String), - ] + # Align column order and types with the new rows before concatenating + aligned_current_log = current_log.select( + pl.col(name).cast(dtype) + for name, dtype in CORRECTION_LOG_SCHEMA.items() ) - updated_log = pl.concat([aligned_current_log, new_entry_df]) + updated_log = pl.concat([aligned_current_log, new_rows]) duckdb_save_table( project_id=self.project_id, @@ -218,10 +431,118 @@ def add_correction_entry( alias=f"corr_log_{alias}", db_name="logs", ) - # Clear correction log cache so the new entry shows immediately + # Clear correction log cache so the new entries show immediately self.get_correction_log.clear() self.get_correction_summary.clear() + def accept_value( + self, + alias: str, + key_col: str, + key_value: str, + check_type: str, + column: str | None, + current_value: Any, + reason: str, + survey_id_value: Any | None = None, + source: str | None = None, + ) -> None: + """Record that a flagged value was reviewed and is correct. + + The entry never changes the data. It stays active only while the + data still holds `current_value` (see `get_active_acceptances`). + + Parameters + ---------- + alias : str + The data alias/table name + key_col : str + The key column name + key_value : str + The KEY of the accepted record + check_type : str + The check whose flag is accepted, one of `ACCEPT_CHECK_TYPES` + column : str | None + The accepted column. None for GPS, which accepts a coordinate pair. + current_value : Any + The value being accepted. For GPS, a mapping of the latitude and + longitude column names to their values. + reason : str + Why the value is correct. Required. + survey_id_value : Any | None + The Survey ID value for this KEY, if a Survey ID column is + configured, recorded in the log's ID column + source : str | None + The page that produced the entry. Defaults to `check_type`. + + Raises + ------ + ValueError + If the check type is unknown, the reason is blank, the column and + value do not fit the check type, or the corrected data does not + hold `current_value` for the KEY (for example, because the value + changed after it was flagged). + """ + _validate_acceptance(check_type, column, current_value) + if not reason or not reason.strip(): + raise ValueError("A reason is required to accept a value") + _check_acceptance_against_data( + self.get_corrected_data(alias), key_col, key_value, column, current_value + ) + + self.add_correction_entry( + alias=alias, + key_value=key_value, + current_id=survey_id_value, + action=ACCEPT_ACTION, + column=column, + current_value=current_value, + new_value=None, + reason=reason, + source=source or check_type, + check_type=check_type, + ) + + def get_active_acceptances( + self, alias: str, check_type: str, key_col: str + ) -> pl.DataFrame: + """Return the acceptances for a check that still apply to the data. + + An acceptance is active only while the corrected data still holds the + value recorded when it was accepted. For GPS, both the latitude and + the longitude must still match. + + Parameters + ---------- + alias : str + The data alias/table name + check_type : str + The check to return acceptances for + key_col : str + The Survey KEY column name + + Returns + ------- + pl.DataFrame + The active "accept" rows from the correction log, in log order + """ + log = self.get_correction_log(alias) + if log.width == 0: + return empty_correction_log() + + acceptances = log.filter( + (pl.col("action") == ACCEPT_ACTION) & (pl.col("check_type") == check_type) + ) + if acceptances.is_empty(): + return acceptances + + data = self.get_corrected_data(alias) + is_active = [ + _acceptance_is_active(data, key_col, row) + for row in acceptances.iter_rows(named=True) + ] + return acceptances.filter(pl.Series(is_active, dtype=pl.Boolean)) + def apply_correction( self, alias: str, @@ -263,18 +584,14 @@ def apply_correction( pl.DataFrame The corrected data """ - corrected_data = self.get_corrected_data(alias) - - if action == "modify value" and column and new_value is not None: - corrected_data = self._apply_modify_value( - corrected_data, key_col, key_value, column, new_value - ) - elif action == "remove value" and column: - corrected_data = self._apply_remove_value( - corrected_data, key_col, key_value, column - ) - elif action == "remove row": - corrected_data = self._apply_remove_row(corrected_data, key_col, key_value) + corrected_data = self._apply_action( + self.get_corrected_data(alias), + key_col, + key_value, + action, + column, + new_value, + ) self.save_corrected_data(alias, corrected_data) @@ -293,6 +610,127 @@ def apply_correction( return corrected_data + def apply_corrections( + self, + alias: str, + key_col: str, + entries: Sequence[CorrectionEntry], + source: str = CORRECTIONS_PAGE_SOURCE, + ) -> pl.DataFrame: + """Apply several corrections and acceptances as one all-or-nothing step. + + Entries are validated and applied in order against the result of the + entries before them. If any entry is invalid, nothing is saved and + nothing is logged. + + Parameters + ---------- + alias : str + The data alias/table name + key_col : str + The key column name + entries : Sequence[CorrectionEntry] + The corrections and acceptances to apply, in order + source : str + The page that produced the entries, recorded on every log row + + Returns + ------- + pl.DataFrame + The corrected data + + Raises + ------ + ValueError + If any entry is invalid; the message names the entry's problem. + Storage errors are re-raised after the corrected data is restored. + """ + original_data = self.get_corrected_data(alias) + corrected_data = original_data + for entry in entries: + corrected_data = self._apply_entry(corrected_data, key_col, entry) + + log_rows = [ + _build_log_row( + key_value=entry.key_value, + current_id=entry.survey_id_value, + action=entry.action, + column=entry.column, + current_value=entry.current_value, + new_value=entry.new_value, + reason=entry.reason, + source=source, + check_type=entry.check_type, + ) + for entry in entries + ] + + self.save_corrected_data(alias, corrected_data) + try: + self._append_log_rows(alias, log_rows) + except Exception: + # Keep data and log in step: undo the data change, then re-raise. + self.save_corrected_data(alias, original_data) + raise + return corrected_data + + def _apply_entry( + self, data: pl.DataFrame, key_col: str, entry: CorrectionEntry + ) -> pl.DataFrame: + """Validate one `CorrectionEntry` against `data` and apply it. + + Raises + ------ + ValueError + If the entry is invalid for `data`. + """ + if not entry.reason or not entry.reason.strip(): + raise ValueError( + f"A reason is required for {entry.action} on {entry.key_value}" + ) + + if entry.action == ACCEPT_ACTION: + _validate_acceptance(entry.check_type, entry.column, entry.current_value) + _check_acceptance_against_data( + data, key_col, entry.key_value, entry.column, entry.current_value + ) + return data + + if entry.action not in ( + MODIFY_VALUE_ACTION, + REMOVE_VALUE_ACTION, + REMOVE_ROW_ACTION, + ): + raise ValueError(f"Unknown correction action '{entry.action}'") + + is_valid, error_msg = self.validate_correction_input( + data, key_col, entry.key_value, entry.action, entry.column, entry.new_value + ) + if not is_valid: + raise ValueError(error_msg) + + return self._apply_action( + data, key_col, entry.key_value, entry.action, entry.column, entry.new_value + ) + + def _apply_action( + self, + data: pl.DataFrame, + key_col: str, + key_value: str, + action: str, + column: str | None, + new_value: Any | None, + ) -> pl.DataFrame: + """Apply one correction action to `data`; unknown actions leave it as is.""" + if action == MODIFY_VALUE_ACTION and column and new_value is not None: + return self._apply_modify_value(data, key_col, key_value, column, new_value) + if action == REMOVE_VALUE_ACTION and column: + return self._apply_remove_value(data, key_col, key_value, column) + if action == REMOVE_ROW_ACTION: + return self._apply_remove_row(data, key_col, key_value) + return data + def _apply_modify_value( self, data: pl.DataFrame, @@ -407,8 +845,8 @@ def _apply_remove_row( """ return data.filter(pl.col(key_col) != key_value) - @st.cache_data(ttl=60, show_spinner=False) - def get_data_summary(_self, data: pl.DataFrame) -> dict[str, Any]: + @st.cache_data(ttl=60, show_spinner=False, hash_funcs=_PROCESSOR_HASH_FUNCS) + def get_data_summary(self, data: pl.DataFrame) -> dict[str, Any]: """Get summary statistics for the data. Parameters @@ -481,13 +919,13 @@ def validate_correction_input( if key_value not in data[key_col].to_list(): return False, f"Key value '{key_value}' not found in data" - if action in ["modify value", "remove value"]: + if action in [MODIFY_VALUE_ACTION, REMOVE_VALUE_ACTION]: if not column: return False, "Column must be specified for modify/remove value actions" if column not in data.columns: return False, f"Column '{column}' not found in data" - if action == "modify value" and new_value is None: + if action == MODIFY_VALUE_ACTION and new_value is None: return False, "New value must be provided for modify value action" return True, "" @@ -522,18 +960,7 @@ def remove_correction_entry( # Remove the correction entry at the specified index if correction_index == 0 and correction_log.height == 1: # If removing the only entry, create an empty DataFrame with proper schema - updated_log = pl.DataFrame( - { - "date": pl.Series([], dtype=pl.Datetime("us")), - "KEY": pl.Series([], dtype=pl.String), - "ID": pl.Series([], dtype=pl.String), - "action": pl.Series([], dtype=pl.String), - "column": pl.Series([], dtype=pl.String), - "current_value": pl.Series([], dtype=pl.String), - "new_value": pl.Series([], dtype=pl.String), - "reason": pl.Series([], dtype=pl.String), - } - ) + updated_log = empty_correction_log() else: # Build list of parts to concatenate parts = [] @@ -542,22 +969,8 @@ def remove_correction_entry( if correction_index < len(correction_log) - 1: parts.append(correction_log[correction_index + 1 :]) - if parts: - updated_log = pl.concat(parts) - else: - # Should not happen given the conditions above, but handle it - updated_log = pl.DataFrame( - { - "date": pl.Series([], dtype=pl.Datetime("us")), - "KEY": pl.Series([], dtype=pl.String), - "ID": pl.Series([], dtype=pl.String), - "action": pl.Series([], dtype=pl.String), - "column": pl.Series([], dtype=pl.String), - "current_value": pl.Series([], dtype=pl.String), - "new_value": pl.Series([], dtype=pl.String), - "reason": pl.Series([], dtype=pl.String), - } - ) + # parts is never empty given the conditions above, but handle it + updated_log = pl.concat(parts) if parts else empty_correction_log() # Save the updated log duckdb_save_table( @@ -654,7 +1067,8 @@ def _apply_correction_row( ) -> tuple[pl.DataFrame, str | None]: """Apply one correction-log row to data. - Returns `data` unchanged if the row's key can't be located, or if + Returns `data` unchanged for "accept" rows, if the row's key can't be + located, or if applying the correction fails (the underlying data may have changed since the correction was logged). @@ -672,17 +1086,21 @@ def _apply_correction_row( message describing why the correction was skipped, or None on success. """ + action = row["action"] + if action == ACCEPT_ACTION: + # Acceptances record a decision; they never change the data. + return data, None + key_value = row["KEY"] key_col = self._find_key_column(data, key_value) if not key_col: return data, f"Key '{key_value}' not found in current data" - action = row["action"] column = row["column"] recorded_value = row["current_value"] new_value = row["new_value"] - if action in ("modify value", "remove value") and column: + if action in (MODIFY_VALUE_ACTION, REMOVE_VALUE_ACTION) and column: if column not in data.columns: return data, f"Column '{column}' no longer available in the data" @@ -693,16 +1111,16 @@ def _apply_correction_row( return data, mismatch try: - if action == "modify value" and column and new_value is not None: + if action == MODIFY_VALUE_ACTION and column and new_value is not None: return ( self._apply_modify_value( data, key_col, key_value, column, new_value ), None, ) - if action == "remove value" and column: + if action == REMOVE_VALUE_ACTION and column: return self._apply_remove_value(data, key_col, key_value, column), None - if action == "remove row": + if action == REMOVE_ROW_ACTION: return self._apply_remove_row(data, key_col, key_value), None except Exception as e: # Skip corrections that fail (data may have changed) @@ -778,8 +1196,8 @@ def _find_key_column(data: pl.DataFrame, key_value: str) -> str | None: continue return None - @st.cache_data(ttl=30, show_spinner=False) - def get_correction_summary(_self, alias: str) -> list[dict[str, Any]]: + @st.cache_data(ttl=30, show_spinner=False, hash_funcs=_PROCESSOR_HASH_FUNCS) + def get_correction_summary(self, alias: str) -> list[dict[str, Any]]: """Get a summary of all correction entries for display. Parameters @@ -792,7 +1210,7 @@ def get_correction_summary(_self, alias: str) -> list[dict[str, Any]]: list[dict[str, Any]] List of correction summaries with index, description, and details """ - correction_log = _self.get_correction_log(alias) + correction_log = self.get_correction_log(alias) if correction_log.is_empty(): return [] @@ -813,6 +1231,7 @@ def get_correction_summary(_self, alias: str) -> list[dict[str, Any]]: "index": index, "action_index": f"{index} - {action} - {description}", "action": action, + "check_type": row["check_type"], "description": description, "key_value": key_value, "column": column, diff --git a/src/datasure/replication/package_builder.py b/src/datasure/replication/package_builder.py index 93273099..57a382f5 100644 --- a/src/datasure/replication/package_builder.py +++ b/src/datasure/replication/package_builder.py @@ -10,6 +10,11 @@ import polars as pl +from datasure.processing.correction_log import ( + ACCEPT_ACTION, + CORRECTION_LOG_SCHEMA, + ensure_log_columns, +) from datasure.replication.codebook import generate_codebook from datasure.replication.prep_script_generator import ( generate_prepare_data_script, @@ -73,7 +78,9 @@ def _load_prep_log(project_id: str, alias: str) -> pl.DataFrame: def _load_correction_log(project_id: str, alias: str) -> pl.DataFrame: try: - return duckdb_get_table(project_id, f"corr_log_{alias}", "logs") + return ensure_log_columns( + duckdb_get_table(project_id, f"corr_log_{alias}", "logs") + ) except Exception: logger.warning( "Correction log for %s not found; returning empty DataFrame", alias @@ -81,7 +88,15 @@ def _load_correction_log(project_id: str, alias: str) -> pl.DataFrame: return pl.DataFrame() +def _applied_corrections(correction_log: pl.DataFrame) -> pl.DataFrame: + """Return the log rows that change data, dropping "accept" review records.""" + if "action" not in correction_log.columns: + return correction_log + return correction_log.filter(pl.col("action") != ACCEPT_ACTION) + + def _action_summary(correction_log: pl.DataFrame) -> dict[str, int]: + correction_log = _applied_corrections(correction_log) if correction_log.is_empty(): return {} counts = ( @@ -249,7 +264,7 @@ def _step(msg: str) -> None: project_name=project_name, survey_name=survey_name, datasure_version=datasure_version, - correction_count=correction_log.height, + correction_count=_applied_corrections(correction_log).height, prep_count=prep_log.height, raw_rows=raw_df.height, prepped_rows=prepped_df.height if not prepped_df.is_empty() else 0, @@ -269,7 +284,7 @@ def _step(msg: str) -> None: correction_log_csv = ( correction_log.write_csv() if not correction_log.is_empty() - else "date,KEY,ID,action,column,current_value,new_value,reason\n" + else ",".join(CORRECTION_LOG_SCHEMA) + "\n" ) prep_log_csv = ( prep_log.with_columns( diff --git a/src/datasure/replication/script_generators.py b/src/datasure/replication/script_generators.py index fb4fae33..c334bbc9 100644 --- a/src/datasure/replication/script_generators.py +++ b/src/datasure/replication/script_generators.py @@ -6,6 +6,13 @@ import polars as pl +from datasure.processing.correction_log import ( + ACCEPT_ACTION, + MODIFY_VALUE_ACTION, + REMOVE_ROW_ACTION, + REMOVE_VALUE_ACTION, +) + SCRIPT_EXT = "do" _C = "*" # Stata comment character _LOG_CLOSE = "cap log close" @@ -69,17 +76,17 @@ def _emit_stata( escaped_key = _escape(key_val) escaped_val = _escape(new_val) if new_val is not None else None - if action == "modify value" and col and new_val is not None: + if action == MODIFY_VALUE_ACTION and col and new_val is not None: if _is_numeric(new_val): stmt = f'replace {col} = {new_val} if {key_col} == "{escaped_key}"' else: stmt = f'replace {col} = "{escaped_val}" if {key_col} == "{escaped_key}"' - elif action == "remove value" and col: + elif action == REMOVE_VALUE_ACTION and col: if _is_numeric(new_val): stmt = f'replace {col} = . if {key_col} == "{escaped_key}"' else: stmt = f'replace {col} = "" if {key_col} == "{escaped_key}"' - elif action == "remove row": + elif action == REMOVE_ROW_ACTION: stmt = f'drop if {key_col} == "{escaped_key}"' else: return [] @@ -298,6 +305,7 @@ def generate_corrections_script( ---------- correction_log : pl.DataFrame Correction log with columns: action, KEY, column, new_value, reason. + "accept" rows are skipped. key_col : str The survey key column name. project_name : str @@ -323,6 +331,11 @@ def generate_corrections_script( + "\n" ) + # "accept" rows record that a flagged value was reviewed and is correct; + # they never change the data, so they have no Stata equivalent. + if "action" in correction_log.columns: + correction_log = correction_log.filter(pl.col("action") != ACCEPT_ACTION) + if correction_log.is_empty(): return ( header diff --git a/src/datasure/utils/correction_form.py b/src/datasure/utils/correction_form.py new file mode 100644 index 00000000..d9d53e49 --- /dev/null +++ b/src/datasure/utils/correction_form.py @@ -0,0 +1,592 @@ +"""Shared correction form: the action, new-value and reason inputs. + +The Corrections page and the check pages all correct or accept values with +the same inputs. `render_correction_form` renders them for one KEY (and +optionally a prefilled column and current value) with an Apply button; +`render_correction_inputs` renders them without a button so a caller can +collect several entries and save them together with +`apply_correction_entries`, which applies them all or none. + +Every widget key is suffixed with `key_namespace`, so the form can appear on +several pages and tabs at once. + +Streamlit is imported inside each function rather than at module level, as +in `ui_utils`: view tests swap ``sys.modules["streamlit"]`` for a mock, and +resolving it per call honors that swap whatever the import order. +""" + +import logging +from collections.abc import Callable, Sequence +from datetime import date, datetime +from typing import Any + +import polars as pl +from pydantic import BaseModel, Field + +from datasure.processing.correction_log import ( + ACCEPT_ACTION, + CORRECTION_ACTIONS, + MODIFY_VALUE_ACTION, + REMOVE_ROW_ACTION, + REMOVE_VALUE_ACTION, +) +from datasure.processing.corrections import CorrectionEntry, CorrectionProcessor + +logger = logging.getLogger(__name__) + + +class CorrectionFormState(BaseModel): + """State management for correction form inputs.""" + + key_value: str = Field(..., description="The selected key value for correction") + action: str = Field(..., description="The correction action type") + column: str | None = Field(None, description="The column to modify (if applicable)") + current_value: Any | None = Field( + None, description="The current value (if applicable)" + ) + new_value: Any | None = Field(None, description="The new value (if applicable)") + validation_error: str | None = Field( + None, description="Validation error message (if any)" + ) + reason: str = Field("", description="The reason entered for the entry") + check_type: str | None = Field( + None, description="For 'accept', the check whose flag is accepted" + ) + survey_id_value: Any | None = Field( + None, description="The Survey ID value for the KEY, if configured" + ) + + def to_entry(self) -> CorrectionEntry: + """Return the entry to pass to `CorrectionProcessor.apply_corrections`.""" + return CorrectionEntry( + key_value=self.key_value, + action=self.action, + reason=self.reason, + column=self.column, + current_value=self.current_value, + new_value=self.new_value, + survey_id_value=self.survey_id_value, + check_type=self.check_type, + ) + + +def get_current_value( + data: pl.DataFrame, key_col: str, key_value: str, column: str +) -> Any: + """ + Retrieve the current value for a specific key and column. + + Parameters + ---------- + data : pl.DataFrame + The dataset to query. + key_col : str + The name of the key column. + key_value : str + The key value to filter by. + column : str + The column to retrieve the value from. + + Returns + ------- + Any + The current value, or None if not found. + """ + try: + return data.filter(pl.col(key_col) == key_value).select(column)[0, 0] + except (pl.exceptions.PolarsError, IndexError): + logger.debug( + "No value for %s=%r in column %r", key_col, key_value, column, exc_info=True + ) + return None + + +def parse_date_value(value: Any) -> date | None: + """ + Parse a datetime value to a date object. + + Parameters + ---------- + value : Any + The value to parse (can be string or datetime). + + Returns + ------- + date | None + Parsed date or None if parsing fails. + """ + if not value: + return None + + try: + if isinstance(value, str): + return datetime.fromisoformat(value).date() + return value.date() + except (ValueError, TypeError, AttributeError): + logger.debug("Could not parse %r as a date", value, exc_info=True) + return None + + +def validate_numeric_input(value: str, dtype: pl.DataType) -> tuple[bool, str | None]: + """ + Validate numeric input based on column data type. + + Parameters + ---------- + value : str + The input value to validate. + dtype : pl.DataType + The expected data type. + + Returns + ------- + tuple[bool, str | None] + A tuple of (is_valid, error_message). + """ + if dtype in [pl.Int64, pl.Int32, pl.Float64, pl.Float32]: + try: + float(value) + return True, None # noqa: TRY300 + except ValueError: + return False, "New value must be a number." + return True, None + + +def should_enable_apply_button(action: str, reason: str, new_value: Any = None) -> bool: + """ + Determine if the apply button should be enabled. + + Parameters + ---------- + action : str + The correction action type. + reason : str + The reason for correction. + new_value : Any, optional + The new value (required for modify action). Falsy values such as + "0" or 0 are valid. An empty string is treated as missing: to blank + a cell, use the "remove value" action. + + Returns + ------- + bool + True if apply button should be enabled. + """ + if not reason: + return False + + if action == MODIFY_VALUE_ACTION: + return new_value is not None and new_value != "" + return action in [REMOVE_VALUE_ACTION, REMOVE_ROW_ACTION, ACCEPT_ACTION] + + +def render_value_input_widget( + column: str, + col_dtype: pl.DataType, + current_value: Any, + key_namespace: str | int, +) -> tuple[Any, str | None]: + """ + Render appropriate input widget based on column data type. + + Parameters + ---------- + column : str + The column name being modified. + col_dtype : pl.DataType + The column data type. + current_value : Any + The current value in the column. + key_namespace : str | int + Suffix for unique widget keys. + + Returns + ------- + tuple[Any, str | None] + A tuple of (new_value, error_message). + """ + import streamlit as st + + if col_dtype == pl.Datetime: + current_date = parse_date_value(current_value) + new_value = st.date_input( + label="New Value", + key=f"correction_new_value_{key_namespace}", + value=current_date, + help="Select a date for the new value.", + ) + return new_value, None + + # Text input for other types + new_value = st.text_input( + label="New Value", + key=f"correction_new_value_{key_namespace}", + placeholder="Enter new value", + ) + + if new_value: + is_valid, error_msg = validate_numeric_input(new_value, col_dtype) + if not is_valid: + return None, error_msg + + return new_value, None + + +def _render_column_selector( + data: pl.DataFrame, + key_col: str, + key_value: str, + key_namespace: str | int, + column: str | None = None, + current_value: Any = None, +) -> tuple[str | None, Any]: + """Render the column selector (unless prefilled) and the current value. + + Returns the column and its current value. A prefilled `current_value` is + shown as is; otherwise it is looked up in `data`. + """ + import streamlit as st + + if column is None: + column = st.selectbox( + label="Select Column to Modify", + options=data.columns, + key=f"correction_col_to_modify_{key_namespace}", + ) + if not column: + return None, None + else: + st.write(f"**Column:** {column}") + + if current_value is None: + current_value = get_current_value(data, key_col, key_value, column) + + st.write(f"**Current Value:** {current_value}") + + return column, current_value + + +def _render_modify_value_action( + data: pl.DataFrame, + key_col: str, + key_value: str, + key_namespace: str | int, + column: str | None = None, + current_value: Any = None, +) -> CorrectionFormState: + """Render the 'modify value' inputs and return the collected state.""" + import streamlit as st + + column, current_value = _render_column_selector( + data, key_col, key_value, key_namespace, column, current_value + ) + + if not column: + return CorrectionFormState( + key_value=key_value, action=MODIFY_VALUE_ACTION, column=None + ) + + col_dtype = data.schema[column] + new_value, validation_error = render_value_input_widget( + column, col_dtype, current_value, key_namespace + ) + + if validation_error: + st.error(validation_error) + + return CorrectionFormState( + key_value=key_value, + action=MODIFY_VALUE_ACTION, + column=column, + current_value=current_value, + new_value=new_value, + validation_error=validation_error, + ) + + +def _render_remove_value_action( + data: pl.DataFrame, + key_col: str, + key_value: str, + key_namespace: str | int, + column: str | None = None, + current_value: Any = None, +) -> CorrectionFormState: + """Render the 'remove value' inputs and return the collected state.""" + column, current_value = _render_column_selector( + data, key_col, key_value, key_namespace, column, current_value + ) + + return CorrectionFormState( + key_value=key_value, + action=REMOVE_VALUE_ACTION, + column=column, + current_value=current_value, + ) + + +def _render_remove_row_action(key_value: str) -> CorrectionFormState: + """Render the 'remove row' warning and return the collected state.""" + import streamlit as st + + st.warning("This will remove the row with the selected key value from the dataset.") + + return CorrectionFormState(key_value=key_value, action=REMOVE_ROW_ACTION) + + +def _render_accept_action( + data: pl.DataFrame, + key_col: str, + key_value: str, + key_namespace: str | int, + column: str | None = None, + current_value: Any = None, + check_type: str | None = None, +) -> CorrectionFormState: + """Render the 'accept' inputs and return the collected state. + + A GPS acceptance has no column: its prefilled current value maps the + latitude and longitude columns to their values. + """ + import streamlit as st + + if column is None and isinstance(current_value, dict): + st.write(f"**Current Value:** {current_value}") + else: + column, current_value = _render_column_selector( + data, key_col, key_value, key_namespace, column, current_value + ) + + return CorrectionFormState( + key_value=key_value, + action=ACCEPT_ACTION, + column=column, + current_value=current_value, + check_type=check_type, + ) + + +def _render_action_ui( + action: str, + data: pl.DataFrame, + key_col: str, + key_value: str, + key_namespace: str | int, + column: str | None = None, + current_value: Any = None, + check_type: str | None = None, +) -> CorrectionFormState: + """Render the inputs for `action` and return the collected state.""" + if action == MODIFY_VALUE_ACTION: + return _render_modify_value_action( + data, key_col, key_value, key_namespace, column, current_value + ) + + if action == REMOVE_VALUE_ACTION: + return _render_remove_value_action( + data, key_col, key_value, key_namespace, column, current_value + ) + + if action == ACCEPT_ACTION: + return _render_accept_action( + data, key_col, key_value, key_namespace, column, current_value, check_type + ) + + if action == REMOVE_ROW_ACTION: + return _render_remove_row_action(key_value) + + # Never fall back to a destructive action for an unrecognized value. + raise ValueError(f"Unsupported correction action: {action!r}") + + +def render_correction_inputs( + data: pl.DataFrame, + key_col: str, + key_value: str, + *, + key_namespace: str | int, + actions: Sequence[str] = CORRECTION_ACTIONS, + column: str | None = None, + current_value: Any = None, + check_type: str | None = None, + survey_id_value: Any = None, +) -> CorrectionFormState: + """ + Render the action, new-value and reason inputs for one KEY. + + Parameters + ---------- + data : pl.DataFrame + The corrected dataset, used for column options, types and lookups. + key_col : str + The Survey KEY column name. + key_value : str + The KEY the entry applies to. + key_namespace : str | int + Suffix for unique widget keys, so several forms can render at once. + actions : Sequence[str] + The actions to offer. May include "accept". + column : str | None + A prefilled column. If None, the user picks one when the action + needs it (except a GPS acceptance, which has no column). + current_value : Any + A prefilled current value. If None, it is looked up in `data`. For a + GPS acceptance, a mapping of the latitude and longitude columns to + their values. + check_type : str | None + The check an "accept" entry accepts. Required if `actions` includes + "accept". + survey_id_value : Any + The Survey ID value for the KEY, recorded with the entry. + + Returns + ------- + CorrectionFormState + The collected inputs; `to_entry()` turns them into a CorrectionEntry. + """ + import streamlit as st + + action = st.selectbox( + label="Select Action", + options=actions, + key=f"correction_action_{key_namespace}", + ) + + state = _render_action_ui( + action, + data, + key_col, + key_value, + key_namespace, + column, + current_value, + check_type, + ) + + reason = st.text_input( + label="Reason for Correction", + key=f"correction_reason_{key_namespace}", + placeholder="Enter reason for correction", + ) + + return state.model_copy( + update={"reason": reason, "survey_id_value": survey_id_value} + ) + + +def apply_correction_entries( + correction_processor: CorrectionProcessor, + alias: str, + key_col: str, + entries: Sequence[CorrectionEntry], + source: str, +) -> bool: + """ + Apply several entries all-or-nothing and report the outcome. + + Parameters + ---------- + correction_processor : CorrectionProcessor + The correction processor instance. + alias : str + The data alias/table name. + key_col : str + The Survey KEY column name. + entries : Sequence[CorrectionEntry] + The corrections and acceptances to apply, in order. + source : str + The page making the entries, recorded in the log. + + Returns + ------- + bool + True if every entry was applied; False if none were. + """ + import streamlit as st + + try: + correction_processor.apply_corrections( + alias=alias, key_col=key_col, entries=list(entries), source=source + ) + except Exception as e: + # UI boundary: report any failure to the user instead of crashing the page. + logger.exception( + "Failed to apply %d correction entries to %s", len(entries), alias + ) + st.error(f"Error applying correction: {e!s}") + return False + + st.success( + "Correction applied successfully!" + if len(entries) == 1 + else f"{len(entries)} corrections applied successfully!" + ) + return True + + +def render_correction_form( + correction_processor: CorrectionProcessor, + alias: str, + key_col: str, + data: pl.DataFrame, + key_value: str, + *, + key_namespace: str | int, + source: str, + actions: Sequence[str] = CORRECTION_ACTIONS, + column: str | None = None, + current_value: Any = None, + check_type: str | None = None, + survey_id_value: Any = None, + on_apply: Callable[[CorrectionFormState], None] | None = None, +) -> None: + """ + Render the correction inputs for one KEY with an Apply button. + + Parameters are as for `render_correction_inputs`, plus: + + Parameters + ---------- + correction_processor : CorrectionProcessor + The correction processor instance. + alias : str + The data alias/table name. + source : str + The page making the entry, recorded in the log. + on_apply : Callable[[CorrectionFormState], None] | None + Called with the collected state when Apply is clicked, in place of + the default save through `apply_correction_entries`. + """ + import streamlit as st + + state = render_correction_inputs( + data, + key_col, + key_value, + key_namespace=key_namespace, + actions=actions, + column=column, + current_value=current_value, + check_type=check_type, + survey_id_value=survey_id_value, + ) + + apply_enabled = should_enable_apply_button( + state.action, state.reason, state.new_value + ) + + if not st.button( + label="Apply", + key=f"correction_apply_{key_namespace}", + width="stretch", + disabled=not apply_enabled or bool(state.validation_error), + type="primary", + ): + return + + if on_apply is not None: + on_apply(state) + elif apply_correction_entries( + correction_processor, alias, key_col, [state.to_entry()], source + ): + st.rerun() diff --git a/src/datasure/views/correction_view.py b/src/datasure/views/correction_view.py index 16f958b0..75ef3919 100644 --- a/src/datasure/views/correction_view.py +++ b/src/datasure/views/correction_view.py @@ -6,14 +6,24 @@ removed or modified as needed. """ -from datetime import datetime from typing import Any import polars as pl import streamlit as st from pydantic import BaseModel, Field +from datasure.processing.correction_log import ( + CORRECTIONS_PAGE_SOURCE, + ensure_log_columns, +) from datasure.processing.corrections import CorrectionProcessor +from datasure.utils.correction_form import ( + get_current_value, + render_correction_form, +) +from datasure.utils.correction_form import ( + render_value_input_widget as render_shared_value_input_widget, +) from datasure.utils.duckdb_utils import duckdb_get_table from datasure.utils.navigations_utils import ( add_demo_navigation, @@ -30,9 +40,6 @@ section_header, ) -# DEFINE CONSTANTS FOR CORRECTION -CORRECTION_ACTIONS = ("modify value", "remove value", "remove row") - class TabConfig(BaseModel): """Configuration for a correction tab.""" @@ -45,19 +52,32 @@ class TabConfig(BaseModel): ) -class CorrectionFormState(BaseModel): - """State management for correction form inputs.""" +def render_value_input_widget( + column: str, + col_dtype: pl.DataType, + current_value: Any, + tab_index: int, +) -> tuple[Any, str | None]: + """ + Render the new-value input for a column on a Corrections page tab. - key_value: str = Field(..., description="The selected key value for correction") - action: str = Field(..., description="The correction action type") - column: str | None = Field(None, description="The column to modify (if applicable)") - current_value: Any | None = Field( - None, description="The current value (if applicable)" - ) - new_value: Any | None = Field(None, description="The new value (if applicable)") - validation_error: str | None = Field( - None, description="Validation error message (if any)" - ) + Parameters + ---------- + column : str + The column name being modified. + col_dtype : pl.DataType + The column data type. + current_value : Any + The current value in the column. + tab_index : int + The tab index for unique widget keys. + + Returns + ------- + tuple[Any, str | None] + A tuple of (new_value, error_message). + """ + return render_shared_value_input_widget(column, col_dtype, current_value, tab_index) def load_hfc_config(project_id: str) -> tuple[pl.DataFrame, list[str]]: @@ -102,113 +122,6 @@ def get_key_options(data: pl.DataFrame, key_col: str) -> list: return data.select(key_col).unique(maintain_order=True).to_series().to_list() -def get_current_value( - data: pl.DataFrame, key_col: str, key_value: str, column: str -) -> Any: - ( - """ - Retrieve the current value for a specific key and column. - - Parameters - ---------- - data : pl.DataFrame - The dataset to query. - key_col : str - The name of the key column. - key_value : str - The key value to filter by. - column : str - The column to retrieve the value from. - - Returns - ------- - Any - The current value, or None if not found. - """ - "" - ) - try: - return data.filter(pl.col(key_col) == key_value).select(column)[0, 0] - except Exception: - return None - - -def parse_date_value(value: Any) -> datetime.date: - """ - Parse a datetime value to a date object. - - Parameters - ---------- - value : Any - The value to parse (can be string or datetime). - - Returns - ------- - datetime.date | None - Parsed date or None if parsing fails. - """ - if not value: - return None - - try: - if isinstance(value, str): - return datetime.fromisoformat(value).date() - return value.date() - except Exception: - return None - - -def validate_numeric_input(value: str, dtype: pl.DataType) -> tuple[bool, str | None]: - """ - Validate numeric input based on column data type. - - Parameters - ---------- - value : str - The input value to validate. - dtype : pl.DataType - The expected data type. - - Returns - ------- - tuple[bool, str | None] - A tuple of (is_valid, error_message). - """ - if dtype in [pl.Int64, pl.Int32, pl.Float64, pl.Float32]: - try: - float(value) - return True, None # noqa: TRY300 - except ValueError: - return False, "New value must be a number." - return True, None - - -def should_enable_apply_button(action: str, reason: str, new_value: Any = None) -> bool: - """ - Determine if the apply button should be enabled. - - Parameters - ---------- - action : str - The correction action type. - reason : str - The reason for correction. - new_value : Any, optional - The new value (required for modify action). - - Returns - ------- - bool - True if apply button should be enabled. - """ - if not reason: - return False - - if action == "modify value": - return bool(new_value) - return action in ["remove value", "remove row"] - - def load_tab_config(project_id: str, tab_index: int) -> TabConfig | None: """ Load configuration for a specific correction tab. @@ -281,205 +194,6 @@ def validate_prerequisites(project_id: str | None) -> tuple[pl.DataFrame, list[s return hfc_config_logs, hfc_pages -def render_value_input_widget( - column: str, - col_dtype: pl.DataType, - current_value: Any, - tab_index: int, -) -> tuple[Any, str | None]: - """ - Render appropriate input widget based on column data type. - - Parameters - ---------- - column : str - The column name being modified. - col_dtype : pl.DataType - The column data type. - current_value : Any - The current value in the column. - tab_index : int - The tab index for unique widget keys. - - Returns - ------- - tuple[Any, str | None] - A tuple of (new_value, error_message). - """ - if col_dtype == pl.Datetime: - current_date = parse_date_value(current_value) - new_value = st.date_input( - label="New Value", - key=f"correction_new_value_{tab_index}", - value=current_date, - help="Select a date for the new value.", - ) - return new_value, None - - # Text input for other types - new_value = st.text_input( - label="New Value", - key=f"correction_new_value_{tab_index}", - placeholder="Enter new value", - ) - - if new_value: - is_valid, error_msg = validate_numeric_input(new_value, col_dtype) - if not is_valid: - return None, error_msg - - return new_value, None - - -def _render_column_selector( - corrected_data: pl.DataFrame, - key_col: str, - key_value: str, - tab_index: int, -) -> tuple[str | None, Any]: - """ - Render column selector and display current value. - - Parameters - ---------- - corrected_data : pl.DataFrame - The corrected dataset. - key_col : str - The key column name. - key_value : str - The selected key value. - tab_index : int - The tab index for unique widget keys. - - Returns - ------- - tuple[str | None, Any] - Column name and current value. - """ - column = st.selectbox( - label="Select Column to Modify", - options=corrected_data.columns, - key=f"correction_col_to_modify_{tab_index}", - ) - - if not column: - return None, None - - current_value = get_current_value(corrected_data, key_col, key_value, column) - - st.write(f"**Current Value:** {current_value}") - - return column, current_value - - -def _render_modify_value_action( - corrected_data: pl.DataFrame, - key_col: str, - key_value: str, - tab_index: int, -) -> CorrectionFormState: - """ - Render UI elements for 'modify value' action. - - Parameters - ---------- - corrected_data : pl.DataFrame - The corrected dataset. - key_col : str - The key column name. - key_value : str - The selected key value. - tab_index : int - The tab index for unique widget keys. - - Returns - ------- - CorrectionFormState - Form state with column, values, and validation errors. - """ - column, current_value = _render_column_selector( - corrected_data, key_col, key_value, tab_index - ) - - if not column: - return CorrectionFormState( - key_value=key_value, action="modify value", column=None - ) - - col_dtype = corrected_data.schema[column] - new_value, validation_error = render_value_input_widget( - column, col_dtype, current_value, tab_index - ) - - if validation_error: - st.error(validation_error) - - return CorrectionFormState( - key_value=key_value, - action="modify value", - column=column, - current_value=current_value, - new_value=new_value, - validation_error=validation_error, - ) - - -def _render_remove_value_action( - corrected_data: pl.DataFrame, - key_col: str, - key_value: str, - tab_index: int, -) -> CorrectionFormState: - """ - Render UI elements for 'remove value' action. - - Parameters - ---------- - corrected_data : pl.DataFrame - The corrected dataset. - key_col : str - The key column name. - key_value : str - The selected key value. - tab_index : int - The tab index for unique widget keys. - - Returns - ------- - CorrectionFormState - Form state with column and current value. - """ - column, current_value = _render_column_selector( - corrected_data, key_col, key_value, tab_index - ) - - return CorrectionFormState( - key_value=key_value, - action="remove value", - column=column, - current_value=current_value, - ) - - -def _render_remove_row_action(key_value: str) -> CorrectionFormState: - """ - Render UI elements for 'remove row' action. - - Parameters - ---------- - key_value : str - The selected key value. - - Returns - ------- - CorrectionFormState - Form state for row removal. - """ - st.warning("This will remove the row with the selected key value from the dataset.") - - return CorrectionFormState(key_value=key_value, action="remove row") - - def render_add_correction_form( correction_processor: CorrectionProcessor, key_col: str, @@ -490,8 +204,10 @@ def render_add_correction_form( """ Render the add correction step form. - This function orchestrates the correction form UI, delegating action-specific - rendering to helper functions to maintain low cognitive complexity. + The page picks the KEY and shows its Survey ID; the shared correction + form renders the action, new-value and reason inputs and the Apply + button. Apply goes through `_handle_apply_correction`, so the page keeps + offering any KEY, any column and every correction action. Parameters ---------- @@ -535,138 +251,29 @@ def render_add_correction_form( ) st.write(f"**Survey ID:** {survey_id_value}") - # Step 2: Select action - corr_action = st.selectbox( - label="Select Action", - options=CORRECTION_ACTIONS, - key=f"correction_action_{tab_index}", - ) - - # Step 3: Render action-specific UI and collect form state - form_state = _render_action_ui( - corr_action, corrected_data, key_col, corr_key_val, tab_index - ) - - # Step 4: Collect reason - reason = st.text_input( - label="Reason for Correction", - key=f"correction_reason_{tab_index}", - placeholder="Enter reason for correction", - ) - - # Step 5: Render apply button - _render_apply_button( + # Steps 2-5: action, new value, reason and Apply + render_correction_form( correction_processor=correction_processor, - corrected_data=corrected_data, alias=alias, key_col=key_col, - form_state=form_state, - reason=reason, - tab_index=tab_index, - survey_id_value=survey_id_value, - ) - - -def _render_action_ui( - action: str, - corrected_data: pl.DataFrame, - key_col: str, - key_value: str, - tab_index: int, -) -> CorrectionFormState: - """ - Render UI elements based on selected action type. - - Parameters - ---------- - action : str - The correction action type. - corrected_data : pl.DataFrame - The corrected dataset. - key_col : str - The key column name. - key_value : str - The selected key value. - tab_index : int - The tab index for unique widget keys. - - Returns - ------- - CorrectionFormState - Form state containing collected values. - """ - if action == "modify value": - return _render_modify_value_action( - corrected_data, key_col, key_value, tab_index - ) - - if action == "remove value": - return _render_remove_value_action( - corrected_data, key_col, key_value, tab_index - ) - - # action == "remove row" - return _render_remove_row_action(key_value) - - -def _render_apply_button( - correction_processor: CorrectionProcessor, - corrected_data: pl.DataFrame, - alias: str, - key_col: str, - form_state: CorrectionFormState, - reason: str, - tab_index: int, - survey_id_value: Any = None, -) -> None: - """ - Render apply button and handle correction application. - - Parameters - ---------- - correction_processor : CorrectionProcessor - The correction processor instance. - corrected_data : pl.DataFrame - The corrected dataset. - alias : str - The data alias/table name. - key_col : str - The key column name. - form_state : CorrectionFormState - Current form state. - reason : str - Reason for correction. - tab_index : int - The tab index for unique widget keys. - survey_id_value : Any - The Survey ID value for the selected KEY, if a Survey ID column is - configured, to record alongside the correction log entry. - """ - apply_enabled = should_enable_apply_button( - form_state.action, reason, form_state.new_value - ) - - has_validation_error = bool(form_state.validation_error) - - if st.button( - label="Apply", - key=f"correction_apply_{tab_index}", - width="stretch", - disabled=not apply_enabled or has_validation_error, - type="primary", - ): - _handle_apply_correction( - correction_processor=correction_processor, - corrected_data=corrected_data, - alias=alias, - key_col=key_col, - key_value=form_state.key_value, - action=form_state.action, - column=form_state.column, - current_value=form_state.current_value, - new_value=form_state.new_value, - reason=reason, + data=corrected_data, + key_value=corr_key_val, + key_namespace=tab_index, + source=CORRECTIONS_PAGE_SOURCE, survey_id_value=survey_id_value, + on_apply=lambda state: _handle_apply_correction( + correction_processor=correction_processor, + corrected_data=corrected_data, + alias=alias, + key_col=key_col, + key_value=state.key_value, + action=state.action, + column=state.column, + current_value=state.current_value, + new_value=state.new_value, + reason=state.reason, + survey_id_value=survey_id_value, + ), ) @@ -871,6 +478,8 @@ def _display_correction_details( s for s in correction_summaries if s["action_index"] == selected_action ) st.write(f"**Action:** {selected_summary['action']}") + if selected_summary.get("check_type"): + st.write(f"**Check type:** {selected_summary['check_type']}") st.write(f"**Key:** {selected_summary['key_value']}") if selected_summary["column"]: st.write(f"**Column:** {selected_summary['column']}") @@ -923,9 +532,11 @@ def _handle_remove_correction( def _build_correction_log_display(correction_log: pl.DataFrame) -> pl.DataFrame: """Prepare a correction log for display in the Correction Log table. - Backfills the status columns for logs saved before they existed, orders + Backfills columns missing from logs saved before they existed, orders columns so status/status_reason sit right after action, and relabels the - "ID" column as "Survey ID" for display. + "ID" column as "Survey ID" for display. "accept" rows carry the check + whose flag was accepted in check_type; source names the page that made + each entry. Parameters ---------- @@ -938,14 +549,7 @@ def _build_correction_log_display(correction_log: pl.DataFrame) -> pl.DataFrame: The log with status columns present, in display column order, ready for display. """ - if "status" not in correction_log.columns: - correction_log = correction_log.with_columns( - pl.lit("Successful").alias("status") - ) - if "status_reason" not in correction_log.columns: - correction_log = correction_log.with_columns( - pl.lit(None, dtype=pl.String).alias("status_reason") - ) + correction_log = ensure_log_columns(correction_log) display_columns = [ "date", @@ -954,10 +558,12 @@ def _build_correction_log_display(correction_log: pl.DataFrame) -> pl.DataFrame: "action", "status", "status_reason", + "check_type", "column", "current_value", "new_value", "reason", + "source", ] # "ID" holds the Survey ID value recorded for the KEY, if one was # configured - rename it for display so the column reads clearly. diff --git a/tests/processing/test_corrections.py b/tests/processing/test_corrections.py index 10ad2e37..ab1824d7 100644 --- a/tests/processing/test_corrections.py +++ b/tests/processing/test_corrections.py @@ -6,7 +6,7 @@ import polars as pl import pytest -from datasure.processing.corrections import CorrectionProcessor +from datasure.processing.corrections import CorrectionEntry, CorrectionProcessor @pytest.fixture(autouse=True) @@ -153,7 +153,10 @@ def test_get_correction_log(self, correction_processor, sample_corrections_log): result = processor.get_correction_log("test_alias") - assert result.equals(sample_corrections_log) + # Persisted columns come back unchanged; later columns are backfilled. + assert result.select(sample_corrections_log.columns).equals( + sample_corrections_log + ) mock_get.assert_called_once_with( project_id="test_project", alias="corr_log_test_alias", db_name="logs" ) @@ -1043,3 +1046,732 @@ def test_private_apply_remove_row(self, correction_processor, sample_data): assert len(result) == 2 assert not (result["survey_key"] == "key1").any() + + +# --------------------------------------------------------------------------- +# Behavior tests against an in-memory store +# +# Storage is the only collaborator replaced here: the fake keeps tables in a +# dict keyed by (project_id, db_name, alias), so these tests observe behavior +# through the processor's public methods rather than inspecting save calls. +# --------------------------------------------------------------------------- + + +@pytest.fixture +def store(): + """Patch DuckDB storage with an in-memory dict and reset processor caches.""" + tables: dict[tuple[str, str, str], pl.DataFrame] = {} + + def fake_get(project_id, alias, db_name, type="pl"): + return tables.get((project_id, db_name, alias), pl.DataFrame()) + + def fake_save(project_id, table_data, alias, db_name="raw"): + tables[(project_id, db_name, alias)] = table_data + + with ( + patch("datasure.processing.corrections.duckdb_get_table", fake_get), + patch("datasure.processing.corrections.duckdb_save_table", fake_save), + ): + _clear_processor_caches() + yield tables + _clear_processor_caches() + + +def _clear_processor_caches(): + processor = CorrectionProcessor("any") + processor.get_corrected_data.clear() + processor.get_correction_log.clear() + processor.get_data_summary.clear() + processor.get_correction_summary.clear() + + +def _seed_prep(store, data: pl.DataFrame, project_id="p1", alias="survey"): + store[(project_id, "prep", alias)] = data + + +class TestCorrectionLogSource: + """Every log entry records which page produced it.""" + + def test_new_entry_defaults_source_to_corrections_page(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + processor.apply_correction( + alias="survey", + key_col="survey_key", + key_value="key1", + action="modify value", + column="name", + current_value="John", + new_value="Johnny", + reason="typo", + ) + + log = processor.get_correction_log("survey") + assert log["source"].to_list() == ["corrections_page"] + + def test_new_entry_records_given_source(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + processor.add_correction_entry( + alias="survey", + key_value="key1", + current_id=None, + action="remove row", + column=None, + current_value=None, + new_value=None, + reason="duplicate", + source="duplicates", + ) + + log = processor.get_correction_log("survey") + assert log["source"].to_list() == ["duplicates"] + + def test_legacy_log_loads_with_backfilled_source_and_no_data_loss( + self, store, sample_corrections_log + ): + store[("p1", "logs", "corr_log_survey")] = sample_corrections_log + processor = CorrectionProcessor("p1") + + log = processor.get_correction_log("survey") + + assert log["source"].to_list() == ["corrections_page"] * 3 + assert log["status"].to_list() == ["Successful"] * 3 + assert log.select(sample_corrections_log.columns).equals(sample_corrections_log) + + def test_removing_the_only_entry_leaves_an_empty_log_with_full_schema( + self, store, sample_corrections_log + ): + # No prep table, so removal does not replay (and re-save) the log. + store[("p1", "logs", "corr_log_survey")] = sample_corrections_log[:1] + processor = CorrectionProcessor("p1") + + processor.remove_correction_entry("survey", 0) + + persisted = store[("p1", "logs", "corr_log_survey")] + assert persisted.is_empty() + assert persisted.columns == [ + "date", + "KEY", + "ID", + "action", + "column", + "current_value", + "new_value", + "reason", + "status", + "status_reason", + "source", + "check_type", + ] + + +class TestAcceptAction: + """Accepting a flagged value records a decision without changing data.""" + + def _accept_age(self, processor, key="key1", value=25, reason="verified"): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value=key, + check_type="outliers", + column="age", + current_value=value, + reason=reason, + ) + + def test_accept_logs_entry_and_leaves_corrected_data_unchanged( + self, store, sample_data + ): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + self._accept_age(processor) + + assert processor.get_corrected_data("survey").equals(sample_data) + log = processor.get_correction_log("survey") + assert log.select( + "KEY", "action", "check_type", "column", "current_value", "reason" + ).rows() == [("key1", "accept", "outliers", "age", "25", "verified")] + + def test_accept_requires_a_reason(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="reason"): + self._accept_age(processor, reason=" ") + + assert processor.get_correction_log("survey").is_empty() + + def test_accept_rejects_unknown_check_type(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="check type"): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value="key1", + check_type="missing", + column="age", + current_value=25, + reason="ok", + ) + + def test_acceptance_is_active_while_value_is_unchanged(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + self._accept_age(processor) + + active = processor.get_active_acceptances("survey", "outliers", "survey_key") + + assert active.select("KEY", "column").rows() == [("key1", "age")] + + def test_acceptance_is_inactive_once_value_changes(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + self._accept_age(processor) + + processor.apply_correction( + alias="survey", + key_col="survey_key", + key_value="key1", + action="modify value", + column="age", + current_value=25, + new_value=26, + reason="re-interview", + ) + + active = processor.get_active_acceptances("survey", "outliers", "survey_key") + assert active.is_empty() + + def test_acceptances_are_scoped_to_their_check_type(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + self._accept_age(processor) + + active = processor.get_active_acceptances("survey", "constraints", "survey_key") + + assert active.is_empty() + + def _gps_data(self): + return pl.DataFrame( + { + "survey_key": ["key1", "key2"], + "gps_lat": [5.6037, 6.6885], + "gps_lon": [-0.187, -1.6244], + } + ) + + def _accept_gps(self, processor): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value="key1", + check_type="gps", + column=None, + current_value={"gps_lat": 5.6037, "gps_lon": -0.187}, + reason="Household relocated", + ) + + def test_gps_acceptance_is_active_while_both_coordinates_match(self, store): + _seed_prep(store, self._gps_data()) + processor = CorrectionProcessor("p1") + self._accept_gps(processor) + + active = processor.get_active_acceptances("survey", "gps", "survey_key") + + assert active["KEY"].to_list() == ["key1"] + assert active["column"].to_list() == [None] + + def test_gps_acceptance_is_inactive_when_one_coordinate_changes(self, store): + _seed_prep(store, self._gps_data()) + processor = CorrectionProcessor("p1") + self._accept_gps(processor) + + processor.apply_correction( + alias="survey", + key_col="survey_key", + key_value="key1", + action="modify value", + column="gps_lon", + current_value=-0.187, + new_value=-0.2, + reason="re-recorded", + ) + + assert processor.get_active_acceptances( + "survey", "gps", "survey_key" + ).is_empty() + + def test_gps_acceptance_requires_a_coordinate_mapping(self, store): + _seed_prep(store, self._gps_data()) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="GPS"): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value="key1", + check_type="gps", + column="gps_lat", + current_value=5.6037, + reason="ok", + ) + + @pytest.mark.parametrize( + "current_value", + [{}, {"gps_lat": 5.6037}], + ids=["empty", "latitude_only"], + ) + def test_gps_acceptance_requires_both_coordinates(self, store, current_value): + _seed_prep(store, self._gps_data()) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="GPS"): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value="key1", + check_type="gps", + column=None, + current_value=current_value, + reason="ok", + ) + + assert processor.get_correction_log("survey").is_empty() + + def test_accept_rejects_a_value_that_changed_since_it_was_flagged( + self, store, sample_data + ): + # key1's age is 25; the check page still shows the stale value 99 + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="has changed since it was flagged"): + self._accept_age(processor, value=99) + + assert processor.get_correction_log("survey").is_empty() + + def test_gps_accept_rejects_a_changed_coordinate(self, store): + _seed_prep(store, self._gps_data()) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match=r"'gps_lon'.*has changed"): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value="key1", + check_type="gps", + column=None, + current_value={"gps_lat": 5.6037, "gps_lon": -0.2}, + reason="ok", + ) + + assert processor.get_correction_log("survey").is_empty() + + @pytest.mark.parametrize( + ("key", "column", "message"), + [ + ("missing", "age", "Key value 'missing' not found"), + ("key1", "height", "Column 'height' not found"), + ], + ) + def test_accept_rejects_unknown_key_or_column( + self, store, sample_data, key, column, message + ): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match=message): + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value=key, + check_type="outliers", + column=column, + current_value=25, + reason="ok", + ) + + assert processor.get_correction_log("survey").is_empty() + + def test_apply_corrections_rejects_a_stale_acceptance_atomically( + self, store, sample_data + ): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="has changed since it was flagged"): + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry( + key_value="key2", + action="modify value", + column="age", + current_value=30, + new_value=31, + reason="typo", + ), + CorrectionEntry( + key_value="key1", + action="accept", + check_type="outliers", + column="age", + current_value=99, + reason="verified", + ), + ], + source="outliers", + ) + + assert processor.get_corrected_data("survey").equals(sample_data) + assert processor.get_correction_log("survey").is_empty() + + def test_replay_skips_accept_rows(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + self._accept_age(processor, key="key1", value=25) + processor.apply_correction( + alias="survey", + key_col="survey_key", + key_value="key2", + action="modify value", + column="name", + current_value="Jane", + new_value="Janet", + reason="typo", + ) + # The accepted record is later dropped from prep, so a replayed + # accept row would have no KEY to match. + _seed_prep(store, sample_data.filter(pl.col("survey_key") != "key1")) + + failures = processor.refresh_corrected_data("survey") + + assert failures == [] + corrected = processor.get_corrected_data("survey") + assert corrected["name"].to_list() == ["Janet", "Bob"] + assert processor.get_correction_log("survey")["status"].to_list() == [ + "Successful", + "Successful", + ] + + def test_accept_rows_can_be_removed(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + self._accept_age(processor) + + summaries = processor.get_correction_summary("survey") + processor.remove_correction_entry("survey", summaries[0]["index"]) + + assert processor.get_correction_log("survey").is_empty() + assert processor.get_active_acceptances( + "survey", "outliers", "survey_key" + ).is_empty() + + +class TestCacheIsScopedToProject: + """Cached reads are keyed on project as well as alias.""" + + def test_projects_sharing_an_alias_do_not_share_corrected_data(self, store): + _seed_prep(store, pl.DataFrame({"KEY": ["a"]}), project_id="p1") + _seed_prep(store, pl.DataFrame({"KEY": ["b"]}), project_id="p2") + + first = CorrectionProcessor("p1").get_corrected_data("survey") + second = CorrectionProcessor("p2").get_corrected_data("survey") + + assert first["KEY"].to_list() == ["a"] + assert second["KEY"].to_list() == ["b"] + + def test_projects_sharing_an_alias_do_not_share_correction_logs( + self, store, sample_corrections_log + ): + store[("p1", "logs", "corr_log_survey")] = sample_corrections_log + + assert CorrectionProcessor("p1").get_correction_log("survey").height == 3 + assert CorrectionProcessor("p2").get_correction_log("survey").is_empty() + assert CorrectionProcessor("p2").get_correction_summary("survey") == [] + + +class TestApplyCorrectionsAtomically: + """Several entries apply together, or not at all.""" + + def test_applies_every_entry_and_logs_each_with_the_source( + self, store, sample_data + ): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry( + key_value="key1", + action="modify value", + column="name", + current_value="John", + new_value="Jon", + reason="keep first", + ), + CorrectionEntry( + key_value="key1", + action="modify value", + column="age", + current_value=25, + new_value=26, + reason="keep first", + ), + CorrectionEntry( + key_value="key2", action="remove row", reason="duplicate" + ), + ], + source="duplicates", + ) + + corrected = processor.get_corrected_data("survey") + assert corrected.select("survey_key", "name", "age").rows() == [ + ("key1", "Jon", 26), + ("key3", "Bob", 35), + ] + log = processor.get_correction_log("survey") + assert log["action"].to_list() == ["modify value", "modify value", "remove row"] + assert log["source"].to_list() == ["duplicates"] * 3 + + def test_an_invalid_entry_applies_nothing(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="no_such_column"): + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry( + key_value="key1", + action="modify value", + column="name", + current_value="John", + new_value="Jon", + reason="fix", + ), + CorrectionEntry( + key_value="key2", + action="modify value", + column="no_such_column", + new_value="x", + reason="fix", + ), + ], + ) + + assert processor.get_corrected_data("survey").equals(sample_data) + assert processor.get_correction_log("survey").is_empty() + + def test_an_entry_invalidated_by_an_earlier_entry_applies_nothing( + self, store, sample_data + ): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="key1"): + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry( + key_value="key1", action="remove row", reason="dup" + ), + CorrectionEntry( + key_value="key1", + action="remove value", + column="name", + reason="dup", + ), + ], + ) + + assert processor.get_corrected_data("survey").equals(sample_data) + assert processor.get_correction_log("survey").is_empty() + + def test_can_mix_an_acceptance_with_corrections(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry( + key_value="key1", + action="accept", + check_type="outliers", + column="age", + current_value=25, + reason="verified", + ), + CorrectionEntry( + key_value="key2", + action="modify value", + column="age", + current_value=30, + new_value=31, + reason="typo", + ), + ], + source="outliers", + ) + + assert processor.get_corrected_data("survey")["age"].to_list() == [25, 31, 35] + active = processor.get_active_acceptances("survey", "outliers", "survey_key") + assert active["KEY"].to_list() == ["key1"] + + def test_every_entry_needs_a_reason(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="reason"): + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry(key_value="key1", action="remove row", reason="") + ], + ) + + +class TestCorrectionSummaryDescribesAcceptances: + def test_accept_rows_are_listed_with_their_check_type(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + processor.accept_value( + alias="survey", + key_col="survey_key", + key_value="key1", + check_type="outliers", + column="age", + current_value=25, + reason="verified", + ) + + (summary,) = processor.get_correction_summary("survey") + + assert summary["action"] == "accept" + assert summary["check_type"] == "outliers" + assert summary["description"] == "Accept outliers flag on age for key key1" + + def test_gps_accept_rows_describe_the_coordinates(self, store): + _seed_prep(store, pl.DataFrame({"KEY": ["k1"], "lat": [1.0], "lon": [2.0]})) + processor = CorrectionProcessor("p1") + processor.accept_value( + alias="survey", + key_col="KEY", + key_value="k1", + check_type="gps", + column=None, + current_value={"lat": 1.0, "lon": 2.0}, + reason="verified", + ) + + (summary,) = processor.get_correction_summary("survey") + + assert summary["description"] == "Accept gps flag on coordinates for key k1" + + +class TestAcceptanceMatching: + """Acceptances compare values, not their pandas/polars string forms.""" + + def _accept(self, processor, key, column, value, key_col="KEY"): + processor.accept_value( + alias="survey", + key_col=key_col, + key_value=key, + check_type="constraints", + column=column, + current_value=value, + reason="verified", + ) + + def test_nan_accepted_value_matches_a_missing_cell(self, store): + _seed_prep( + store, + pl.DataFrame( + {"KEY": ["k1"], "age": [None]}, + schema={"KEY": pl.String, "age": pl.Int64}, + ), + ) + processor = CorrectionProcessor("p1") + self._accept(processor, "k1", "age", float("nan")) + + active = processor.get_active_acceptances("survey", "constraints", "KEY") + + assert active["KEY"].to_list() == ["k1"] + + def test_float_form_of_an_integer_matches(self, store): + # pandas turns an int column with nulls into floats: 25 arrives as 25.0 + _seed_prep(store, pl.DataFrame({"KEY": ["k1"], "age": [25]})) + processor = CorrectionProcessor("p1") + self._accept(processor, "k1", "age", 25.0) + + active = processor.get_active_acceptances("survey", "constraints", "KEY") + + assert active["KEY"].to_list() == ["k1"] + + def test_a_different_number_does_not_match(self, store): + _seed_prep(store, pl.DataFrame({"KEY": ["k1"], "age": [25]})) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="has changed since it was flagged"): + self._accept(processor, "k1", "age", 25.5) + + assert processor.get_correction_log("survey").is_empty() + + def test_non_string_key_column_is_supported(self, store): + _seed_prep(store, pl.DataFrame({"hhid": [101, 102], "age": [25, 30]})) + processor = CorrectionProcessor("p1") + self._accept(processor, "102", "age", 30, key_col="hhid") + + active = processor.get_active_acceptances("survey", "constraints", "hhid") + + assert active["KEY"].to_list() == ["102"] + + def test_duplicate_keys_must_all_hold_the_accepted_value(self, store): + _seed_prep(store, pl.DataFrame({"KEY": ["k1", "k1"], "age": [25, 40]})) + processor = CorrectionProcessor("p1") + + with pytest.raises(ValueError, match="has changed since it was flagged"): + self._accept(processor, "k1", "age", 25) + + assert processor.get_correction_log("survey").is_empty() + + +class TestApplyCorrectionsStorageFailure: + def test_failed_log_save_leaves_corrected_data_unchanged(self, store, sample_data): + _seed_prep(store, sample_data) + processor = CorrectionProcessor("p1") + processor.get_corrected_data("survey") # materialize the corrected table + + def failing_save(project_id, table_data, alias, db_name="raw"): + if db_name == "logs": + raise OSError("disk full") + store[(project_id, db_name, alias)] = table_data + + with ( + patch("datasure.processing.corrections.duckdb_save_table", failing_save), + pytest.raises(OSError, match="disk full"), + ): + processor.apply_corrections( + alias="survey", + key_col="survey_key", + entries=[ + CorrectionEntry(key_value="key1", action="remove row", reason="dup") + ], + ) + + assert processor.get_corrected_data("survey").equals(sample_data) + assert processor.get_correction_log("survey").is_empty() diff --git a/tests/replication/test_package_builder.py b/tests/replication/test_package_builder.py index f85cc447..97804932 100644 --- a/tests/replication/test_package_builder.py +++ b/tests/replication/test_package_builder.py @@ -266,3 +266,147 @@ def _duckdb_get(project_id, table, db_name): on_progress=progress_messages.append, ) assert len(progress_messages) > 0 + + +class TestPackageKeepsAcceptancesInAuditLog: + def test_correction_log_csv_keeps_accept_rows(self): + corr_log = pl.DataFrame( + { + "date": ["2026-01-01", "2026-01-02"], + "KEY": ["k1", "k2"], + "ID": [None, None], + "action": ["accept", "modify value"], + "column": ["age", "age"], + "current_value": ["25", "30"], + "new_value": [None, "31"], + "reason": ["verified", "typo"], + "source": ["outliers", "corrections_page"], + "check_type": ["outliers", None], + } + ) + + def _duckdb_get(project_id, table, db_name): + if "prep_log" in table: + return _PREP_LOG + if "corr_log" in table: + return corr_log + return _mock_loader(project_id, table, db_name) + + with patch( + "datasure.replication.package_builder.duckdb_get_table", + side_effect=_duckdb_get, + ): + zip_bytes = build_replication_package( + project_id="p", + project_name="P", + survey_name="S", + alias="s", + key_col="key", + ) + + with zipfile.ZipFile(BytesIO(zip_bytes)) as zf: + log_csv = zf.read("replication_p_s/4_output/3_logs/correction_log.csv") + corrections_do = zf.read("replication_p_s/2_scripts/4_corrections.do") + + audit = pl.read_csv(BytesIO(log_csv)) + assert audit.select("KEY", "action", "check_type").rows() == [ + ("k1", "accept", "outliers"), + ("k2", "modify value", None), + ] + assert b'"k1"' not in corrections_do + assert b'"k2"' in corrections_do + + def test_empty_correction_log_csv_has_full_header(self): + def _duckdb_get(project_id, table, db_name): + if "prep_log" in table: + return _PREP_LOG + return _mock_loader(project_id, table, db_name) + + with patch( + "datasure.replication.package_builder.duckdb_get_table", + side_effect=_duckdb_get, + ): + zip_bytes = build_replication_package( + project_id="p", + project_name="P", + survey_name="S", + alias="s", + key_col="key", + ) + + with zipfile.ZipFile(BytesIO(zip_bytes)) as zf: + header = ( + zf.read("replication_p_s/4_output/3_logs/correction_log.csv") + .decode() + .strip() + ) + assert header == ( + "date,KEY,ID,action,column,current_value,new_value,reason," + "status,status_reason,source,check_type" + ) + + +class TestPackageCountsAndLegacyLogs: + def _build(self, corr_log: pl.DataFrame) -> zipfile.ZipFile: + def _duckdb_get(project_id, table, db_name): + if "prep_log" in table: + return _PREP_LOG + if "corr_log" in table: + return corr_log + return _mock_loader(project_id, table, db_name) + + with patch( + "datasure.replication.package_builder.duckdb_get_table", + side_effect=_duckdb_get, + ): + zip_bytes = build_replication_package( + project_id="p", + project_name="P", + survey_name="S", + alias="s", + key_col="key", + ) + return zipfile.ZipFile(BytesIO(zip_bytes)) + + def test_readme_counts_exclude_acceptances(self): + corr_log = pl.DataFrame( + { + "date": ["2026-01-01"] * 3, + "KEY": ["k1", "k2", "k3"], + "ID": [None] * 3, + "action": ["accept", "accept", "remove row"], + "column": ["age", "age", None], + "current_value": ["25", "30", None], + "new_value": [None] * 3, + "reason": ["ok"] * 3, + } + ) + + with self._build(corr_log) as zf: + readme = zf.read("replication_p_s/0_README.txt").decode() + + assert "Corrections applied | 1 |" in readme + assert "accept" not in readme + + def test_legacy_log_exports_with_backfilled_columns(self): + with self._build( + pl.DataFrame( + { + "date": ["2026-01-01"], + "KEY": ["k1"], + "ID": [None], + "action": ["remove row"], + "column": [None], + "current_value": [None], + "new_value": [None], + "reason": ["dup"], + } + ) + ) as zf: + audit = pl.read_csv( + BytesIO(zf.read("replication_p_s/4_output/3_logs/correction_log.csv")) + ) + + assert audit["source"].to_list() == ["corrections_page"] + assert audit["status"].to_list() == ["Successful"] + assert "check_type" in audit.columns diff --git a/tests/replication/test_script_generators.py b/tests/replication/test_script_generators.py index d1967c3d..28f9e4a7 100644 --- a/tests/replication/test_script_generators.py +++ b/tests/replication/test_script_generators.py @@ -363,3 +363,35 @@ def test_contains_save_command(self, script): def test_header_present(self, script): assert "Import Script" in script + + +class TestCorrectionsScriptSkipsAcceptances: + """Accept rows record a review decision and never become Stata commands.""" + + def _log(self, actions: list[str]) -> pl.DataFrame: + n = len(actions) + return pl.DataFrame( + { + "action": actions, + "KEY": [f"k{i}" for i in range(n)], + "column": ["age"] * n, + "new_value": ["30" if a == "modify value" else None for a in actions], + "reason": ["checked"] * n, + "check_type": ["outliers" if a == "accept" else None for a in actions], + } + ) + + def test_accept_rows_emit_no_commands(self): + script = generate_corrections_script( + self._log(["accept", "modify value"]), "KEY", "P", "S", "0.1" + ) + + assert 'replace age = 30 if KEY == "k1"' in script + assert "k0" not in script + + def test_log_with_only_acceptances_has_no_corrections(self): + script = generate_corrections_script( + self._log(["accept", "accept"]), "KEY", "P", "S", "0.1" + ) + + assert "No corrections recorded" in script diff --git a/tests/views/test_correction_form.py b/tests/views/test_correction_form.py new file mode 100644 index 00000000..390b9f12 --- /dev/null +++ b/tests/views/test_correction_form.py @@ -0,0 +1,330 @@ +"""Tests for the shared correction form component.""" + +import sys +from contextlib import contextmanager +from types import SimpleNamespace +from typing import Any +from unittest.mock import MagicMock + +import polars as pl + +from datasure.processing.corrections import CorrectionEntry +from datasure.utils.correction_form import ( + apply_correction_entries, + get_current_value, + parse_date_value, + render_correction_form, + render_correction_inputs, +) + +_st = sys.modules["streamlit"] + +_DATA = pl.DataFrame( + {"KEY": ["k1", "k2"], "age": [25, 99], "lat": [5.6, 5.7], "lon": [-0.1, -0.2]} +) + + +@contextmanager +def _widgets(values: dict[str, Any]): + """Mock the form's widgets, answering each by its widget key. + + `values` maps a widget key to what the widget returns. Unlisted + selectboxes return their first option, text inputs "", buttons False. + Yields the mocks, which keep their recorded calls after the block exits. + """ + names = ("selectbox", "text_input", "button", "write", "warning", "error") + names += ("success", "rerun", "date_input") + originals = {name: getattr(_st, name) for name in names} + + def selectbox(label, options, key, **kwargs): + return values.get(key, next(iter(options)) if options else None) + + def text_input(label, key, **kwargs): + return values.get(key, "") + + def button(label, key, **kwargs): + return values.get(key, False) + + mocks = SimpleNamespace( + selectbox=MagicMock(side_effect=selectbox), + text_input=MagicMock(side_effect=text_input), + button=MagicMock(side_effect=button), + **{ + name: MagicMock() + for name in ("write", "warning", "error", "success", "rerun", "date_input") + }, + ) + try: + for name in names: + setattr(_st, name, getattr(mocks, name)) + yield mocks + finally: + for name, value in originals.items(): + setattr(_st, name, value) + + +def _widget_keys(mock: MagicMock) -> list[str]: + return [c.kwargs["key"] for c in mock.call_args_list] + + +class TestRenderCorrectionInputs: + """The inputs collect one entry for a prefilled or chosen target.""" + + def test_prefilled_column_skips_the_column_selector(self): + with _widgets( + { + "correction_action_outliers_0": "modify value", + "correction_new_value_outliers_0": "30", + "correction_reason_outliers_0": "typo", + } + ) as st: + state = render_correction_inputs( + _DATA, + "KEY", + "k2", + key_namespace="outliers_0", + column="age", + current_value=99, + ) + + assert "correction_col_to_modify_outliers_0" not in _widget_keys(st.selectbox) + assert state.to_entry() == CorrectionEntry( + key_value="k2", + action="modify value", + column="age", + current_value=99, + new_value="30", + reason="typo", + ) + + def test_current_value_is_looked_up_when_only_the_column_is_given(self): + with _widgets({"correction_action_x": "remove value"}): + state = render_correction_inputs( + _DATA, "KEY", "k2", key_namespace="x", column="age" + ) + + assert state.current_value == 99 + + def test_offers_only_the_allowed_actions(self): + with _widgets({}) as st: + render_correction_inputs( + _DATA, + "KEY", + "k2", + key_namespace="x", + column="age", + actions=("accept", "modify value"), + check_type="outliers", + ) + + (action_call,) = [ + c + for c in st.selectbox.call_args_list + if c.kwargs["key"] == "correction_action_x" + ] + assert list(action_call.kwargs["options"]) == ["accept", "modify value"] + + def test_widget_keys_are_namespaced(self): + with _widgets({"correction_action_a": "modify value"}) as st: + render_correction_inputs(_DATA, "KEY", "k1", key_namespace="a") + keys_a = _widget_keys(st.selectbox) + _widget_keys(st.text_input) + + with _widgets({"correction_action_b": "modify value"}) as st: + render_correction_inputs(_DATA, "KEY", "k1", key_namespace="b") + keys_b = _widget_keys(st.selectbox) + _widget_keys(st.text_input) + + assert keys_a + assert not set(keys_a) & set(keys_b) + + def test_accept_entry_carries_the_check_type(self): + with _widgets( + {"correction_action_x": "accept", "correction_reason_x": "verified"} + ): + state = render_correction_inputs( + _DATA, + "KEY", + "k2", + key_namespace="x", + column="age", + current_value=99, + actions=("accept",), + check_type="outliers", + ) + + assert state.to_entry() == CorrectionEntry( + key_value="k2", + action="accept", + column="age", + current_value=99, + reason="verified", + check_type="outliers", + ) + + def test_gps_accept_keeps_the_coordinate_pair(self): + with _widgets( + {"correction_action_x": "accept", "correction_reason_x": "moved"} + ) as st: + state = render_correction_inputs( + _DATA, + "KEY", + "k1", + key_namespace="x", + current_value={"lat": 5.6, "lon": -0.1}, + actions=("accept",), + check_type="gps", + ) + + assert "correction_col_to_modify_x" not in _widget_keys(st.selectbox) + assert state.column is None + assert state.current_value == {"lat": 5.6, "lon": -0.1} + + +class TestRenderCorrectionForm: + """The form applies its entry atomically, tagged with its source.""" + + def _render(self, processor, values, **overrides): + kwargs = dict( + correction_processor=processor, + alias="survey", + key_col="KEY", + data=_DATA, + key_value="k2", + key_namespace="x", + source="outliers", + column="age", + current_value=99, + actions=("accept", "modify value"), + check_type="outliers", + ) + kwargs.update(overrides) + with _widgets(values) as st: + render_correction_form(**kwargs) + return st + + def test_apply_saves_the_entry_with_the_source(self): + processor = MagicMock() + + st = self._render( + processor, + { + "correction_action_x": "accept", + "correction_reason_x": "verified", + "correction_apply_x": True, + }, + ) + + processor.apply_corrections.assert_called_once() + kwargs = processor.apply_corrections.call_args.kwargs + assert kwargs["source"] == "outliers" + assert kwargs["entries"][0].action == "accept" + assert st.rerun.called + + def test_apply_is_disabled_without_a_reason(self): + processor = MagicMock() + + st = self._render(processor, {"correction_action_x": "accept"}) + + (apply_call,) = [ + c + for c in st.button.call_args_list + if c.kwargs["key"] == "correction_apply_x" + ] + assert apply_call.kwargs["disabled"] is True + processor.apply_corrections.assert_not_called() + + def test_on_apply_replaces_the_default_save(self): + processor = MagicMock() + submitted = [] + + self._render( + processor, + { + "correction_action_x": "accept", + "correction_reason_x": "verified", + "correction_apply_x": True, + }, + on_apply=submitted.append, + ) + + assert [s.action for s in submitted] == ["accept"] + processor.apply_corrections.assert_not_called() + + +class TestApplyCorrectionEntries: + """Several entries are saved in one all-or-nothing call.""" + + _ENTRIES = ( + CorrectionEntry(key_value="k1", action="remove row", reason="dup"), + CorrectionEntry( + key_value="k2", + action="modify value", + column="age", + new_value="30", + reason="dup", + ), + ) + + def test_saves_every_entry_in_one_call(self): + processor = MagicMock() + + with _widgets({}) as st: + ok = apply_correction_entries( + processor, "survey", "KEY", self._ENTRIES, source="duplicates" + ) + + assert ok is True + processor.apply_corrections.assert_called_once_with( + alias="survey", + key_col="KEY", + entries=list(self._ENTRIES), + source="duplicates", + ) + assert st.success.called + + def test_reports_an_invalid_entry_and_saves_nothing(self): + processor = MagicMock() + processor.apply_corrections.side_effect = ValueError("Column 'x' not found") + + with _widgets({}) as st: + ok = apply_correction_entries( + processor, "survey", "KEY", self._ENTRIES, source="duplicates" + ) + message = st.error.call_args.args[0] + + assert ok is False + assert "Column 'x' not found" in message + assert not st.success.called + + +class TestFormErrorLogging: + """Failures are logged, not swallowed.""" + + def test_failed_apply_logs_the_traceback(self, caplog): + processor = MagicMock() + processor.apply_corrections.side_effect = OSError("disk full") + + with _widgets({}), caplog.at_level("ERROR"): + apply_correction_entries( + processor, + "survey", + "KEY", + [CorrectionEntry(key_value="k1", action="remove row", reason="dup")], + source="duplicates", + ) + + (record,) = [r for r in caplog.records if r.levelname == "ERROR"] + assert record.exc_info is not None + + def test_missing_lookup_is_logged_before_returning_none(self, caplog): + with caplog.at_level("DEBUG"): + value = get_current_value(_DATA, "KEY", "k1", "no_such_column") + + assert value is None + assert any("no_such_column" in r.getMessage() for r in caplog.records) + + def test_unparseable_date_is_logged_before_returning_none(self, caplog): + with caplog.at_level("DEBUG"): + value = parse_date_value("not-a-date") + + assert value is None + assert any("not-a-date" in r.getMessage() for r in caplog.records) diff --git a/tests/views/test_correction_view.py b/tests/views/test_correction_view.py index 811d8cf7..94adeb02 100644 --- a/tests/views/test_correction_view.py +++ b/tests/views/test_correction_view.py @@ -8,30 +8,32 @@ import polars as pl import pytest -from datasure.views.correction_view import ( +from datasure.utils.correction_form import ( CorrectionFormState, - _build_correction_log_display, - _display_correction_details, - _handle_apply_correction, - _handle_remove_correction, _render_action_ui, _render_column_selector, _render_modify_value_action, _render_remove_row_action, _render_remove_value_action, + parse_date_value, + should_enable_apply_button, + validate_numeric_input, +) +from datasure.views.correction_view import ( + _build_correction_log_display, + _display_correction_details, + _handle_apply_correction, + _handle_remove_correction, get_current_value, get_key_options, load_hfc_config, load_tab_config, main, - parse_date_value, render_add_correction_form, render_correction_input_form, render_page_header, render_page_navigation, render_value_input_widget, - should_enable_apply_button, - validate_numeric_input, validate_prerequisites, ) @@ -462,10 +464,32 @@ def test_status_columns_ordered_right_after_action(self): "action", "status", "status_reason", + "check_type", "column", "current_value", "new_value", "reason", + "source", + ] + + def test_legacy_log_shows_corrections_page_source(self): + result = _build_correction_log_display(self._base_log()) + + assert result["source"].to_list() == ["corrections_page"] + assert result["check_type"].to_list() == [None] + + def test_accept_rows_show_their_check_type_and_source(self): + log = self._base_log( + action=["accept"], + new_value=[None], + check_type=["outliers"], + source=["outliers"], + ) + + result = _build_correction_log_display(log) + + assert result.select("action", "check_type", "source").rows() == [ + ("accept", "outliers", "outliers") ] @@ -719,6 +743,11 @@ def test_modify_value_requires_new_value(self): assert should_enable_apply_button("modify value", "reason", "") is False assert should_enable_apply_button("modify value", "reason", "new") is True + def test_modify_value_accepts_zero(self): + """Zero, as text or a number, is a real value, not a missing one.""" + assert should_enable_apply_button("modify value", "reason", "0") is True + assert should_enable_apply_button("modify value", "reason", 0) is True + def test_remove_value_enabled_with_reason(self): assert should_enable_apply_button("remove value", "reason") is True @@ -931,6 +960,17 @@ def test_dispatches_remove_row(self): state = _render_action_ui("remove row", data, "KEY", "k1", 0) assert state.action == "remove row" + @pytest.mark.parametrize("action", ["modify ID", "", None]) + def test_rejects_unknown_action_instead_of_removing_row(self, action): + data = pl.DataFrame({"KEY": ["k1"], "name": ["Alice"]}) + warning = MagicMock() + with ( + _patched_st(warning=warning), + pytest.raises(ValueError, match="Unsupported correction action"), + ): + _render_action_ui(action, data, "KEY", "k1", 0) + warning.assert_not_called() + class TestHandleApplyCorrection: """Test _handle_apply_correction: validation failure, success, exception.""" @@ -1082,6 +1122,26 @@ def test_skips_absent_optional_fields(self): assert not any("Column" in t for t in written) assert not any("New Value" in t for t in written) + assert not any("Check type" in t for t in written) + + def test_shows_check_type_for_accept_rows(self): + summaries = [ + { + "action_index": "0 - accept - x", + "action": "accept", + "check_type": "outliers", + "key_value": "k1", + "column": "age", + "new_value": None, + "reason": "verified", + } + ] + + with _patched_st(write=MagicMock()): + _display_correction_details(summaries, "0 - accept - x") + written = [str(c.args[0]) for c in _st.write.call_args_list] + + assert any("Check type" in t and "outliers" in t for t in written) class TestHandleRemoveCorrection: