import argparse
import json
import platform
import sys
from pathlib import Path

from .common import ROOT, DEFAULT_DATA, DEFAULT_RUN, device_for, read_json, read_jsonl, write_json


def parser():
    p = argparse.ArgumentParser(prog="poet", description="从一字到一首诗：本地古诗 GPT 学习工具")
    sub = p.add_subparsers(dest="command", required=True)
    sub.add_parser("doctor", help="检查环境、材料和模型")
    prepare = sub.add_parser("prepare", help="整理诗词，生成训练材料")
    prepare.add_argument("--source", type=Path, default=ROOT)
    prepare.add_argument("--output", type=Path, default=DEFAULT_DATA)
    prepare.add_argument("--min-count", type=int, default=2)
    train = sub.add_parser("train", help="从零训练，或继续已有进度")
    train.add_argument("--data", type=Path, default=DEFAULT_DATA)
    train.add_argument("--run-dir", type=Path, default=DEFAULT_RUN)
    train.add_argument("--config", type=Path, default=ROOT/"configs/poet.json")
    train.add_argument("--device", choices=["auto","cpu","mps","cuda"], default="auto")
    train.add_argument("--steps", type=int)
    train.add_argument("--batch-size", type=int)
    train.add_argument("--resume", type=Path)
    for name in ("write", "chat"):
        w = sub.add_parser(name, help="按首字或关键词写诗" if name=="write" else "按中文提示连续写诗")
        w.add_argument("--checkpoint", type=Path, default=DEFAULT_RUN/"best.pt")
        w.add_argument("--data", type=Path, default=DEFAULT_DATA)
        w.add_argument("--device", choices=["auto","cpu","mps","cuda"], default="auto")
        if name == "write":
            w.add_argument("--start", default="")
            w.add_argument("--keywords", default="")
            w.add_argument("--form", choices=["five","seven"], default="five")
            w.add_argument("--count", type=int, default=1)
            w.add_argument("--temperature", type=float, default=0.8)
            w.add_argument("--top-k", type=int, default=40)
            w.add_argument("--candidates", type=int, default=12)
            w.add_argument("--seed", type=int, default=2026)
            w.add_argument("--json", action="store_true")
            w.add_argument("--output", type=Path)
    e = sub.add_parser("evaluate", help="在保留材料上检查预测误差")
    e.add_argument("--checkpoint", type=Path, default=DEFAULT_RUN/"best.pt")
    e.add_argument("--data", type=Path, default=DEFAULT_DATA)
    e.add_argument("--split", choices=["val","test"], default="val")
    e.add_argument("--batches", type=int, default=30)
    e.add_argument("--device", default="auto", choices=["auto","cpu","mps","cuda"])
    e.add_argument("--output", type=Path)
    lesson = sub.add_parser("lesson", help="运行某节课的小实验")
    lesson.add_argument("number", type=int, choices=range(1,27))
    lesson.add_argument("--exercise", action="store_true", help="显示本课练习要求")
    course = sub.add_parser("course", help="显示离线课程入口或启动本地阅读服务")
    mode = course.add_mutually_exclusive_group()
    mode.add_argument("--serve", action="store_true", help="前台运行，关闭终端后停止")
    mode.add_argument("--start", action="store_true", help="后台启动，返回后仍可阅读")
    mode.add_argument("--status", action="store_true", help="查看课程服务是否运行")
    mode.add_argument("--stop", action="store_true", help="停止本项目的课程服务")
    course.add_argument("--port", type=int, default=8766)
    return p


def doctor():
    import torch
    report = {"Python": platform.python_version(), "PyTorch": str(torch.__version__),
              "平台": platform.platform(), "Mac加速可用": torch.backends.mps.is_available(),
              "NVIDIA加速可用": torch.cuda.is_available(), "默认设备": str(device_for()),
              "材料已准备": (DEFAULT_DATA/"manifest.json").exists(),
              "模型已训练": (DEFAULT_RUN/"best.pt").exists(), "课程入口": str(ROOT/"course/dist/index.html")}
    print(json.dumps(report, ensure_ascii=False, indent=2))


def main(argv=None):
    args = parser().parse_args(argv)
    try:
        if args.command == "doctor": doctor()
        elif args.command == "prepare":
            from .data import prepare
            report = prepare(args.source,args.output,args.min_count)
            print(json.dumps({k:report[k] for k in ("source_counts","split_counts","forms","vocab_size","rejections")},ensure_ascii=False,indent=2))
        elif args.command == "train":
            from .train import train
            train(args.data,args.run_dir,args.config,args.device,args.steps,args.batch_size,args.resume)
        elif args.command in ("write","chat"):
            from .generate import Writer, parse_keywords
            writer = Writer(args.checkpoint,args.device,args.data)
            if args.command == "chat":
                print("古诗写作。输入 /quit 退出；输入 /start 春 指定开头；其他文字作为关键词。")
                while True:
                    try: text = input("关键词或开头 > ").strip()
                    except EOFError: break
                    if text == "/quit": break
                    if not text: continue
                    try:
                        start = text[7:].strip() if text.startswith("/start ") else ""
                        words = [] if start else parse_keywords(text)
                        result = writer.write(start=start,keywords=words)
                        print("\n".join(result[0]["lines"]))
                    except ValueError as error: print(f"提示：{error}")
            else:
                result = writer.write(args.start,parse_keywords(args.keywords),5 if args.form=="five" else 7,
                                      args.count,args.temperature,args.top_k,args.candidates,args.seed)
                if args.output: write_json(args.output,result)
                if args.json: print(json.dumps(result,ensure_ascii=False,indent=2))
                else:
                    for i, poem in enumerate(result):
                        if len(result)>1: print(f"\n第 {i+1} 首")
                        print("\n".join(poem["lines"]))
        elif args.command == "evaluate":
            from .train import load_model,evaluate_loss
            if args.batches < 1: raise ValueError("batches 必须大于 0。")
            device = device_for(args.device)
            model,tokenizer,meta = load_model(args.checkpoint,device)
            manifest = read_json(args.data/"manifest.json")
            if meta["data_fingerprint"] != manifest["split_sha256"]: raise ValueError("模型与材料版本不匹配。")
            loss = evaluate_loss(model,read_jsonl(args.data/f"{args.split}.jsonl"),tokenizer,32,args.batches,device)
            report = {"split":args.split,"loss":loss,"batches":args.batches,"step":meta["step"],
                      "note":"固定种子抽样、按有效字数加权；不是逐篇穷举整卷。"}
            if args.output: write_json(args.output,report)
            print(json.dumps(report,ensure_ascii=False,indent=2))
        elif args.command == "lesson":
            from .labs import run
            run(args.number,args.exercise)
        elif args.command == "course":
            from . import course_server
            if not 1024 <= args.port <= 65535:
                raise ValueError("课程端口请使用 1024—65535 之间的整数。")
            if args.serve:
                print("前台服务运行中，按 Ctrl+C 停止。",flush=True)
                course_server.serve(args.port)
            elif args.start:
                course_server.start(args.port)
                course_server.status(args.port)
            elif args.status: course_server.status(args.port)
            elif args.stop: course_server.stop(args.port)
            else: print(ROOT/"course/dist/index.html")
    except (ValueError, FileNotFoundError, ModuleNotFoundError) as error:
        print(f"提示：{error}", file=sys.stderr)
        return_code=2
        raise SystemExit(return_code)
    except KeyboardInterrupt:
        print("\n已结束。",file=sys.stderr)
        raise SystemExit(130)
