import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

import pandas as pd
import yaml

from common.paths import resolve_path


REQUIRED_COLUMNS = [
    "image_id",
    "image_path",
    "meter_value",
    "label_text",
    "label_id",
    "store_id",
    "shoot_date",
    "batch_id",
    "steam_level",
    "memo",
]


def load_label_rules(config_path: str | Path = "configs/density_classifier.yaml") -> list[dict]:
    with resolve_path(config_path).open("r", encoding="utf-8") as f:
        return yaml.safe_load(f)["label_rules"]


def label_from_meter_value(value: float, rules: list[dict] | None = None) -> tuple[str, int]:
    rules = rules or load_label_rules()
    for rule in rules:
        min_value = rule.get("min_value")
        max_value = rule.get("max_value")
        include_min = bool(rule.get("include_min", True))
        include_max = bool(rule.get("include_max", True))
        lower_ok = True if min_value is None else value >= min_value if include_min else value > min_value
        upper_ok = True if max_value is None else value <= max_value if include_max else value < max_value
        if lower_ok and upper_ok:
            return str(rule["label_text"]), int(rule["label_id"])
    raise ValueError(f"meter_value does not match any label rule: {value}")


def validate_manifest(manifest: str | Path, check_labels: bool = True) -> list[str]:
    path = resolve_path(manifest)
    df = pd.read_csv(path)
    errors: list[str] = []
    missing = [col for col in REQUIRED_COLUMNS if col not in df.columns]
    if missing:
        errors.append(f"missing required columns: {missing}")
        return errors

    if df["image_id"].duplicated().any():
        errors.append("image_id contains duplicates")

    if check_labels:
        rules = load_label_rules()
        for row_number, row in df.iterrows():
            expected_text, expected_id = label_from_meter_value(float(row["meter_value"]), rules)
            if str(row["label_text"]) != expected_text or int(row["label_id"]) != expected_id:
                errors.append(
                    f"row {row_number + 2}: expected ({expected_text}, {expected_id}) "
                    f"but got ({row['label_text']}, {row['label_id']})"
                )
    return errors


def main() -> int:
    parser = argparse.ArgumentParser(description="Validate soup density manifest CSV.")
    parser.add_argument("--manifest", required=True)
    parser.add_argument("--skip-label-check", action="store_true")
    args = parser.parse_args()
    errors = validate_manifest(args.manifest, check_labels=not args.skip_label_check)
    if errors:
        for error in errors:
            print(f"ERROR: {error}")
        return 1
    print("Manifest validation passed.")
    return 0


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