#!/usr/bin/env -S uv run --script
# /// script
# requires-python = "==3.13"
# dependencies = [
#     "pydantic==2.11.7",
# ]
# ///

import json
import random
from collections import Counter
from pathlib import Path
from pprint import pprint

from models import (
    Distribution,
    DistributionKind,
    FileStatus,
    FileReproductionResult,
    Package,
    ReproductionResult,
    DistributionStatus,
)


def print_most_common(title: str, items: list):
    print(f"### {title} ###")
    pprint(Counter(items).most_common())
    print()


def extract_tag(filename: str):
    p = Path(filename)
    stem = p.stem
    try:
        idx = stem.find("cpython-")
        idx2 = stem[idx + 8 :].find("-")
        return stem[idx + 8 + idx2 + 1 :]
    except Exception:
        return None
    return None


def extract_input(
    packages: list[Package],
    package_name: str,
    package_version: str,
    distribution_kind: DistributionKind,
    path: str,
) -> Package:
    print("Looking for package ", package_name, package_version)
    for package in packages:
        if package.name == package_name and package.version == package_version:
            print("Looking for dist", distribution_kind.value)
            for dist in package.distributions:
                if dist.kind == distribution_kind:
                    for file in dist.cache_files:
                        print("Looking for file", path)
                        if str(file.cache_path).endswith(path):
                            return Package(
                                name=package.name,
                                version=package.version,
                                proton_path=package.proton_path,
                                distributions=[
                                    Distribution(
                                        kind=distribution_kind,
                                        dist_path=Path(
                                            "/home/user/Coding/uni/python-cache-poisoning/cache-reproduction-tests/test-package/"
                                        ).joinpath(dist.dist_path.name),
                                        cache_files=[file],
                                    )
                                ],
                            )


def main():
    with open("./results/2025-09-09-compact.json") as fp:
        input_data = json.load(fp)
    packages = [Package(**package) for package in input_data]
    with open("./results/reproducibility-results.json") as fp:
        reproduction_results: list[ReproductionResult] = [
            ReproductionResult(**e) for e in json.load(fp)
        ]
    status = []
    version_choices = []
    versions = []
    cache_validity = []
    invalidation_modes = []
    reproducibility_x_validity = []
    validity_x_validation_mode = []
    # Irreproducibility classes:
    # 1. files where the cache is invalid
    # 2. files where the cache is older than the source file
    # 3. files with additional tags like `-pycache` that are probably created in a different way
    # than regular pycache files
    # 4.

    irreproducible_files: list[FileReproductionResult] = []
    irreproducible_files_invalid_cache = []
    irreproducible_files_old_cache = []
    irreproducible_files_additional_tag = []
    tags = []

    for result in reproduction_results:
        for distribution in result.distributions:
            status.append(distribution.status)
            if distribution.status == DistributionStatus.UnsupportedKind:
                continue
            for file in distribution.files:
                status.append(file.status)
                cache_validity.append(file.cache_valid)
                invalidation_modes.append(file.invalidation_mode)
                reproducibility_x_validity.append((file.status, file.cache_valid))
                validity_x_validation_mode.append(
                    (file.cache_valid, file.invalidation_mode)
                )
                tags.append(extract_tag(file.file_name))
                if file.status == FileStatus.Reproducible:
                    version_choices.append(file.python_version_choice)
                    versions.append(file.python_version)
                if file.status == FileStatus.Irreproducible:
                    irreproducible_files.append(file)
    irreproducible_files_invalid_cache = [
        f for f in irreproducible_files if not f.cache_valid
    ]
    irreproducible_files = [f for f in irreproducible_files if f.cache_valid]
    irreproducible_files_old_cache = [
        f
        for f in irreproducible_files
        if f.source_modification_time
        and f.cache_modification_time
        and f.source_modification_time > f.cache_modification_time
    ]
    irreproducible_files = [
        f
        for f in irreproducible_files
        if f.source_modification_time
        and f.cache_modification_time
        and f.source_modification_time <= f.cache_modification_time
    ]
    irreproducible_files_additional_tag = [
        f for f in irreproducible_files if f.additional_tag is not None
    ]
    irreproducible_files = [f for f in irreproducible_files if f.additional_tag is None]

    # General stats
    print_most_common("Reproducibility status", status)
    print_most_common("Python version choice methods", version_choices)
    print_most_common("Python versions used for successful reproduction", versions)
    print_most_common("Cache validities", cache_validity)
    print_most_common("Used invalidation modes", invalidation_modes)
    print_most_common(
        "Relation of reproducibility status and cache validity",
        reproducibility_x_validity,
    )
    print_most_common(
        "Relation of cache validity and invalidation mode", validity_x_validation_mode
    )
    print_most_common("Most used tags: ", tags)

    # Classification
    print(
        "Irreproducible files with invalid cache:",
        len(irreproducible_files_invalid_cache),
    )
    print(
        "Irreproducible files with old cache:",
        len(irreproducible_files_old_cache),
        irreproducible_files_old_cache,
    )
    print(
        "Irreproducible files additional tag:", len(irreproducible_files_additional_tag)
    )
    print("Unclassified irreproducible files:", len(irreproducible_files))

    if len(irreproducible_files) > 0:
        print("### Samples:")
        samples = [
            extract_input(packages, *r).model_dump()
            for r in random.sample(
                irreproducible_files, min(15, len(irreproducible_files))
            )
        ]
        with open("./results/manual-investigation-samples.json", "w") as fp:
            json.dump(samples, fp, indent=2)


if __name__ == "__main__":
    main()
