from __future__ import annotations

import argparse
import datetime as dt
import json
import math
import sys
from pathlib import Path
from typing import Any

from okstra_ctl.time_report import (
    _duration_ms,
    _read_state,
    wall_clock_ms,
    load_runs,
)
from okstra_ctl.json_boundary import load_owned_object
from okstra_ctl.fixed_text import line
from okstra_project import ResolverError, list_project_tasks, resolve_project_root

UTC = dt.timezone.utc
UNAVAILABLE_KEYS = (
    "invalidRunTimestamp",
    "teamStateMissing",
    "usageSummaryMissing",
    "durationMissing",
)


def _iso_z(value: dt.datetime) -> str:
    return value.astimezone(UTC).isoformat().replace("+00:00", "Z")


def _parse_timestamp(value: Any) -> dt.datetime | None:
    if not isinstance(value, str) or not value:
        return None
    try:
        parsed = dt.datetime.fromisoformat(value.replace("Z", "+00:00"))
        if parsed.tzinfo is None:
            parsed = parsed.replace(tzinfo=UTC)
        return parsed.astimezone(UTC)
    except (OverflowError, ValueError):
        return None


def _project_id(project_root: Path) -> str:
    try:
        data = load_owned_object(
            project_root / ".okstra" / "project.json", artifact="project config"
        )
    except (OSError, ValueError):
        return ""
    return str(data.get("projectId") or "") if isinstance(data, dict) else ""


def _int_value(value: Any) -> int:
    try:
        return int(value or 0)
    except (OverflowError, TypeError, ValueError):
        return 0


def _float_value(value: Any) -> float:
    try:
        result = float(value or 0)
    except (OverflowError, TypeError, ValueError):
        return 0.0
    return result if math.isfinite(result) else 0.0


def _usage_block(value: Any) -> dict:
    if not isinstance(value, dict):
        return {}
    block = {"source": value["source"]} if "source" in value else {}
    if "durationMs" in value:
        block["durationMs"] = _int_value(value["durationMs"])
    for key in ("startedAt", "endedAt"):
        if isinstance(value.get(key), str):
            block[key] = value[key]
    return block


def _usage_blocks(state: dict) -> tuple[dict, list[dict]]:
    workers = state.get("workers")
    worker_entries = workers if isinstance(workers, list) else []
    return _usage_block(state.get("leadUsage")), [
        _usage_block(worker.get("usage"))
        for worker in worker_entries
        if isinstance(worker, dict)
    ]


def _wall_clock_from_blocks(lead: dict, workers: list[dict]) -> int:
    state = {
        "leadUsage": lead,
        "workers": [{"usage": usage} for usage in workers],
    }
    try:
        return wall_clock_ms(state)
    except TypeError:
        return 0


def _run_metrics(run: dict, project_root: Path) -> tuple[dict | None, str | None]:
    state = _read_state(project_root, str(run.get("teamStatePath") or ""))
    if state is None:
        return None, "teamStateMissing"
    summary = state.get("usageSummary")
    if not isinstance(summary, dict) or _int_value(summary.get("grandTotalTokens")) <= 0:
        return None, "usageSummaryMissing"
    lead, workers = _usage_blocks(state)
    durations = [_duration_ms(lead), *(_duration_ms(worker) for worker in workers)]
    positive_durations = [duration for duration in durations if duration > 0]
    if not positive_durations:
        return None, "durationMissing"
    costs = summary.get("estimatedCostUsd")
    cost = costs.get("grandTotal", 0) if isinstance(costs, dict) else 0
    unmatched = summary.get("unmatchedModels")
    unmatched_models = unmatched if isinstance(unmatched, list) else []
    return {
        "rawTokens": _int_value(summary.get("grandTotalTokens")),
        "billableTokens": _int_value(summary.get("grandBillableEquivalentTokens")),
        "costUsd": _float_value(cost),
        "cpuSumMs": sum(positive_durations),
        "wallClockMs": _wall_clock_from_blocks(lead, workers),
        "unmatchedModels": [
            model for model in unmatched_models if isinstance(model, str) and model
        ],
    }, None


def _new_bucket(task_type: str, first_run_at: dt.datetime) -> dict:
    return {
        "taskType": task_type,
        "firstRunAt": _iso_z(first_run_at),
        "runs": 0,
        "collectedRuns": 0,
        "rawTokens": 0,
        "billableTokens": 0,
        "costUsd": 0.0,
        "cpuSumMs": 0,
        "wallClockMs": 0,
    }


def _collect_usage(
    project_root: Path,
    start_at: dt.datetime,
    end_at: dt.datetime,
    unavailable: dict[str, int],
    warnings: dict[str, int],
) -> tuple[dict[str, dict], set[str], int, int, list[dict]]:
    from okstra_ctl.usage_identity import identity_rows_from_state

    buckets: dict[str, dict] = {}
    unmatched_models: set[str] = set()
    identity_rows: list[dict] = []
    seen_identity: set[str] = set()
    total_runs = 0
    collected_runs = 0
    for entry in list_project_tasks(project_root):
        runs, timeline_warning = load_runs(Path(entry["_resolvedTaskRoot"]))
        warnings["missingOrInvalidTimelines"] += int(timeline_warning)
        for run in runs:
            if not isinstance(run, dict):
                unavailable["invalidRunTimestamp"] += 1
                continue
            run_at = _parse_timestamp(run.get("runTimestamp"))
            if run_at is None:
                unavailable["invalidRunTimestamp"] += 1
                continue
            if run_at < start_at or run_at > end_at:
                continue
            task_type = str(run.get("taskType") or "unknown")
            bucket = buckets.setdefault(task_type, _new_bucket(task_type, run_at))
            run_at_text = _iso_z(run_at)
            if run_at_text < bucket["firstRunAt"]:
                bucket["firstRunAt"] = run_at_text
            total_runs += 1
            bucket["runs"] += 1
            metrics, reason = _run_metrics(run, project_root)
            state = _read_state(project_root, str(run.get("teamStatePath") or ""))
            for row in identity_rows_from_state(state):
                key = str(row.get("roleExecutionRef") or row.get("executionLabel"))
                if key in seen_identity:
                    continue
                seen_identity.add(key)
                identity_rows.append(row)
            if reason is not None:
                unavailable[reason] += 1
                continue
            collected_runs += 1
            bucket["collectedRuns"] += 1
            for key in ("rawTokens", "billableTokens", "cpuSumMs", "wallClockMs"):
                bucket[key] += metrics[key]
            bucket["costUsd"] += metrics["costUsd"]
            unmatched_models.update(metrics["unmatchedModels"])
    return buckets, unmatched_models, total_runs, collected_runs, identity_rows


def build_usage_report(project_root: Path, days: int, now: dt.datetime) -> dict:
    if days <= 0:
        raise ValueError("days must be a positive integer")
    project_root = Path(project_root).resolve()
    end_at = now.astimezone(UTC)
    try:
        start_at = end_at - dt.timedelta(days=days)
    except OverflowError:
        start_at = dt.datetime.min.replace(tzinfo=UTC)
    unavailable = {key: 0 for key in UNAVAILABLE_KEYS}
    warnings = {"missingOrInvalidTimelines": 0}
    buckets, unmatched_models, total_runs, collected_runs, identity_rows = _collect_usage(
        project_root, start_at, end_at, unavailable, warnings
    )

    rows = sorted(buckets.values(), key=lambda row: row["firstRunAt"])
    for row in rows:
        row["costUsd"] = round(_float_value(row["costUsd"]), 4)
        row["collectionRate"] = (
            row["collectedRuns"] / row["runs"] if row["runs"] else 0.0
        )
    totals = {
        "runs": total_runs,
        "collectedRuns": collected_runs,
        "rawTokens": sum(row["rawTokens"] for row in rows),
        "billableTokens": sum(row["billableTokens"] for row in rows),
        "costUsd": round(_float_value(sum(row["costUsd"] for row in rows)), 4),
        "cpuSumMs": sum(row["cpuSumMs"] for row in rows),
        "wallClockMs": sum(row["wallClockMs"] for row in rows),
    }
    return {
        "ok": True,
        "projectId": _project_id(project_root),
        "range": {"days": days, "startAt": _iso_z(start_at), "endAt": _iso_z(end_at)},
        "runs": {
            "total": total_runs,
            "collected": collected_runs,
            "collectionRate": collected_runs / total_runs if total_runs else 0.0,
        },
        "byTaskType": rows,
        "totals": totals,
        "unavailable": unavailable,
        "warnings": warnings,
        "unmatchedModels": sorted(unmatched_models),
        "identityRows": identity_rows,
    }


def utc_now() -> dt.datetime:
    return dt.datetime.now(UTC)


def _positive_days(value: str) -> int:
    try:
        days = int(value)
    except ValueError as exc:
        raise argparse.ArgumentTypeError("--days must be a positive integer") from exc
    if days <= 0:
        raise argparse.ArgumentTypeError("--days must be a positive integer")
    return days


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(
        prog="okstra usage-report",
        description="Aggregate recent project usage by task type.",
    )
    parser.add_argument("--days", type=_positive_days, default=30)
    parser.add_argument("--project-root", default="")
    parser.add_argument("--cwd", default=".")
    parser.add_argument("--json", action="store_true", help="emit JSON (always on)")
    parser.add_argument("--text", action="store_true", help="emit fixed text fields")
    args = parser.parse_args(argv)
    try:
        project_root = Path(
            resolve_project_root(explicit_root=args.project_root, cwd=args.cwd)
        ).resolve()
    except ResolverError as exc:
        payload = {"ok": False, "stage": "resolve", "reason": str(exc)}
        print(render_usage_text(payload) if args.text else json.dumps(payload))
        return 2
    result = build_usage_report(project_root, args.days, utc_now())
    print(
        render_usage_text(result)
        if args.text
        else json.dumps(result, ensure_ascii=False, indent=2)
    )
    return 0


def render_usage_text(payload: dict) -> str:
    """usage 모델 표면에 승인된 집계만 고정 순서로 투영한다."""
    rows = ["Okstra usage\n"]
    rows.append(line("Status", "ready" if payload.get("ok") else "error"))
    rows.append(line("Project ID", payload.get("projectId")))
    period = payload.get("range") if isinstance(payload.get("range"), dict) else {}
    rows.append(line("Range days", period.get("days")))
    runs = payload.get("runs") if isinstance(payload.get("runs"), dict) else {}
    rows.append(line("Runs total", runs.get("total")))
    rows.append(line("Runs collected", runs.get("collected")))
    rows.append(line("Runs collection rate", runs.get("collectionRate")))
    task_types = payload.get("byTaskType") if isinstance(payload.get("byTaskType"), list) else []
    for index, item in enumerate(task_types, 1):
        if not isinstance(item, dict):
            continue
        rows.append(line(f"Task type {index} name", item.get("taskType")))
        rows.append(line(f"Task type {index} runs", item.get("runs")))
        rows.append(line(f"Task type {index} collected runs", item.get("collectedRuns")))
        rows.append(line(f"Task type {index} collection rate", item.get("collectionRate")))
        rows.append(line(f"Task type {index} raw tokens", item.get("rawTokens")))
        rows.append(line(f"Task type {index} billable tokens", item.get("billableTokens")))
        rows.append(line(f"Task type {index} cost USD", item.get("costUsd")))
        rows.append(line(f"Task type {index} CPU sum ms", item.get("cpuSumMs")))
        rows.append(line(f"Task type {index} wall clock ms", item.get("wallClockMs")))
    totals = payload.get("totals") if isinstance(payload.get("totals"), dict) else {}
    for label, key in (("Raw tokens", "rawTokens"), ("Billable tokens", "billableTokens"),
                       ("Cost USD", "costUsd"), ("CPU sum ms", "cpuSumMs"),
                       ("Wall clock ms", "wallClockMs")):
        rows.append(line(label, totals.get(key)))
    for prefix, key in (("Unavailable", "unavailable"), ("Warning", "warnings")):
        values = payload.get(key) if isinstance(payload.get(key), dict) else {}
        for index, (name, count) in enumerate(sorted(values.items()), 1):
            rows.append(line(f"{prefix} {index} reason", name))
            rows.append(line(f"{prefix} {index} count", count))
    models = payload.get("unmatchedModels") if isinstance(payload.get("unmatchedModels"), list) else []
    for index, model in enumerate(models, 1):
        rows.append(line(f"Unmatched model {index} name", model))
    if not payload.get("ok"):
        rows.append(line("Failure stage", payload.get("stage")))
        rows.append(line("Failure reason", payload.get("reason")))
    return "".join(rows)


if __name__ == "__main__":
    raise SystemExit(main(sys.argv[1:]))
