#!/usr/bin/env python3
"""
explore_segment.py
==================

探索 Kafka Topic-Partition 目录里所有 segment 文件，打印：
  - 每个 segment 的 base offset / .log 大小 / .index 项数 / .timeindex 项数 / .snapshot 是否存在
  - leader-epoch-checkpoint 内容
  - 所有文件的总占用
  - 演示「按 offset 找 segment」的二分定位

不依赖 Kafka 服务端，仅靠文件名 + 文件大小解析。

用法：
    python3 explore_segment.py --part-dir /var/lib/kafka/data/learn.07.storage-0
    python3 explore_segment.py --part-dir /var/lib/kafka/data/learn.07.storage-0 --lookup 7050

也可以与 init.sh 配合：
    bash init.sh
    python3 code/explore_segment.py \
        --part-dir /var/lib/kafka/data/learn.07.storage-0 --lookup 5000
"""

from __future__ import annotations

import argparse
import os
import struct
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import List, Optional


# ---------------------------------------------------------------------------
# 常量
# ---------------------------------------------------------------------------

# .index 每项 8 字节： relativeOffset (int32) + position (int32)
INDEX_ENTRY_SIZE = 8
# .timeindex 每项 12 字节： timestamp (int64) + relativeOffset (int32)
TIMEINDEX_ENTRY_SIZE = 12


# ---------------------------------------------------------------------------
# Segment 描述
# ---------------------------------------------------------------------------

@dataclass
class Segment:
    base_offset: int
    log_path: Path
    index_path: Optional[Path]
    timeindex_path: Optional[Path]
    snapshot_path: Optional[Path]

    @property
    def log_size(self) -> int:
        return self.log_path.stat().st_size if self.log_path.exists() else 0

    @property
    def index_entries(self) -> int:
        if self.index_path and self.index_path.exists():
            return self.index_path.stat().st_size // INDEX_ENTRY_SIZE
        return 0

    @property
    def timeindex_entries(self) -> int:
        if self.timeindex_path and self.timeindex_path.exists():
            return self.timeindex_path.stat().st_size // TIMEINDEX_ENTRY_SIZE
        return 0

    @property
    def has_snapshot(self) -> bool:
        return self.snapshot_path is not None and self.snapshot_path.exists()


# ---------------------------------------------------------------------------
# 扫描目录
# ---------------------------------------------------------------------------

def scan_segments(part_dir: Path) -> List[Segment]:
    if not part_dir.exists():
        raise FileNotFoundError(f"分区目录不存在: {part_dir}")
    log_files = sorted(part_dir.glob("*.log"))
    segs: List[Segment] = []
    for log in log_files:
        base_str = log.stem  # 文件名（去后缀）
        try:
            base = int(base_str)
        except ValueError:
            continue   # 不是 segment 命名（例如 leader-epoch-checkpoint）
        index = part_dir / f"{base_str}.index"
        timeindex = part_dir / f"{base_str}.timeindex"
        snapshot = part_dir / f"{base_str}.snapshot"
        segs.append(Segment(
            base_offset=base,
            log_path=log,
            index_path=index if index.exists() else None,
            timeindex_path=timeindex if timeindex.exists() else None,
            snapshot_path=snapshot if snapshot.exists() else None,
        ))
    return segs


# ---------------------------------------------------------------------------
# 打印 Segment 概要
# ---------------------------------------------------------------------------

def print_summary(part_dir: Path, segs: List[Segment]) -> None:
    print(f"\n📁 {part_dir}")
    print(f"   共发现 {len(segs)} 个 segment\n")
    if not segs:
        return

    print(f"  {'idx':<4} {'base offset':>13}  {'.log':>10}  {'.index':>8}  "
          f"{'.timeindex':>10}  {'snapshot':>9}  {'state':>8}")
    print("  " + "-" * 78)
    for i, s in enumerate(segs):
        is_active = (i == len(segs) - 1)
        state = "ACTIVE" if is_active else "SEALED"
        print(
            f"  {i:<4} {s.base_offset:>13}  "
            f"{human(s.log_size):>10}  "
            f"{s.index_entries:>8}  "
            f"{s.timeindex_entries:>10}  "
            f"{'✓' if s.has_snapshot else '·':>9}  "
            f"{state:>8}"
        )

    total_log = sum(s.log_size for s in segs)
    total_idx = sum((s.index_path.stat().st_size if s.index_path else 0) for s in segs)
    total_tidx = sum((s.timeindex_path.stat().st_size if s.timeindex_path else 0) for s in segs)
    print()
    print(f"  📊 .log 总大小      : {human(total_log)}")
    print(f"  📊 .index 总大小    : {human(total_idx)}  ({total_idx/total_log*100:.3f}% of .log)" if total_log else "")
    print(f"  📊 .timeindex 总大小: {human(total_tidx)}")


# ---------------------------------------------------------------------------
# 解析 leader-epoch-checkpoint
# ---------------------------------------------------------------------------

def print_leader_epoch(part_dir: Path) -> None:
    f = part_dir / "leader-epoch-checkpoint"
    if not f.exists():
        return
    print(f"\n📜 leader-epoch-checkpoint （{f}）")
    lines = f.read_text().strip().splitlines()
    if len(lines) < 2:
        print("  (文件为空或格式异常)")
        return
    version = lines[0]
    count = lines[1]
    print(f"  version = {version}, entries = {count}")
    print(f"  {'epoch':>6} {'start_offset':>14}")
    for line in lines[2:]:
        parts = line.split()
        if len(parts) == 2:
            print(f"  {parts[0]:>6} {parts[1]:>14}")


# ---------------------------------------------------------------------------
# 解析 .index 文件，打印前 N 项
# ---------------------------------------------------------------------------

def print_index_head(seg: Segment, n: int = 10) -> None:
    if not seg.index_path or not seg.index_path.exists():
        return
    print(f"\n🔎 .index 头部（{seg.index_path.name}，前 {n} 项）")
    print(f"  {'idx':>4}  {'rel_offset':>12} {'abs_offset':>12} {'position':>10}")
    with open(seg.index_path, "rb") as f:
        data = f.read(n * INDEX_ENTRY_SIZE)
    for i in range(0, len(data), INDEX_ENTRY_SIZE):
        rel, pos = struct.unpack(">II", data[i:i + INDEX_ENTRY_SIZE])
        if rel == 0 and pos == 0 and i > 0:
            # 已经到尾部填零段
            break
        print(f"  {i//INDEX_ENTRY_SIZE:>4}  {rel:>12} {seg.base_offset + rel:>12} {pos:>10}")


def print_timeindex_head(seg: Segment, n: int = 10) -> None:
    if not seg.timeindex_path or not seg.timeindex_path.exists():
        return
    print(f"\n⏱  .timeindex 头部（{seg.timeindex_path.name}，前 {n} 项）")
    print(f"  {'idx':>4}  {'timestamp_ms':>16} {'rel_offset':>12} {'abs_offset':>12}")
    with open(seg.timeindex_path, "rb") as f:
        data = f.read(n * TIMEINDEX_ENTRY_SIZE)
    for i in range(0, len(data), TIMEINDEX_ENTRY_SIZE):
        ts, rel = struct.unpack(">QI", data[i:i + TIMEINDEX_ENTRY_SIZE])
        if ts == 0 and rel == 0 and i > 0:
            break
        print(f"  {i//TIMEINDEX_ENTRY_SIZE:>4}  {ts:>16} {rel:>12} {seg.base_offset + rel:>12}")


# ---------------------------------------------------------------------------
# 「按 offset 找 segment」演示
# ---------------------------------------------------------------------------

def lookup_offset(segs: List[Segment], target: int) -> None:
    print(f"\n🔍 查找 offset = {target}")
    if not segs:
        print("  无 segment 可查")
        return

    # 二分：找最大的 base_offset ≤ target
    lo, hi, found = 0, len(segs) - 1, -1
    steps = []
    while lo <= hi:
        mid = (lo + hi) // 2
        steps.append((mid, segs[mid].base_offset))
        if segs[mid].base_offset <= target:
            found = mid
            lo = mid + 1
        else:
            hi = mid - 1
    print(f"  二分步骤（共 {len(steps)} 步）: {steps}")
    if found == -1:
        print(f"  ❌ 找不到合适的 segment（目标 offset 太小）")
        return
    seg = segs[found]
    print(f"  ✅ 命中 segment[{found}]  base_offset={seg.base_offset}")
    print(f"     文件: {seg.log_path.name}")
    print(f"     下一步：在该 .index 里二分 relative_offset = {target - seg.base_offset}")
    print(f"     再到 .log 顺序扫 ≤ index.interval.bytes（默认 4KB）即可命中")


# ---------------------------------------------------------------------------
# 工具
# ---------------------------------------------------------------------------

def human(n: int) -> str:
    for unit in ["B", "KB", "MB", "GB"]:
        if n < 1024:
            return f"{n:.1f}{unit}" if unit != "B" else f"{n}{unit}"
        n /= 1024
    return f"{n:.1f}TB"


# ---------------------------------------------------------------------------
# 主入口
# ---------------------------------------------------------------------------

def main() -> None:
    parser = argparse.ArgumentParser(description="探索 Kafka Topic-Partition 目录的 segment 结构")
    parser.add_argument("--part-dir", type=str, required=True,
                        help="分区目录，如 /var/lib/kafka/data/learn.07.storage-0")
    parser.add_argument("--lookup", type=int, default=None,
                        help="查找某 offset 落在哪个 segment 上（演示二分过程）")
    parser.add_argument("--head", type=int, default=10,
                        help="每个文件打印前 N 项索引，默认 10")
    args = parser.parse_args()

    part_dir = Path(args.part_dir)
    try:
        segs = scan_segments(part_dir)
    except FileNotFoundError as e:
        print(f"❌ {e}", file=sys.stderr)
        sys.exit(1)

    print_summary(part_dir, segs)
    print_leader_epoch(part_dir)

    if segs:
        # 默认展示第一个 segment 的索引头部
        print_index_head(segs[0], args.head)
        print_timeindex_head(segs[0], args.head)

    if args.lookup is not None:
        lookup_offset(segs, args.lookup)

    print()


if __name__ == "__main__":
    main()
