"""Tempo からトレースを取り、サービスをまたいだスパンの木を表示する。

  python3 scripts/trace_tree.py <traceId>
"""

import base64
import json
import sys
import urllib.request

TEMPO = "http://localhost:3200"


def hex_id(v: str) -> str:
    """Tempo の JSON は ID を base64 で返すことがある。16進に揃える。"""
    if not v:
        return ""
    if all(c in "0123456789abcdefABCDEF" for c in v) and len(v) in (16, 32):
        return v.lower()
    return base64.b64decode(v).hex()


def main(trace_id: str) -> None:
    with urllib.request.urlopen(f"{TEMPO}/api/traces/{trace_id}") as r:
        data = json.load(r)
    batches = data.get("batches") or data.get("resourceSpans") or data.get("trace", {}).get("resourceSpans", [])
    spans = []
    for b in batches:
        attrs = {a["key"]: a["value"].get("stringValue") for a in b.get("resource", {}).get("attributes", [])}
        service = attrs.get("service.name", "?")
        for ss in b.get("scopeSpans") or b.get("instrumentationLibrarySpans", []):
            for s in ss["spans"]:
                spans.append({
                    "id": hex_id(s["spanId"]), "parent": hex_id(s.get("parentSpanId", "")), "name": s["name"],
                    "service": service, "start": int(s["startTimeUnixNano"]),
                    "ms": (int(s["endTimeUnixNano"]) - int(s["startTimeUnixNano"])) // 1_000_000,
                    "error": s.get("status", {}).get("code") in ("STATUS_CODE_ERROR", 2),
                })
    ids = {s["id"] for s in spans}
    kids: dict = {}
    for s in spans:
        kids.setdefault(s["parent"] if s["parent"] in ids else None, []).append(s)

    def show(pid, depth):
        for s in sorted(kids.get(pid, []), key=lambda x: x["start"]):
            print(f"{'  ' * depth}{s['service']:<18} {s['name']} ({s['ms']} ms){' [ERROR]' if s['error'] else ''}")
            show(s["id"], depth + 1)

    show(None, 0)
    print(f"\n{len(spans)} spans, services: {', '.join(sorted({s['service'] for s in spans}))}")


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