"""Find exact duplicate sample_id values in a UTF-8 CSV. Python 3.10+."""
import argparse
from contextlib import closing
import csv
from pathlib import Path
import sqlite3
import sys
import tempfile

UPSERT = """
INSERT INTO ids VALUES (?, 1, ?, ?)
ON CONFLICT(sample_id) DO UPDATE SET
    occurrences = ids.occurrences + 1,
    last_record = excluded.last_record
"""


def audit(input_path, output, *, work_dir=None):
    """Validate all input before writing a CSV report; return audit counts."""
    with tempfile.TemporaryDirectory(prefix="csv-ids-", dir=work_dir) as directory:
        with closing(sqlite3.connect(Path(directory) / "ids.sqlite3")) as db:
            db.execute("""CREATE TABLE ids (
                sample_id TEXT PRIMARY KEY NOT NULL,
                occurrences INTEGER NOT NULL,
                first_record INTEGER NOT NULL,
                last_record INTEGER NOT NULL
            )""")
            records = 0
            with db, open(input_path, encoding="utf-8-sig", newline="") as handle:
                reader = csv.reader(handle, strict=True)
                header = next(reader, None)
                if not header or any(not name for name in header):
                    raise ValueError("Expected a nonempty CSV header")
                if len(set(header)) != len(header) or "sample_id" not in header:
                    raise ValueError("Header must be unique and include sample_id")
                column = header.index("sample_id")
                for records, row in enumerate(reader, start=1):
                    if len(row) != len(header):
                        raise ValueError(f"Record {records}: wrong field count")
                    sample_id = row[column]
                    if not sample_id or sample_id != sample_id.strip():
                        raise ValueError(f"Record {records}: empty or padded sample_id")
                    db.execute(UPSERT, (sample_id, records, records))
            unique_ids = db.execute("SELECT COUNT(*) FROM ids").fetchone()[0]
            writer = csv.writer(output, lineterminator="\n")
            writer.writerow(["sample_id", "occurrences", "first_record", "last_record"])
            duplicate_ids = 0
            for row in db.execute("""SELECT * FROM ids WHERE occurrences > 1
                                     ORDER BY sample_id COLLATE BINARY"""):
                writer.writerow(row)
                duplicate_ids += 1
    return records, unique_ids, duplicate_ids


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path)
    parser.add_argument("--work-dir", type=Path, help="existing directory for scratch database")
    args = parser.parse_args()
    try:
        records, unique_ids, duplicates = audit(args.input, sys.stdout, work_dir=args.work_dir)
    except (OSError, UnicodeError, ValueError, csv.Error, sqlite3.Error) as exc:
        parser.exit(1, f"error: {exc}\n")
    print(f"records={records} unique_ids={unique_ids} duplicate_ids={duplicates}", file=sys.stderr)


if __name__ == "__main__":
    main()
