#!/usr/bin/env python3
import argparse
import csv
import json
import os
import re
import sys
from collections import defaultdict


def normalize(value):
    return (value or "").strip().lower()


def safe_filename_tag(tag):
    safe = re.sub(r"[^A-Za-z0-9._-]+", "_", tag.strip())
    return safe.strip("_") or "unknown_tag"


def extract_managed_metadata(additional_info_cell):
    if not additional_info_cell:
        return None
    if "{" not in additional_info_cell:
        return None

    try:
        parsed = json.loads(additional_info_cell)
    except Exception:
        return None

    if not isinstance(parsed, dict):
        return None

    mm_tags = parsed.get("managed_metadata_tags")
    if isinstance(mm_tags, dict):
        return mm_tags

    return None


def collect_row_tags(
    row,
    additional_info_indexes,
    wanted_tags,
):
    row_tags = defaultdict(set)

    def add_tag_values(tags_obj):
        for tag_name, raw_values in tags_obj.items():
            if tag_name is None:
                continue
            tag_name = str(tag_name).strip()
            if not tag_name:
                continue
            if wanted_tags and tag_name not in wanted_tags:
                continue

            if isinstance(raw_values, list):
                values = raw_values
            elif raw_values is None:
                values = ["(empty)"]
            else:
                values = [raw_values]

            if not values:
                values = ["(empty)"]

            for value in values:
                value_text = str(value).strip() if value is not None else ""
                row_tags[tag_name].add(value_text if value_text else "(empty)")

    for idx in additional_info_indexes:
        if idx >= len(row):
            continue

        cell_content = row[idx]
        tags_obj = extract_managed_metadata(cell_content)
        if not isinstance(tags_obj, dict):
            continue

        add_tag_values(tags_obj)

    return row_tags


def classify_states(states):
    flags = {
        "indexed": False,
        "skipped": False,
        "delivered_to_all": False,
        "not_delivered": False,
    }

    had_state = False

    for st in states:
        s = normalize(st)
        if not s:
            continue

        had_state = True

        if "skip" in s:
            flags["skipped"] = True
        if "not delivered" in s or "not present on this source" in s:
            flags["not_delivered"] = True
        if "downloaded" in s or "pre-seeded" in s or "preseeded" in s:
            flags["delivered_to_all"] = True

        if any(token in s for token in ("indexed", "added", "modified", "moved", "renamed", "created", "updated", "archived")):
            flags["indexed"] = True

    if (not flags["indexed"]) and had_state and (not flags["skipped"]) and (not flags["not_delivered"]):
        flags["indexed"] = True

    if flags["not_delivered"]:
        flags["delivered_to_all"] = False

    return flags


def fresh_metrics():
    return {
        "indexed": 0,
        "skipped": 0,
        "delivered_to_all": 0,
        "not_delivered": 0,
    }


def build_parser():
    parser = argparse.ArgumentParser(
        prog="aggregate_managed_metadata.py",
        description=(
            "Aggregate managed metadata tags from all-report CSV by tag value and "
            "generate per-tag output files."
        ),
    )
    parser.add_argument("--input", "-i", required=True, help="Path to the report CSV file")
    parser.add_argument(
        "--tag-name",
        "-t",
        action="append",
        default=[],
        help="Aggregate only this tag name (can be passed multiple times)",
    )
    parser.add_argument(
        "--out-dir",
        "-o",
        default=os.getcwd(),
        help="Output directory for generated files (default: current directory)",
    )
    parser.add_argument(
        "--output",
        help="Single output file path (combined output). If omitted, writes to <out-dir>/overall_by_tags.csv",
    )
    parser.add_argument(
        "--split-per-tag",
        action="store_true",
        help="Write one output file per tag name",
    )
    parser.add_argument(
        "--show-no-tags-files",
        dest="show_no_tags_stats",
        action="store_true",
        default=False,
        help="Show count of files without managed metadata tags",
    )
    return parser


def ensure_dir(path):
    if os.path.isdir(path):
        return

    try:
        os.makedirs(path)
    except OSError:
        if not os.path.isdir(path):
            raise


def main():
    args = build_parser().parse_args()

    input_file = args.input
    out_dir = args.out_dir

    if not os.path.isfile(input_file):
        print("Error: input file not found: {0}".format(input_file), file=sys.stderr)
        return 1

    try:
        ensure_dir(out_dir)
    except OSError as exc:
        print("Error: cannot create output directory {0}: {1}".format(out_dir, exc), file=sys.stderr)
        return 1

    if args.split_per_tag and args.output:
        print("Warning: --output is ignored when --split-per-tag is used", file=sys.stderr)
    if args.split_per_tag and args.show_no_tags_stats:
        print("Warning: --show-no-tags-files is only supported in base flow (ignored with --split-per-tag)", file=sys.stderr)

    wanted_tags = None
    if args.tag_name:
        cleaned = {t.strip() for t in args.tag_name if t and t.strip()}
        if cleaned:
            wanted_tags = cleaned

    counts = defaultdict(lambda: defaultdict(fresh_metrics))
    processed_file_rows = 0
    tagged_file_rows = 0

    # Increase CSV field size limit for large cells (e.g., big JSON in additional_info)
    csv.field_size_limit(int(2**31 - 1))

    try:
        fp = open(input_file, "r", encoding="utf-8-sig", newline="")
    except (IOError, OSError) as exc:
        print("Error: cannot read input file {0}: {1}".format(input_file, exc), file=sys.stderr)
        return 1

    with fp:
        reader = csv.reader(fp)

        try:
            header1 = next(reader)
            header2 = next(reader)
        except StopIteration:
            print("CSV must contain at least two header lines", file=sys.stderr)
            return 1

        type_idx = next((i for i, v in enumerate(header1) if normalize(v) == "type"), None)
        if type_idx is None:
            print('Could not find "type" column', file=sys.stderr)
            return 1

        state_indexes = [i for i, v in enumerate(header2) if normalize(v) == "state"]
        additional_info_indexes = [i for i, v in enumerate(header2) if normalize(v) == "additional info"]
        if not state_indexes:
            print('Could not find "state" columns in second header line', file=sys.stderr)
            return 1
        if not additional_info_indexes:
            print('Could not find "additional info" columns in second header line', file=sys.stderr)
            return 1

        for row in reader:
            if type_idx >= len(row):
                continue

            if normalize(row[type_idx]) != "file":
                continue

            processed_file_rows += 1

            states = [row[i].strip() for i in state_indexes if i < len(row) and row[i].strip()]
            flags = classify_states(states)

            row_tags = collect_row_tags(
                row=row,
                additional_info_indexes=additional_info_indexes,
                wanted_tags=wanted_tags,
            )

            if not row_tags:
                continue

            tagged_file_rows += 1

            for tag_name, tag_values in row_tags.items():
                for tag_value in tag_values:
                    metrics = counts[tag_name][tag_value]
                    if flags["indexed"]:
                        metrics["indexed"] += 1
                    if flags["skipped"]:
                        metrics["skipped"] += 1
                    if flags["delivered_to_all"]:
                        metrics["delivered_to_all"] += 1
                    if flags["not_delivered"]:
                        metrics["not_delivered"] += 1

    if not counts:
        if wanted_tags:
            print(
                "No managed metadata found for requested tag(s): {0}".format(', '.join(sorted(wanted_tags))),
                file=sys.stderr,
            )
        else:
            print("No managed metadata found for file rows", file=sys.stderr)
        return 1

    generated = 0

    if args.split_per_tag:
        for tag_name in sorted(counts.keys()):
            tag_map = counts[tag_name]
            if not tag_map:
                continue

            out_file = os.path.join(out_dir, "overall_{0}_data.csv".format(safe_filename_tag(tag_name)))
            try:
                out_fp = open(out_file, "w", encoding="utf-8", newline="")
            except (IOError, OSError) as exc:
                print("Error: cannot write {0}: {1}".format(out_file, exc), file=sys.stderr)
                return 1
            with out_fp:
                writer = csv.writer(out_fp)
                writer.writerow(["tag_value", "indexed", "skipped", "delivered_to_all", "not_delivered"])
                totals = fresh_metrics()
                for tag_value in sorted(tag_map.keys()):
                    m = tag_map[tag_value]
                    writer.writerow([
                        tag_value,
                        m["indexed"],
                        m["skipped"],
                        m["delivered_to_all"],
                        m["not_delivered"],
                    ])
                    totals["indexed"] += m["indexed"]
                    totals["skipped"] += m["skipped"]
                    totals["delivered_to_all"] += m["delivered_to_all"]
                    totals["not_delivered"] += m["not_delivered"]

                writer.writerow([
                    "Total",
                    totals["indexed"],
                    totals["skipped"],
                    totals["delivered_to_all"],
                    totals["not_delivered"],
                ])

            print("Generated {0}".format(out_file))
            generated += 1
    else:
        out_file = args.output or os.path.join(out_dir, "overall_by_tags.csv")
        try:
            out_fp = open(out_file, "w", encoding="utf-8", newline="")
        except (IOError, OSError) as exc:
            print("Error: cannot write {0}: {1}".format(out_file, exc), file=sys.stderr)
            return 1
        with out_fp:
            writer = csv.writer(out_fp)
            writer.writerow(["tag_name", "tag_value", "indexed", "skipped", "delivered_to_all", "not_delivered"])
            for tag_name in sorted(counts.keys()):
                tag_map = counts[tag_name]
                if not tag_map:
                    continue

                totals = fresh_metrics()
                for tag_value in sorted(tag_map.keys()):
                    m = tag_map[tag_value]
                    writer.writerow([
                        tag_name,
                        tag_value,
                        m["indexed"],
                        m["skipped"],
                        m["delivered_to_all"],
                        m["not_delivered"],
                    ])

                    totals["indexed"] += m["indexed"]
                    totals["skipped"] += m["skipped"]
                    totals["delivered_to_all"] += m["delivered_to_all"]
                    totals["not_delivered"] += m["not_delivered"]

                writer.writerow([
                    tag_name,
                    "Total",
                    totals["indexed"],
                    totals["skipped"],
                    totals["delivered_to_all"],
                    totals["not_delivered"],
                ])

            # Add no tags statistics if requested
            if args.show_no_tags_stats:
                no_tags_count = processed_file_rows - tagged_file_rows
                writer.writerow(["No tags", "Total", no_tags_count, 0, 0, 0])


        generated = 1

    if generated == 0:
        print("No output files generated (check --tag-name values)", file=sys.stderr)
        return 1

    print("Processed file rows: {0}".format(processed_file_rows))
    print("Processed tagged file rows: {0}".format(tagged_file_rows))
    if args.show_no_tags_stats and not args.split_per_tag:
        print("Files without managed metadata tags: {0}".format(processed_file_rows - tagged_file_rows))

    return 0


if __name__ == "__main__":
    raise SystemExit(main())
