"""Beginner-friendly deterministic gesture classifier. Standard library only; no hardware I/O."""
import argparse, hashlib, json
from pathlib import Path

LABELS = ("rock", "paper", "scissors")
FEATURES = ("foregroundRatio", "topSegmentsNorm", "contrastNorm")
THRESHOLD = 0.65

def canonical_json(value):
    return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode("utf-8")

def sha256(value):
    return hashlib.sha256(value).hexdigest().upper()

def validate_dataset(dataset):
    if not isinstance(dataset, dict) or dataset.get("schemaVersion") != "0.1.0" or dataset.get("sourceOrigin") != "SYNTHETIC":
        raise ValueError("DATASET_HEADER_INVALID")
    if dataset.get("featureNames") != list(FEATURES) or dataset.get("splitPolicy") != "GROUP_EXCLUSIVE_PREDEFINED":
        raise ValueError("DATASET_SCHEMA_INVALID")
    samples = dataset.get("samples")
    if not isinstance(samples, list) or not samples:
        raise ValueError("DATASET_SAMPLES_INVALID")
    ids, groups = set(), {}
    for row in samples:
        if not isinstance(row, dict) or set(row) != {"sampleId", "groupId", "split", "label", "features"}:
            raise ValueError("DATASET_ROW_INVALID")
        if row["sampleId"] in ids or row["split"] not in {"train", "validation", "test"} or row["label"] not in LABELS:
            raise ValueError("DATASET_ROW_INVALID")
        ids.add(row["sampleId"])
        if not isinstance(row["features"], list) or len(row["features"]) != 3 or any(not isinstance(v, (int, float)) or not 0 <= v <= 1 for v in row["features"]):
            raise ValueError("DATASET_FEATURE_INVALID")
        groups.setdefault(row["groupId"], set()).add(row["split"])
    if any(len(splits) != 1 for splits in groups.values()):
        raise ValueError("DATASET_GROUP_LEAKAGE")
    for split in ("train", "validation", "test"):
        for label in LABELS:
            if not any(row["split"] == split and row["label"] == label for row in samples):
                raise ValueError("DATASET_CLASS_MISSING")
    return samples

def train(samples):
    result = {}
    for label in LABELS:
        rows = [row["features"] for row in samples if row["split"] == "train" and row["label"] == label]
        result[label] = [round(sum(float(row[index]) for row in rows) / len(rows), 6) for index in range(3)]
    return result

def predict(centroids, features):
    weights = {}
    for label in LABELS:
        distance = sum((float(features[index]) - centroids[label][index]) ** 2 for index in range(3)) ** .5
        weights[label] = 1.0 / (distance + .05)
    total = sum(weights.values())
    scores = {label: round(weights[label] / total, 6) for label in LABELS}
    candidate = max(LABELS, key=lambda item: weights[item])
    confidence = round(weights[candidate] / total, 6)
    return {"label": candidate if confidence >= THRESHOLD else None, "candidateLabel": candidate,
            "confidence": confidence, "scores": scores,
            "quality": "GOOD" if confidence >= THRESHOLD else "MISSING"}

def evaluate(samples, centroids, split):
    rows = [row for row in samples if row["split"] == split]
    matrix = [[0, 0, 0] for _ in LABELS]
    correct = accepted = unknown = 0
    for row in rows:
        prediction = predict(centroids, row["features"])
        if prediction["label"] is None:
            unknown += 1
            continue
        accepted += 1
        actual_index, predicted_index = LABELS.index(row["label"]), LABELS.index(prediction["label"])
        matrix[actual_index][predicted_index] += 1
        correct += int(actual_index == predicted_index)
    recalls = {}
    for index, label in enumerate(LABELS):
        total = sum(1 for row in rows if row["label"] == label)
        recalls[label] = round(matrix[index][index] / total, 4)
    return {"split": split, "sampleCount": len(rows), "acceptedCount": accepted, "unknownCount": unknown,
            "accuracy": round(correct / len(rows), 4), "coverage": round(accepted / len(rows), 4),
            "confusionMatrix": matrix, "recallByClass": recalls}

def build_manifest(dataset):
    samples = validate_dataset(dataset)
    centroids = train(samples)
    core = {"schemaVersion": "0.1.0", "artifactType": "TEACHING_NEAREST_CENTROID_MODEL",
            "modelVersion": "rps-centroid-v1", "algorithm": "NEAREST_CENTROID_NORMALIZED_V1",
            "labels": list(LABELS), "featureNames": list(FEATURES), "threshold": THRESHOLD,
            "datasetId": dataset["datasetId"], "datasetSha256": sha256(canonical_json(dataset)),
            "sourceOrigin": "SYNTHETIC", "splitPolicy": dataset["splitPolicy"], "groupOverlap": [],
            "centroids": centroids, "validation": evaluate(samples, centroids, "validation"),
            "test": evaluate(samples, centroids, "test"), "trainingExecuted": True,
            "trainingRuntime": "PYTHON_STDLIB_DETERMINISTIC", "knownLimitations": dataset["limitations"],
            "realModelClaimed": False, "actionEligible": False, "hardwareAccessed": False,
            "pcanInitialized": False, "canFramesSent": 0, "decision": "NO_GO"}
    return {**core, "modelArtifactId": sha256(canonical_json(core))}

def main():
    root = Path(__file__).resolve().parents[1]
    parser = argparse.ArgumentParser(description="Train the offline teaching gesture classifier")
    parser.add_argument("--data", type=Path, default=root / "data" / "synthetic-rps-features.json")
    parser.add_argument("--output", type=Path, required=True, help="Write generated manifest to an explicit local path")
    args = parser.parse_args()
    dataset = json.loads(args.data.read_text(encoding="utf-8"))
    manifest = build_manifest(dataset)
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_bytes(canonical_json(manifest))
    print(json.dumps({"modelArtifactId": manifest["modelArtifactId"], "test": manifest["test"],
                      "actionEligible": False, "decision": "NO_GO"}, ensure_ascii=False, sort_keys=True))

if __name__ == "__main__":
    main()
