import argparse
import csv
import random
from pathlib import Path

import numpy as np
import sklearn
from sklearn.dummy import DummyRegressor
from sklearn.linear_model import Ridge
from sklearn.metrics import mean_absolute_error
from sklearn.model_selection import GroupShuffleSplit
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

AA = "ACDEFGHIKLMNPQRSTVWY"


def demo_rows():
    rng = random.Random(17)
    rows = []
    for family in range(60):
        length = rng.randrange(60, 101)
        weights = [rng.uniform(1, 15) if a == "A" else 1 for a in AA]
        parent = "".join(rng.choices(AA, weights=weights, k=length))
        for variant in range(3):
            seq = parent[:-3] + "".join(rng.choices(AA, k=3))
            # Artificial target designed to depend on composition and length.
            target = 4 * seq.count("A") / len(seq) + 0.01 * len(seq)
            target += rng.gauss(0, 0.03)
            rows.append(dict(id=f"f{family}_v{variant}", sequence=seq,
                             group=f"f{family}", target=target))
    return rows


def prepare(rows):
    if not rows:
        raise ValueError("No data rows")
    ids, sequences, groups, targets = [], [], [], []
    for row in rows:
        ident, group = str(row["id"]).strip(), str(row["group"]).strip()
        seq = str(row["sequence"]).strip().upper()
        target = float(row["target"])
        if not ident or not group or not seq or set(seq) - set(AA):
            raise ValueError("Each row needs id, group, and a canonical sequence")
        if not np.isfinite(target):
            raise ValueError("Targets must be finite numbers")
        ids.append(ident)
        sequences.append(seq)
        groups.append(group)
        targets.append(target)
    if len(set(ids)) != len(ids):
        raise ValueError("IDs must be unique")
    if len(set(groups)) < 4:
        raise ValueError("Provide at least four groups for this demo workflow")
    X = np.array([[s.count(a) / len(s) for a in AA] + [len(s)]
                  for s in sequences], dtype=float)
    return X, np.array(targets), np.array(groups), ids, sequences


def main():
    parser = argparse.ArgumentParser(description="Grouped composition baseline")
    parser.add_argument("--csv", type=Path, help="id,sequence,group,target columns")
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--split-out", type=Path, default=Path("split.csv"))
    args = parser.parse_args()
    if args.csv:
        with args.csv.open(newline="", encoding="utf-8-sig") as handle:
            reader = csv.DictReader(handle)
            required = {"id", "sequence", "group", "target"}
            if not required <= set(reader.fieldnames or []):
                raise ValueError("Missing CSV columns: id,sequence,group,target")
            rows = list(reader)
    else:
        rows = demo_rows()
    X, y, groups, ids, sequences = prepare(rows)
    splitter = GroupShuffleSplit(n_splits=1, test_size=0.25,
                                 random_state=args.seed)
    train, test = next(splitter.split(X, y, groups))
    overlap = set(groups[train]) & set(groups[test])
    if overlap:
        raise ValueError("Group overlap")
    if {sequences[i] for i in train} & {sequences[i] for i in test}:
        raise ValueError("Identical sequences cross the split; revise groups")
    # Exclusive creation protects an existing split record from replacement.
    with args.split_out.open("x", newline="", encoding="utf-8") as handle:
        writer = csv.writer(handle)
        writer.writerow(["id", "group", "split"])
        for label, indices in [("train", train), ("test", test)]:
            writer.writerows((ids[i], groups[i], label) for i in indices)
    print(f"scikit-learn={sklearn.__version__}; synthetic={args.csv is None}")
    print(f"rows: train={len(train)} test={len(test)}; features={X.shape[1]}")
    print(f"groups: train={len(set(groups[train]))} "
          f"test={len(set(groups[test]))}; overlap={len(overlap)}")
    models = {
        "median": DummyRegressor(strategy="median"),
        "composition_ridge": make_pipeline(StandardScaler(), Ridge(alpha=1.0)),
    }
    for name, model in models.items():
        model.fit(X[train], y[train])
        error = mean_absolute_error(y[test], model.predict(X[test]))
        print(f"{name}: MAE={error:.4f}")
    print(f"saved split: {args.split_out}")


if __name__ == "__main__":
    main()
