#!/usr/bin/env python3
"""
Time window helpers shared by the AP+MD workflow.
"""

from __future__ import annotations

import re
from dataclasses import dataclass
from datetime import datetime, timedelta
from pathlib import Path
from typing import List, Optional, Tuple


_FILENAME_TS_RE = re.compile(r"_(?P<date>\d{8})_(?P<time>\d{6})(?:\d+)?(?=\.)")
_TIME_POINT_RE = re.compile(r"(?P<hour>\d{1,2})点(?P<minute>\d{1,2})\s*(?P<direction>进|出)")


@dataclass(frozen=True)
class NamedTimeWindow:
    index: int
    label: str
    start_text: str
    end_text: str
    time_range: str
    date: Optional[str]


def extract_embedded_timestamp(path: Path) -> Optional[datetime]:
    match = _FILENAME_TS_RE.search(path.name)
    if not match:
        return None
    try:
        return datetime.strptime(
            f"{match.group('date')}_{match.group('time')}",
            "%Y%m%d_%H%M%S",
        )
    except ValueError:
        return None


def resolve_time_range_window(
    time_range: Optional[str],
    date: Optional[str],
    reference_dt: Optional[datetime],
) -> Optional[Tuple[datetime, datetime]]:
    if not time_range:
        return None

    try:
        start_text, end_text = [item.strip() for item in time_range.split("-", 1)]
    except ValueError as exc:
        raise ValueError(f"Bad time range: {time_range}") from exc

    ref = reference_dt or datetime.now()
    start_date, end_date = _resolve_dates(date=date, reference_dt=ref)
    start_dt = datetime.combine(start_date.date(), datetime.strptime(start_text, "%H:%M:%S").time())
    end_dt = datetime.combine(end_date.date(), datetime.strptime(end_text, "%H:%M:%S").time())
    if end_dt < start_dt and date is None:
        end_dt += timedelta(days=1)
    return start_dt, end_dt


def parse_named_time_windows(points_file: Path, date: Optional[str]) -> List[NamedTimeWindow]:
    markers: List[Tuple[str, str]] = []
    for raw_line in points_file.read_text(encoding="utf-8").splitlines():
        line = raw_line.strip().strip("，,")
        if not line:
            continue
        match = _TIME_POINT_RE.search(line)
        if not match:
            continue
        hour = int(match.group("hour"))
        minute = int(match.group("minute"))
        direction = match.group("direction")
        markers.append((f"{hour:02d}:{minute:02d}", direction))

    if not markers:
        raise ValueError(f"No valid time points found: {points_file}")
    if len(markers) % 2 != 0:
        raise ValueError(f"Odd number of time points in {points_file}")

    windows: List[NamedTimeWindow] = []
    for index in range(0, len(markers), 2):
        start_text, start_direction = markers[index]
        end_text, end_direction = markers[index + 1]
        if start_direction != "进" or end_direction != "出":
            raise ValueError(
                f"Unexpected time point order at pair {(index // 2) + 1}: "
                f"{start_text}{start_direction} -> {end_text}{end_direction}"
            )
        label = f"stage{(index // 2) + 1:02d}_{start_text.replace(':', '')}_{end_text.replace(':', '')}"
        windows.append(
            NamedTimeWindow(
                index=(index // 2) + 1,
                label=label,
                start_text=start_text,
                end_text=end_text,
                time_range=f"{start_text}:00-{end_text}:59",
                date=_normalize_cli_date(date),
            )
        )
    return windows


def intervals_overlap(
    left_start: datetime,
    left_end: datetime,
    right_start: datetime,
    right_end: datetime,
) -> bool:
    return left_start <= right_end and right_start <= left_end


def _resolve_dates(date: Optional[str], reference_dt: datetime) -> Tuple[datetime, datetime]:
    normalized = _normalize_cli_date(date)
    if not normalized:
        return reference_dt, reference_dt

    if "," in normalized:
        start_token, end_token = [token.strip() for token in normalized.split(",", 1)]
    else:
        start_token = normalized
        end_token = normalized

    start_dt = datetime.strptime(f"{reference_dt.year}-{start_token}", "%Y-%m-%d")
    end_dt = datetime.strptime(f"{reference_dt.year}-{end_token}", "%Y-%m-%d")
    return start_dt, end_dt


def _normalize_cli_date(date: Optional[str]) -> Optional[str]:
    if not date:
        return None
    normalized = date.strip()
    if not normalized:
        return None
    if "," in normalized:
        start_token, end_token = [token.strip() for token in normalized.split(",", 1)]
        if start_token == end_token:
            return start_token
        return f"{start_token},{end_token}"
    return normalized
