from __future__ import annotations

import csv
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Iterable

from openpyxl import load_workbook


@dataclass(frozen=True)
class CellRange:
    header_row: int
    first_column: int
    end_row: int
    end_column: int

    @property
    def nrows(self) -> int:
        return max(0, self.end_row - self.header_row)

    @property
    def ncolumns(self) -> int:
        return max(0, self.end_column - self.first_column + 1)


@dataclass(frozen=True)
class DataSource:
    path: Path
    extension: str
    sheet_name: str | None = None
    encoding: str | None = None
    delimiter: str | None = None


@dataclass(frozen=True)
class MetadataCell:
    row: int
    column: int
    value: Any
    region: str


@dataclass
class DataTable:
    labels: list[Any]
    columns: list[list[Any]]
    metadata: list[MetadataCell]
    data_range: CellRange
    source: DataSource
    raw_rows: list[list[Any]] = field(repr=False, default_factory=list)

    def find_column(
        self,
        selector: str | int,
        *,
        ignore_case: bool = True,
        regex: bool = False,
    ) -> tuple[Any | None, list[Any] | None]:
        if isinstance(selector, int):
            index = selector
        else:
            target = str(selector)
            index = None
            for i, label in enumerate(self.labels):
                text = "" if label is None else str(label)
                if regex:
                    flags = re.IGNORECASE if ignore_case else 0
                    if re.search(target, text, flags=flags):
                        index = i
                        break
                elif ignore_case:
                    if text.casefold() == target.casefold():
                        index = i
                        break
                elif text == target:
                    index = i
                    break
            if index is None:
                try:
                    index = int(target)
                except ValueError:
                    return None, None

        if index < 0 or index >= len(self.columns):
            return None, None
        return self.labels[index], self.columns[index]

    def metadata_dict(self) -> dict[str, Any]:
        rows: dict[int, list[MetadataCell]] = {}
        for cell in self.metadata:
            rows.setdefault(cell.row, []).append(cell)

        result: dict[str, Any] = {}
        for cells in rows.values():
            cells.sort(key=lambda c: c.column)
            if len(cells) >= 2:
                key = str(cells[0].value).strip()
                if key:
                    result[key] = cells[1].value
        return result

    @property
    def nrows(self) -> int:
        return self.data_range.nrows

    @property
    def ncolumns(self) -> int:
        return self.data_range.ncolumns


def _is_blank(value: Any) -> bool:
    return value is None or (isinstance(value, str) and value.strip() == "")


def _convert_value(value: Any, force_numeric: bool) -> Any:
    if _is_blank(value):
        return None
    if not force_numeric:
        return value.strip() if isinstance(value, str) else value
    if isinstance(value, bool):
        return float(value)
    if isinstance(value, (int, float)):
        return value
    if isinstance(value, str):
        try:
            return float(value.strip())
        except ValueError:
            return None
    try:
        return float(value)
    except (TypeError, ValueError, OverflowError):
        return None


def _read_xlsx(path: Path, *, sheet, data_only: bool):
    wb = load_workbook(path, data_only=data_only, read_only=True)
    if sheet is None:
        ws = wb.active
    elif isinstance(sheet, int):
        ws = wb.worksheets[sheet]
    else:
        ws = wb[sheet]
    rows = [list(row) for row in ws.iter_rows(values_only=True)]
    source = DataSource(path=path, extension=path.suffix.lower(), sheet_name=ws.title)
    wb.close()
    return rows, source


def _try_read_text(path: Path, encodings: Iterable[str]):
    last_error = None
    for encoding in encodings:
        try:
            return path.read_text(encoding=encoding), encoding
        except UnicodeDecodeError as exc:
            last_error = exc
    raise last_error or ValueError("No encoding candidates were supplied.")


def _detect_delimiter(text: str, extension: str, delimiter: str | None):
    if delimiter is not None:
        return delimiter
    lines = [line for line in text.splitlines() if line.strip()]
    sample = "\n".join(lines[:20])
    if extension == ".csv":
        try:
            return csv.Sniffer().sniff(sample, delimiters=",;\t").delimiter
        except csv.Error:
            return ","
    if "\t" in sample:
        return "\t"
    if "," in sample:
        return ","
    if ";" in sample:
        return ";"
    return None


def _read_text_table(path: Path, *, encoding: str | None, delimiter: str | None):
    candidates = [encoding] if encoding else ["utf-8-sig", "cp932", "utf-8"]
    text, used_encoding = _try_read_text(path, candidates)
    used_delimiter = _detect_delimiter(text, path.suffix.lower(), delimiter)
    if used_delimiter is None:
        rows = [line.split() for line in text.splitlines()]
    else:
        rows = [list(row) for row in csv.reader(text.splitlines(), delimiter=used_delimiter)]
    source = DataSource(
        path=path,
        extension=path.suffix.lower(),
        encoding=used_encoding,
        delimiter=used_delimiter,
    )
    return rows, source


def _rectangularize(rows):
    if not rows:
        return []
    width = max(len(row) for row in rows)
    return [list(row) + [None] * (width - len(row)) for row in rows]


def _metadata_region(row, column, *, header_row, first_column, end_row, end_column):
    if row < header_row:
        return "above"
    if row > end_row:
        return "below"
    if column < first_column:
        return "left"
    if column > end_column:
        return "right"
    return "inside"


def read_data_table(
    path: str | Path,
    *,
    header_row: int = 1,
    first_column: int = 1,
    force_numeric: bool = True,
    sheet: str | int | None = None,
    data_only: bool = True,
    encoding: str | None = None,
    delimiter: str | None = None,
) -> DataTable:
    """
    .xlsx/.xlsm/.csv/.txtからminimum-matrix形式の表を読む。

    - header_row、first_columnは1始まり。
    - ヘッダーはfirst_columnから最初の空セル直前まで。
    - データはヘッダー直下から開始。
    - 第1データ列が空になった行の直前で終了。
    - 表範囲外の全非空セルをmetadataとして返す。
    """
    source_path = Path(path).expanduser()
    if not source_path.exists():
        raise FileNotFoundError(f"Input file not found: {source_path}")
    if header_row < 1 or first_column < 1:
        raise ValueError("header_row and first_column must be positive 1-based integers")

    ext = source_path.suffix.lower()
    if ext in {".xlsx", ".xlsm"}:
        rows, source = _read_xlsx(source_path, sheet=sheet, data_only=data_only)
    elif ext in {".csv", ".txt"}:
        rows, source = _read_text_table(
            source_path, encoding=encoding, delimiter=delimiter
        )
    else:
        raise ValueError(
            f"Unsupported extension [{ext}]. Supported: .xlsx, .xlsm, .csv, .txt"
        )

    rows = _rectangularize(rows)
    if header_row > len(rows):
        raise ValueError(f"header_row={header_row} exceeds row count {len(rows)}")

    header = rows[header_row - 1]
    if first_column > len(header):
        raise ValueError(
            f"first_column={first_column} exceeds column count {len(header)}"
        )

    labels = []
    col = first_column - 1
    while col < len(header) and not _is_blank(header[col]):
        value = header[col]
        labels.append(value.strip() if isinstance(value, str) else value)
        col += 1
    if not labels:
        raise ValueError(
            f"No labels found at header_row={header_row}, first_column={first_column}"
        )

    end_column = first_column + len(labels) - 1
    columns = [[] for _ in labels]
    end_row = header_row

    for row_number in range(header_row + 1, len(rows) + 1):
        row = rows[row_number - 1]
        first_value = row[first_column - 1]
        if _is_blank(first_value):
            break
        for i in range(len(labels)):
            value = row[first_column - 1 + i]
            columns[i].append(_convert_value(value, force_numeric))
        end_row = row_number

    data_range = CellRange(
        header_row=header_row,
        first_column=first_column,
        end_row=end_row,
        end_column=end_column,
    )

    metadata = []
    for row_index, row in enumerate(rows, start=1):
        for column_index, value in enumerate(row, start=1):
            if _is_blank(value):
                continue
            inside = (
                header_row <= row_index <= end_row
                and first_column <= column_index <= end_column
            )
            if inside:
                continue
            metadata.append(
                MetadataCell(
                    row=row_index,
                    column=column_index,
                    value=value,
                    region=_metadata_region(
                        row_index,
                        column_index,
                        header_row=header_row,
                        first_column=first_column,
                        end_row=end_row,
                        end_column=end_column,
                    ),
                )
            )

    return DataTable(
        labels=labels,
        columns=columns,
        metadata=metadata,
        data_range=data_range,
        source=source,
        raw_rows=rows,
    )
