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
main: 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