audit-labs/audit-tools

A collection of scripts, queries, and other goodies you can use in an audit.

clone: git clone https://gitbay.org/audit-labs/audit-tools.git

v1.0.0: sampling/sampling_tool/cli.py · raw

  1"""CLI orchestration for the audit sampling tool."""
  2
  3from __future__ import annotations
  4
  5import sys
  6from argparse import ArgumentParser, Namespace
  7from datetime import datetime, timezone
  8from pathlib import Path
  9from types import SimpleNamespace
 10
 11import pandas as pd
 12
 13from .filters import apply_filters, parse_filters
 14from .io import AuditSamplingError, load_population, sha256_file, write_csv
 15from .manifest import build_manifest, write_manifest
 16from .methods import (
 17    ensure_seed,
 18    largest_remainder_allocation,
 19    random_sample,
 20    stratified_sample,
 21)
 22from .reconciliation import build_reconciliation, build_strata_summary
 23from .reporting import RunLogger, build_methodology
 24from .validation import validate_and_prepare
 25
 26# Output filenames, defined once so writer and tracker never drift.
 27_POPULATION_VALIDATED_CSV = "population_validated.csv"
 28_EXCLUDED_ROWS_CSV = "excluded_rows.csv"
 29_DUPLICATE_IDS_CSV = "duplicate_ids.csv"
 30
 31
 32def build_parser() -> ArgumentParser:
 33    parser = ArgumentParser(description="Generate documented audit samples.")
 34    parser.add_argument("--input")
 35    parser.add_argument("--sheet")
 36    parser.add_argument("--id-column")
 37    parser.add_argument("--method", choices=["random", "stratified", "validate-only"])
 38    parser.add_argument("--sample-size", type=int)
 39    parser.add_argument("--stratify-column")
 40    parser.add_argument("--strata-counts")
 41    parser.add_argument("--strata-proportions")
 42    parser.add_argument("--seed", type=int)
 43    parser.add_argument("--out")
 44    parser.add_argument("--exclude-blank-id", action="store_true", default=None)
 45    parser.add_argument("--dedupe-id", choices=["fail", "first", "last"])
 46    parser.add_argument("--filter", action="append", dest="filter_values")
 47    parser.add_argument("--config")
 48    parser.add_argument("--allow-shortfall", action="store_true", default=None)
 49    return parser
 50
 51
 52def load_config(path: str | None) -> dict[str, object]:
 53    if not path:
 54        return {}
 55    try:
 56        import yaml
 57    except ImportError as exc:
 58        raise AuditSamplingError(
 59            "YAML config support requires PyYAML. Install requirements.txt."
 60        ) from exc
 61
 62    config_path = Path(path).resolve()
 63    if not config_path.is_file():
 64        raise AuditSamplingError(f"Config file not found: {path}")
 65    with config_path.open("r", encoding="utf-8") as handle:
 66        data = yaml.safe_load(handle) or {}
 67    if not isinstance(data, dict):
 68        raise AuditSamplingError("Config file must contain a YAML mapping.")
 69    return data
 70
 71
 72def merge_options(args: Namespace, config: dict[str, object]) -> SimpleNamespace:
 73    mapping = {
 74        "input": "input",
 75        "sheet": "sheet",
 76        "id_column": "id_column",
 77        "method": "method",
 78        "sample_size": "sample_size",
 79        "stratify_column": "stratify_column",
 80        "strata_counts": "strata_counts",
 81        "strata_proportions": "strata_proportions",
 82        "seed": "seed",
 83        "out": "out",
 84        "exclude_blank_id": "exclude_blank_id",
 85        "dedupe_id": "dedupe_id",
 86        "allow_shortfall": "allow_shortfall",
 87    }
 88    merged: dict[str, object] = {}
 89    for attr, key in mapping.items():
 90        cli_value = getattr(args, attr)
 91        merged[attr] = cli_value if cli_value is not None else config.get(key)
 92
 93    config_filters = config.get("filters", {})
 94    if isinstance(config_filters, list):
 95        config_filters = parse_filters(config_filters)
 96    if not isinstance(config_filters, dict):
 97        raise AuditSamplingError("Config filters must be a mapping or list.")
 98    cli_filters = parse_filters(args.filter_values)
 99    merged["filters"] = {**config_filters, **cli_filters}
100
101    merged["out"] = merged["out"] or "./output"
102    merged["dedupe_id"] = merged["dedupe_id"] or "fail"
103    merged["exclude_blank_id"] = bool(merged["exclude_blank_id"])
104    merged["allow_shortfall"] = bool(merged["allow_shortfall"])
105    if merged["input"] is None:
106        raise AuditSamplingError("--input is required unless provided by --config.")
107    if merged["method"] is None:
108        raise AuditSamplingError("--method is required unless provided by --config.")
109    if merged["method"] not in {"random", "stratified", "validate-only"}:
110        raise AuditSamplingError(
111            "--method must be random, stratified, or validate-only."
112        )
113    if merged["sample_size"] is not None:
114        merged["sample_size"] = int(merged["sample_size"])
115    if merged["seed"] is not None:
116        merged["seed"] = int(merged["seed"])
117    return SimpleNamespace(**merged)
118
119
120def create_run_dir(out_dir: Path, now: datetime) -> Path:
121    run_dir = out_dir / f"sample_{now.strftime('%Y-%m-%d_%H%M%S')}"
122    suffix = 1
123    while True:
124        candidate = run_dir if suffix == 1 else out_dir / f"{run_dir.name}_{suffix}"
125        try:
126            candidate.mkdir(parents=True, exist_ok=False)
127            return candidate
128        except FileExistsError:
129            suffix += 1
130
131
132def main(argv: list[str] | None = None) -> int:
133    parser = build_parser()
134    args = parser.parse_args(argv)
135    try:
136        options = merge_options(args, load_config(args.config))
137        run(options)
138        return 0
139    except AuditSamplingError as exc:
140        print(f"ERROR: {exc}", file=sys.stderr)
141        return 1
142
143
144def run(options) -> Path:
145    now = datetime.now(timezone.utc)
146    timestamp = now.replace(microsecond=0).isoformat().replace("+00:00", "Z")
147    run_dir = create_run_dir(Path(options.out), now)
148    logger = RunLogger()
149    logger.log(f"start timestamp: {timestamp}")
150    logger.log(f"input path: {options.input}")
151    logger.log(f"method: {options.method}")
152    logger.log(f"output folder: {run_dir}")
153    print(f"Writing audit sample package to {run_dir}")
154
155    output_files: list[str] = []
156    input_path = Path(options.input)
157    input_hash = ""
158    source = pd.DataFrame()
159    filtered = pd.DataFrame()
160    validated = pd.DataFrame()
161    excluded_rows = pd.DataFrame()
162    duplicate_rows = pd.DataFrame()
163    sample = pd.DataFrame()
164    strata_rows: list[dict[str, object]] = []
165    random_seed: int | None = options.seed
166    effective_id_column = None
167    id_column_omitted = False
168    blank_id_count = 0
169    duplicate_id_count = 0
170
171    try:
172        input_hash = sha256_file(input_path)
173        source = load_population(input_path, options.sheet)
174        filtered, filter_excluded = apply_filters(source, options.filters)
175        validated = filtered.copy()
176        excluded_rows = filter_excluded.copy()
177        validation = validate_and_prepare(filtered, options)
178        validated = validation.population
179        excluded_rows = _concat_nonempty([filter_excluded, validation.excluded_rows])
180        duplicate_rows = validation.duplicate_rows
181        effective_id_column = validation.effective_id_column
182        id_column_omitted = validation.id_column_omitted
183        blank_id_count = validation.blank_id_count
184        duplicate_id_count = validation.duplicate_id_count
185        for warning in validation.warnings:
186            print(f"WARNING: {warning}")
187            logger.warning(warning)
188
189        if options.method in {"random", "stratified"}:
190            random_seed, generated = ensure_seed(options.seed)
191            if generated:
192                warning = f"No seed provided; generated seed {random_seed}."
193                print(f"WARNING: {warning}")
194                logger.warning(warning)
195
196        if options.method == "random":
197            sample = random_sample(validated, options.sample_size, random_seed)
198            sample = add_sample_metadata(sample, options, random_seed, timestamp, None)
199        elif options.method == "stratified":
200            counts = validation.strata_counts
201            if validation.strata_proportions:
202                counts = largest_remainder_allocation(
203                    options.sample_size, validation.strata_proportions
204                )
205            sampled, strata_rows = stratified_sample(
206                validated,
207                options.stratify_column,
208                counts,
209                random_seed,
210                options.allow_shortfall,
211            )
212            sample = add_sample_metadata(
213                sampled, options, random_seed, timestamp, options.stratify_column
214            )
215
216        _write_outputs(
217            run_dir,
218            options,
219            validated,
220            excluded_rows,
221            duplicate_rows,
222            sample,
223            strata_rows,
224            output_files,
225        )
226        logger.log("final status: success")
227        print("Audit sampling run completed.")
228    except AuditSamplingError as exc:
229        duplicate_rows = exc.artifacts.get("duplicate_rows", duplicate_rows)
230        logger.error(str(exc))
231        print(f"ERROR: {exc}", file=sys.stderr)
232        _write_failure_outputs(
233            run_dir,
234            filtered,
235            excluded_rows,
236            duplicate_rows,
237            output_files,
238        )
239        logger.log("final status: failed")
240        raise
241    finally:
242        reconciliation = build_reconciliation(
243            source_rows=len(source),
244            blank_id_count=blank_id_count,
245            duplicate_id_count=duplicate_id_count or len(duplicate_rows),
246            excluded_rows=len(excluded_rows),
247            validated_rows=len(validated),
248            requested_sample_size=_requested_sample_size(options, strata_rows),
249            final_sample_size=len(sample),
250        )
251        write_csv(reconciliation, run_dir / "population_reconciliation.csv")
252        _track(output_files, "population_reconciliation.csv")
253        methodology = build_methodology(
254            input_file=input_path,
255            input_sheet=options.sheet,
256            input_sha256=input_hash,
257            source_row_count=len(source),
258            options=options,
259            effective_id_column=effective_id_column,
260            id_column_omitted=id_column_omitted,
261            duplicate_id_count=duplicate_id_count or len(duplicate_rows),
262            blank_id_count=blank_id_count,
263            sample_size_actual=len(sample),
264            random_seed=random_seed,
265            strata_summary=strata_rows,
266        )
267        (run_dir / "methodology.txt").write_text(methodology)
268        _track(output_files, "methodology.txt")
269        manifest = build_manifest(
270            run_timestamp_utc=timestamp,
271            input_file=input_path,
272            input_sha256=input_hash,
273            input_sheet=options.sheet,
274            options=options,
275            effective_id_column=effective_id_column,
276            id_column_omitted=id_column_omitted,
277            sample_size_actual=len(sample),
278            random_seed=random_seed,
279            source_row_count=len(source),
280            validated_population_count=len(validated),
281            excluded_row_count=len(excluded_rows),
282            duplicate_id_count=duplicate_id_count or len(duplicate_rows),
283            blank_id_count=blank_id_count,
284            output_files=sorted(output_files + ["run.log", "manifest.json"]),
285        )
286        write_manifest(manifest, run_dir / "manifest.json")
287        logger.write(run_dir / "run.log")
288    return run_dir
289
290
291def add_sample_metadata(
292    sample: pd.DataFrame,
293    options,
294    seed: int,
295    timestamp: str,
296    stratum_column: str | None,
297) -> pd.DataFrame:
298    source_columns = [column for column in sample.columns if column != "_audit_row_id"]
299    output = sample[source_columns].copy()
300    output.insert(0, "_selected_at_utc", timestamp)
301    output.insert(0, "_random_seed", seed)
302    output.insert(0, "_stratum", output[stratum_column] if stratum_column else "")
303    output.insert(0, "_selection_method", options.method)
304    output.insert(0, "_source_row_number", output.pop("_source_row_number"))
305    output.insert(0, "_sample_id", range(1, len(output) + 1))
306    return output
307
308
309def _write_outputs(
310    run_dir: Path,
311    options,
312    validated: pd.DataFrame,
313    excluded_rows: pd.DataFrame,
314    duplicate_rows: pd.DataFrame,
315    sample: pd.DataFrame,
316    strata_rows: list[dict[str, object]],
317    output_files: list[str],
318) -> None:
319    write_csv(validated, run_dir / _POPULATION_VALIDATED_CSV)
320    _track(output_files, _POPULATION_VALIDATED_CSV)
321    if options.method in {"random", "stratified"}:
322        write_csv(sample, run_dir / "sample.csv")
323        _track(output_files, "sample.csv")
324    if not excluded_rows.empty:
325        write_csv(excluded_rows, run_dir / _EXCLUDED_ROWS_CSV)
326        _track(output_files, _EXCLUDED_ROWS_CSV)
327    if not duplicate_rows.empty:
328        write_csv(duplicate_rows, run_dir / _DUPLICATE_IDS_CSV)
329        _track(output_files, _DUPLICATE_IDS_CSV)
330    if options.method == "stratified":
331        write_csv(build_strata_summary(strata_rows), run_dir / "strata_summary.csv")
332        _track(output_files, "strata_summary.csv")
333
334
335def _write_failure_outputs(
336    run_dir: Path,
337    filtered: pd.DataFrame,
338    excluded_rows: pd.DataFrame,
339    duplicate_rows: pd.DataFrame,
340    output_files: list[str],
341) -> None:
342    if not filtered.empty:
343        write_csv(filtered, run_dir / _POPULATION_VALIDATED_CSV)
344        _track(output_files, _POPULATION_VALIDATED_CSV)
345    if not excluded_rows.empty:
346        write_csv(excluded_rows, run_dir / _EXCLUDED_ROWS_CSV)
347        _track(output_files, _EXCLUDED_ROWS_CSV)
348    if not duplicate_rows.empty:
349        write_csv(duplicate_rows, run_dir / _DUPLICATE_IDS_CSV)
350        _track(output_files, _DUPLICATE_IDS_CSV)
351
352
353def _concat_nonempty(frames: list[pd.DataFrame]) -> pd.DataFrame:
354    nonempty = [frame for frame in frames if frame is not None and not frame.empty]
355    if not nonempty:
356        return pd.DataFrame()
357    return pd.concat(nonempty, ignore_index=True)
358
359
360def _track(output_files: list[str], filename: str) -> None:
361    if filename not in output_files:
362        output_files.append(filename)
363
364
365def _requested_sample_size(options, strata_rows: list[dict[str, object]]) -> int | None:
366    if options.sample_size is not None:
367        return options.sample_size
368    if strata_rows:
369        return int(sum(row["Requested Sample Count"] for row in strata_rows))
370    return None