#!/usr/bin/env python3
"""Identify which story parts/batches participate in duplication violations."""

import argparse
import collections
import json
from pathlib import Path


def iter_groups(report):
    for group in report.get("exact_sentence_duplicates", []):
        yield "exact_sentence", group.get("locations", [])
    for group in report.get("exact_paragraph_duplicates", []):
        yield "exact_paragraph", group.get("locations", [])
    for size, groups in report.get("repeated_ngrams", {}).items():
        for group in groups:
            yield f"ngram_{size}", group.get("locations", [])
    for group in report.get("near_duplicates", []):
        yield "near_duplicate", group.get("locations", [])


def location_part(location):
    if isinstance(location, (list, tuple)) and location:
        return int(location[0])
    if isinstance(location, dict) and "part" in location:
        return int(location["part"])
    raise ValueError(f"unsupported duplication location: {location!r}")


def batch_for_part(part, ranges):
    for start, end in ranges:
        if start <= part <= end:
            return f"{start}-{end}"
    return "unmapped"


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("report", type=Path)
    parser.add_argument("--ranges", default="1-4,5-8,9-12")
    args = parser.parse_args()
    ranges = []
    for item in args.ranges.split(","):
        start, end = map(int, item.split("-", 1))
        ranges.append((start, end))

    report = json.loads(args.report.read_text())
    part_locations = collections.Counter()
    batch_locations = collections.Counter()
    group_part_sets = collections.Counter()
    group_types = collections.Counter()

    for kind, locations in iter_groups(report):
        parts = []
        for location in locations:
            part = location_part(location)
            parts.append(part)
            part_locations[part] += 1
            batch_locations[batch_for_part(part, ranges)] += 1
        if parts:
            group_types[kind] += 1
            group_part_sets[tuple(sorted(set(parts)))] += 1

    touched = sorted(batch for batch, count in batch_locations.items() if count)
    clean = [f"{start}-{end}" for start, end in ranges if f"{start}-{end}" not in touched]
    result = {
        "status": "passed",
        "report": str(args.report),
        "story_sha256": report.get("story_sha256"),
        "violation_count": report.get("violation_count"),
        "part_location_counts": dict(sorted(part_locations.items())),
        "batch_location_counts": dict(sorted(batch_locations.items())),
        "touched_batches": touched,
        "clean_batches": clean,
        "group_type_counts": dict(sorted(group_types.items())),
        "top_part_sets": [
            {"parts": list(parts), "groups": count}
            for parts, count in group_part_sets.most_common(20)
        ],
        "remediation_rule": "revise only touched batches; retain clean parent batches byte-for-byte",
    }
    print(json.dumps(result, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
