#!/usr/bin/env python3
from __future__ import annotations

import csv
import hashlib
import json
import sys
from collections import Counter, defaultdict
from pathlib import Path
from typing import Iterable

ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from scripts.legacy_dump_assessment.keepass_analyzer import (
    extract_keepass_hash,
    parse_kdbx_header,
)

FINDINGS_DIR = ROOT / "findings"
DATA_DIR = FINDINGS_DIR / "data"
LEGACY_DUMPS_DIR = FINDINGS_DIR / "legacy_dump_assessment" / "dumps"
CURRENT_DUMPS_DIR = FINDINGS_DIR / "dumps"
CURRENT_KEEPASS_DIR = FINDINGS_DIR / "keepass"

LEGACY_ANALYSIS_JSON = DATA_DIR / "legacy_keepass_analysis.json"
CURRENT_HASHES_TXT = DATA_DIR / "kdbx_hashes.txt"

OUT_INVENTORY_CSV = DATA_DIR / "kdbx_inventory.csv"
OUT_SUMMARY_JSON = DATA_DIR / "kdbx_inventory_summary.json"
OUT_BUCKET_DIR = DATA_DIR / "kdbx_hash_buckets"

EMPTY_THRESHOLD = 1500

CSV_FIELDS = [
    "inventory_id",
    "scope",
    "source_namespace",
    "source_id",
    "host_name",
    "username",
    "relative_path",
    "filename",
    "size_bytes",
    "sha256",
    "keepass_label",
    "kdbx_version",
    "kdbx_major",
    "cipher",
    "kdf",
    "kdf_cost",
    "kdf_cost_unit",
    "argon2_memory",
    "argon2_parallelism",
    "header_size",
    "hashcat_mode",
    "bucket",
    "is_non_empty_estimate",
    "password_only",
    "key_file_used",
    "corporate_classification",
    "corporate_score",
    "priority_score",
    "is_duplicate",
    "duplicate_of",
    "config_paths",
    "cracked_status",
    "keepass_hash",
    "notes",
]


def sha256_file(path: Path) -> str:
    return hashlib.sha256(path.read_bytes()).hexdigest()


def load_current_hash_catalog() -> dict[str, str]:
    mapping: dict[str, str] = {}
    if not CURRENT_HASHES_TXT.exists():
        return mapping
    for raw_line in CURRENT_HASHES_TXT.read_text(encoding="utf-8").splitlines():
        line = raw_line.strip()
        if not line or ":" not in line:
            continue
        label, keepass_hash = line.split(":", 1)
        mapping[label] = keepass_hash
    return mapping


def parse_hash_metadata(keepass_hash: str) -> dict[str, object]:
    parts = keepass_hash.split("*")
    cost = int(parts[2]) if len(parts) > 2 and parts[2].isdigit() else 0
    field4 = int(parts[3]) if len(parts) > 3 and parts[3].isdigit() else 0
    is_argon2 = field4 >= 1_000_000

    if is_argon2:
        memory = field4
        parallelism = int(parts[4]) if len(parts) > 4 and parts[4].isdigit() else 0
        if cost <= 100:
            bucket = "kdbx4_argon2_easy"
        elif cost <= 300:
            bucket = "kdbx4_argon2_medium"
        else:
            bucket = "kdbx4_argon2_hard"
        return {
            "hashcat_mode": 29700,
            "bucket": bucket,
            "kdf_cost": cost,
            "kdf_cost_unit": "iterations",
            "argon2_memory": memory,
            "argon2_parallelism": parallelism,
        }

    if cost <= 10_000:
        bucket = "kdbx3_aes_6000"
    elif cost <= 1_000_000:
        bucket = "kdbx3_aes_600k"
    else:
        bucket = "kdbx3_aes_48m"
    return {
        "hashcat_mode": 13400,
        "bucket": bucket,
        "kdf_cost": cost,
        "kdf_cost_unit": "rounds",
        "argon2_memory": "",
        "argon2_parallelism": "",
    }


def dedupe_records(records: Iterable[dict[str, object]]) -> list[dict[str, object]]:
    seen: set[tuple[object, ...]] = set()
    deduped: list[dict[str, object]] = []
    for record in records:
        key = (
            record["scope"],
            record["source_namespace"],
            record["source_id"],
            record["relative_path"],
            record["keepass_hash"],
        )
        if key in seen:
            continue
        seen.add(key)
        deduped.append(record)
    return deduped


def current_cracked_status(relative_path: str) -> str:
    if relative_path == "findings/keepass/multi.kdbx":
        return "cracked"
    if relative_path == "findings/keepass/Passwords.kdbx":
        return "pending"
    return "uncracked"


def build_current_record(
    path: Path,
    *,
    scope: str,
    source_namespace: str,
    source_id: str,
    host_name: str,
    username: str,
    current_hash_catalog: dict[str, str],
) -> dict[str, object] | None:
    header = parse_kdbx_header(path)
    keepass_hash = extract_keepass_hash(path)
    if not header or not keepass_hash:
        return None

    relative_path = path.relative_to(ROOT).as_posix()
    hash_meta = parse_hash_metadata(keepass_hash)
    historical_hash = current_hash_catalog.get(source_id)
    label = source_id if source_id in current_hash_catalog else ""
    if not label:
        label = path.stem

    inventory_id = f"{source_namespace}:{source_id}:{path.name}:{sha256_file(path)[:12]}"
    notes = []
    if source_namespace == "current_dumps":
        notes.append("current dump lane")
    if source_namespace == "current_keepass":
        notes.append("standalone current keepass lane")
    if historical_hash:
        if historical_hash == keepass_hash:
            notes.append("matches findings/data/kdbx_hashes.txt")
        else:
            notes.append("differs from findings/data/kdbx_hashes.txt")

    return {
        "inventory_id": inventory_id,
        "scope": scope,
        "source_namespace": source_namespace,
        "source_id": source_id,
        "host_name": host_name,
        "username": username,
        "relative_path": relative_path,
        "filename": path.name,
        "size_bytes": path.stat().st_size,
        "sha256": sha256_file(path),
        "keepass_label": label,
        "kdbx_version": header.get("version", ""),
        "kdbx_major": header.get("major", ""),
        "cipher": header.get("cipher", ""),
        "kdf": header.get("kdf", ""),
        "kdf_cost": hash_meta["kdf_cost"],
        "kdf_cost_unit": hash_meta["kdf_cost_unit"],
        "argon2_memory": hash_meta["argon2_memory"],
        "argon2_parallelism": hash_meta["argon2_parallelism"],
        "header_size": header.get("header_size", ""),
        "hashcat_mode": hash_meta["hashcat_mode"],
        "bucket": hash_meta["bucket"],
        "is_non_empty_estimate": path.stat().st_size > EMPTY_THRESHOLD,
        "password_only": "",
        "key_file_used": "",
        "corporate_classification": "",
        "corporate_score": "",
        "priority_score": "",
        "is_duplicate": False,
        "duplicate_of": "",
        "config_paths": "",
        "cracked_status": current_cracked_status(relative_path),
        "keepass_hash": keepass_hash,
        "notes": "; ".join(notes),
    }


def load_current_records() -> list[dict[str, object]]:
    current_hash_catalog = load_current_hash_catalog()
    records: list[dict[str, object]] = []

    if CURRENT_DUMPS_DIR.exists():
        for path in sorted(CURRENT_DUMPS_DIR.rglob("*.kdbx")):
            dump_id = path.parts[-3] if len(path.parts) >= 3 else path.parent.name
            record = build_current_record(
                path,
                scope="current",
                source_namespace="current_dumps",
                source_id=dump_id,
                host_name=dump_id,
                username="",
                current_hash_catalog=current_hash_catalog,
            )
            if record:
                records.append(record)

    if CURRENT_KEEPASS_DIR.exists():
        for path in sorted(CURRENT_KEEPASS_DIR.rglob("*.kdbx")):
            record = build_current_record(
                path,
                scope="current",
                source_namespace="current_keepass",
                source_id=path.stem,
                host_name="",
                username="",
                current_hash_catalog=current_hash_catalog,
            )
            if record:
                records.append(record)

    return records


def load_legacy_records() -> list[dict[str, object]]:
    if not LEGACY_ANALYSIS_JSON.exists():
        return []

    data = json.loads(LEGACY_ANALYSIS_JSON.read_text(encoding="utf-8"))
    records: list[dict[str, object]] = []

    for host in data.get("hosts", []):
        dump_id = host.get("dump_id")
        dump_folder = "dump" if dump_id == 0 else f"dump ({dump_id})"
        config_paths = " | ".join(host.get("db_paths_from_config", []))
        for db in host.get("databases", []):
            keepass_hash = db.get("keepass_hash", "")
            if not keepass_hash:
                continue
            hash_meta = parse_hash_metadata(keepass_hash)
            rel_path = (
                Path("findings")
                / "legacy_dump_assessment"
                / "dumps"
                / dump_folder
                / db.get("file", "")
            ).as_posix()
            records.append(
                {
                    "inventory_id": f"legacy:{dump_id}:{db.get('file', '')}",
                    "scope": "legacy",
                    "source_namespace": "legacy_dump_assessment",
                    "source_id": str(dump_id),
                    "host_name": host.get("computer_name", ""),
                    "username": host.get("username", ""),
                    "relative_path": rel_path,
                    "filename": db.get("filename", ""),
                    "size_bytes": db.get("size_bytes", ""),
                    "sha256": db.get("hash_sha256", ""),
                    "keepass_label": f"{host.get('computer_name', '')}_{host.get('username', '')}_{db.get('filename', '')}",
                    "kdbx_version": db.get("header", {}).get("version", ""),
                    "kdbx_major": db.get("header", {}).get("major", ""),
                    "cipher": db.get("header", {}).get("cipher", ""),
                    "kdf": db.get("header", {}).get("kdf", ""),
                    "kdf_cost": hash_meta["kdf_cost"],
                    "kdf_cost_unit": hash_meta["kdf_cost_unit"],
                    "argon2_memory": hash_meta["argon2_memory"],
                    "argon2_parallelism": hash_meta["argon2_parallelism"],
                    "header_size": db.get("header", {}).get("header_size", ""),
                    "hashcat_mode": hash_meta["hashcat_mode"],
                    "bucket": hash_meta["bucket"],
                    "is_non_empty_estimate": db.get("size_bytes", 0) > EMPTY_THRESHOLD,
                    "password_only": host.get("password_only", ""),
                    "key_file_used": host.get("key_file_used", ""),
                    "corporate_classification": host.get("corporate_classification", ""),
                    "corporate_score": host.get("corporate_score", ""),
                    "priority_score": host.get("priority_score", ""),
                    "is_duplicate": db.get("is_duplicate", False),
                    "duplicate_of": db.get("duplicate_of") or "",
                    "config_paths": config_paths,
                    "cracked_status": "unknown",
                    "keepass_hash": keepass_hash,
                    "notes": "imported from findings/data/legacy_keepass_analysis.json",
                }
            )

    return records


def write_inventory(records: list[dict[str, object]]) -> None:
    OUT_INVENTORY_CSV.parent.mkdir(parents=True, exist_ok=True)
    with OUT_INVENTORY_CSV.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=CSV_FIELDS)
        writer.writeheader()
        writer.writerows(records)


def write_bucket_files(records: list[dict[str, object]]) -> dict[str, int]:
    OUT_BUCKET_DIR.mkdir(parents=True, exist_ok=True)
    buckets: dict[str, list[dict[str, object]]] = defaultdict(list)
    for record in records:
        bucket = str(record["bucket"])
        if bucket:
            buckets[bucket].append(record)

    counts: dict[str, int] = {}
    ordered = [
        "kdbx3_aes_6000",
        "kdbx3_aes_600k",
        "kdbx3_aes_48m",
        "kdbx4_argon2_easy",
        "kdbx4_argon2_medium",
        "kdbx4_argon2_hard",
    ]
    for bucket in ordered:
        bucket_records = buckets.get(bucket, [])
        counts[bucket] = len(bucket_records)
        with (OUT_BUCKET_DIR / f"{bucket}.txt").open("w", encoding="utf-8") as handle:
            for record in bucket_records:
                handle.write(f"{record['keepass_hash']}\n")

    with (OUT_BUCKET_DIR / "all_hashes.txt").open("w", encoding="utf-8") as handle:
        for record in records:
            handle.write(f"{record['keepass_hash']}\n")

    with (OUT_BUCKET_DIR / "all_hashes_labeled.txt").open("w", encoding="utf-8") as handle:
        for record in records:
            handle.write(f"{record['keepass_label']}:{record['keepass_hash']}\n")

    return counts


def build_summary(records: list[dict[str, object]], bucket_counts: dict[str, int]) -> dict[str, object]:
    current_records = [r for r in records if r["scope"] == "current"]
    legacy_records = [r for r in records if r["scope"] == "legacy"]
    kdf_counts = Counter(str(r["kdf"]) for r in records)
    namespace_counts = Counter(str(r["source_namespace"]) for r in records)
    cracked_counts = Counter(str(r["cracked_status"]) for r in records)

    summary = {
        "generated_from": [
            "findings/dumps/**/*.kdbx",
            "findings/keepass/**/*.kdbx",
            "findings/data/legacy_keepass_analysis.json",
        ],
        "summary": {
            "total_records": len(records),
            "current_records": len(current_records),
            "legacy_records": len(legacy_records),
            "non_empty_estimate": sum(1 for r in records if r["is_non_empty_estimate"]),
            "hashcat_mode_13400": sum(1 for r in records if r["hashcat_mode"] == 13400),
            "hashcat_mode_29700": sum(1 for r in records if r["hashcat_mode"] == 29700),
        },
        "by_source_namespace": dict(namespace_counts),
        "by_kdf": dict(kdf_counts),
        "by_bucket": bucket_counts,
        "by_cracked_status": dict(cracked_counts),
        "current_labels_seen": sorted(
            [
                str(r["keepass_label"])
                for r in current_records
                if "matches findings/data/kdbx_hashes.txt" in str(r["notes"])
            ]
        ),
        "current_hash_mismatches": sorted(
            [
                str(r["relative_path"])
                for r in current_records
                if "differs from findings/data/kdbx_hashes.txt" in str(r["notes"])
            ]
        ),
    }
    return summary


def write_summary(summary: dict[str, object]) -> None:
    OUT_SUMMARY_JSON.write_text(json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")


def main() -> int:
    current_records = load_current_records()
    legacy_records = load_legacy_records()
    records = dedupe_records([*current_records, *legacy_records])
    records.sort(key=lambda r: (str(r["scope"]), str(r["source_namespace"]), str(r["source_id"]), str(r["relative_path"])))

    write_inventory(records)
    bucket_counts = write_bucket_files(records)
    summary = build_summary(records, bucket_counts)
    write_summary(summary)

    print(f"Wrote {OUT_INVENTORY_CSV.relative_to(ROOT)} ({len(records)} rows)")
    print(f"Wrote {OUT_SUMMARY_JSON.relative_to(ROOT)}")
    print(f"Wrote {OUT_BUCKET_DIR.relative_to(ROOT)}/")
    for bucket, count in bucket_counts.items():
        print(f"  {bucket}: {count}")
    return 0


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