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
from okstra_token_usage.blocks import accounting_workers, normalize_usage_block, usage_blocks

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]]:
    worker_entries = accounting_workers(state)
    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 []
    original_blocks = usage_blocks(state)
    normalized = [normalize_usage_block(block) for block in original_blocks]
    cache_adjustment = sum(_int_value(block.get("totalTokens")) for block in original_blocks)
    cache_adjustment -= sum(_int_value(block.get("totalTokens")) for block in normalized)
    return {
        "rawTokens": _int_value(summary.get("grandTotalTokens")) - cache_adjustment,
        "reportedRawTokens": _int_value(summary.get("grandTotalTokens")),
        "legacyCacheReadAdjustmentTokens": cache_adjustment,
        "cacheReadTokens": sum(_int_value(block.get("cacheReadTokens")) for block in normalized),
        "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,
        "reportedRawTokens": 0,
        "legacyCacheReadAdjustmentTokens": 0,
        "cacheReadTokens": 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:
            run_at = _parse_timestamp(run.get("runTimestamp") if isinstance(run, dict) else None)
            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", "reportedRawTokens", "legacyCacheReadAdjustmentTokens",
                        "cacheReadTokens", "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),
        "reportedRawTokens": sum(row["reportedRawTokens"] for row in rows),
        "legacyCacheReadAdjustmentTokens": sum(row["legacyCacheReadAdjustmentTokens"] for row in rows),
        "cacheReadTokens": sum(row["cacheReadTokens"] 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


_CLI_EPILOG = r"""Usage:
  okstra usage-report [--days <positive-int>] [--project-root <dir>] [--cwd <dir>] [--json]

Read-only. Prints JSON containing run coverage, raw and billable-equivalent
tokens, known USD cost, CPU-sum time, and wall-clock time grouped by task type.
The default window is the current project's last 30 days.
"""


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(
        epilog=_CLI_EPILOG,
        formatter_class=argparse.RawDescriptionHelpFormatter,
        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"),
                       ("Reported raw tokens", "reportedRawTokens"),
                       ("Legacy cache read adjustment tokens", "legacyCacheReadAdjustmentTokens"),
                       ("Cache read tokens", "cacheReadTokens"),
                       ("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:]))
