#!/usr/bin/env python3

from __future__ import annotations

import argparse
import re
import sys
import tomllib
from dataclasses import dataclass
from pathlib import Path

ROOT = Path(__file__).resolve().parent.parent
SETUP_PY = ROOT / "setup.py"
PYPROJECT_TOML = ROOT / "pyproject.toml"

PYTHON_CLASSIFIER_PREFIX = "Programming Language :: Python :: "
RUFF_TARGET_RE = re.compile(r"^py(\d+)(\d+)$")
REQUIRES_PYTHON_MIN_RE = re.compile(r"(?:>=|==)\s*(\d+)\.(\d+)")
VERSION_SHORT_RE = re.compile(r"^(\d+)\.(\d+)")
SETUP_PYTHON_REQUIRES_RE = re.compile(
	r"""python_requires\s*=\s*["']([^"']+)["']""",
)

GREEN = "\033[32m"
RED = "\033[31m"
RESET = "\033[0m"


def color_status(status: str, *, file: object = sys.stdout) -> str:
	if status == "ok":
		color = GREEN
	elif status == "MISMATCH":
		color = RED
	else:
		return status
	if hasattr(file, "isatty") and file.isatty():
		return f"{color}{status}{RESET}"
	return status


@dataclass(frozen=True, order=True)
class PythonVersion:
	major: int
	minor: int

	@classmethod
	def parse(cls, value: str) -> PythonVersion:
		match = VERSION_SHORT_RE.match(value.strip())
		if match is None:
			raise ValueError(f"invalid Python version: {value!r}")
		return cls(int(match.group(1)), int(match.group(2)))

	@classmethod
	def from_ruff_target(cls, value: str) -> PythonVersion:
		match = RUFF_TARGET_RE.fullmatch(value)
		if match is None:
			raise ValueError(f"invalid ruff target-version: {value!r}")
		return cls(int(match.group(1)), int(match.group(2)))

	@property
	def short(self) -> str:
		return f"{self.major}.{self.minor}"

	@property
	def setup_requires(self) -> str:
		return f">={self.major}.{self.minor}.0"

	@property
	def ruff_target(self) -> str:
		return f"py{self.major}{self.minor}"

	def classifier(self) -> str:
		return f"{PYTHON_CLASSIFIER_PREFIX}{self.short}"


def min_version_from_requires(requires_python: str) -> PythonVersion:
	match = REQUIRES_PYTHON_MIN_RE.search(requires_python)
	if match is None:
		raise ValueError(f"unsupported requires-python specifier: {requires_python!r}")
	return PythonVersion(int(match.group(1)), int(match.group(2)))


def get_nested(data: dict, *keys: str) -> object | None:
	current: object = data
	for key in keys:
		if not isinstance(current, dict) or key not in current:
			return None
		current = current[key]
	return current


def get_str(data: dict, *keys: str) -> str | None:
	value = get_nested(data, *keys)
	return value if isinstance(value, str) else None


def get_str_list(data: dict, *keys: str) -> list[str] | None:
	value = get_nested(data, *keys)
	if not isinstance(value, list):
		return None
	return [item for item in value if isinstance(item, str)]


def read_setup_python_requires() -> str:
	text = SETUP_PY.read_text(encoding="utf-8")
	match = SETUP_PYTHON_REQUIRES_RE.search(text)
	if match is None:
		raise ValueError("python_requires not found in setup.py")
	return match.group(1)


def replace_setup_python_requires(value: str) -> None:
	text = SETUP_PY.read_text(encoding="utf-8")
	new_text, count = SETUP_PYTHON_REQUIRES_RE.subn(
		f'python_requires="{value}"',
		text,
		count=1,
	)
	if count != 1:
		raise ValueError("failed to update python_requires in setup.py")
	SETUP_PY.write_text(new_text, encoding="utf-8")


def read_pyproject() -> dict:
	with PYPROJECT_TOML.open("rb") as file:
		return tomllib.load(file)


def replace_pyproject_line_if_present(key: str, value: str) -> bool:
	prefix = f"{key} = "
	lines: list[str] = []
	updated = False
	for line in PYPROJECT_TOML.read_text(encoding="utf-8").splitlines(keepends=True):
		if line.startswith(prefix):
			indent = line[: len(line) - len(line.lstrip("\t"))]
			lines.append(f'{indent}{key} = "{value}"\n')
			updated = True
		else:
			lines.append(line)
	if updated:
		PYPROJECT_TOML.write_text("".join(lines), encoding="utf-8")
	return updated


def python_classifiers(classifiers: list[str]) -> list[PythonVersion]:
	versions: list[PythonVersion] = []
	for classifier in classifiers:
		if not classifier.startswith(PYTHON_CLASSIFIER_PREFIX):
			continue
		versions.append(PythonVersion.parse(classifier.removeprefix(PYTHON_CLASSIFIER_PREFIX)))
	return sorted(versions, key=lambda item: (item.major, item.minor))


def classifier_versions_display(versions: list[PythonVersion]) -> str:
	return ", ".join(item.short for item in versions)


def contiguous_classifier_versions(
	min_version: PythonVersion,
	max_version: PythonVersion,
) -> list[PythonVersion]:
	versions: list[PythonVersion] = []
	minor = min_version.minor
	while (min_version.major, minor) <= (max_version.major, max_version.minor):
		versions.append(PythonVersion(min_version.major, minor))
		minor += 1
	return versions


def classifier_gap_versions(versions: list[PythonVersion]) -> list[PythonVersion]:
	if len(versions) < 2:
		return []
	expected = contiguous_classifier_versions(versions[0], versions[-1])
	return [item for item in expected if item not in versions]


def fix_classifiers(min_version: PythonVersion) -> None:
	lines = PYPROJECT_TOML.read_text(encoding="utf-8").splitlines(keepends=True)
	new_lines: list[str] = []
	in_classifiers = False
	python_classifier_lines: list[tuple[str, PythonVersion | None]] = []

	for line in lines:
		if line.startswith("classifiers = ["):
			in_classifiers = True
			new_lines.append(line)
			continue
		if in_classifiers:
			if line.strip() == "]":
				in_classifiers = False
				kept_versions = sorted(
					(
						version
						for _, version in python_classifier_lines
						if version is not None and version >= min_version
					),
					key=lambda item: (item.major, item.minor),
				)
				if not any(version == min_version for version in kept_versions):
					kept_versions.insert(0, min_version)
				if kept_versions:
					kept_versions = contiguous_classifier_versions(
						kept_versions[0],
						kept_versions[-1],
					)
				tab = "\t"
				if python_classifier_lines:
					tab_match = re.match(r"^(\t+)", python_classifier_lines[0][0])
					if tab_match:
						tab = tab_match.group(1)
				for version in kept_versions:
					new_lines.append(f'{tab}"{version.classifier()}",\n')
				new_lines.append(line)
				python_classifier_lines = []
				continue
			if PYTHON_CLASSIFIER_PREFIX in line:
				match = re.search(r'"([^"]+)"', line)
				if match is None:
					new_lines.append(line)
					continue
				classifier = match.group(1)
				version = PythonVersion.parse(classifier.removeprefix(PYTHON_CLASSIFIER_PREFIX))
				python_classifier_lines.append((line, version))
				continue
		new_lines.append(line)

	PYPROJECT_TOML.write_text("".join(new_lines), encoding="utf-8")


@dataclass
class CheckResult:
	name: str
	expected: str
	actual: str

	@property
	def ok(self) -> bool:
		return self.expected == self.actual

	@property
	def values_display(self) -> str:
		if self.expected == self.actual:
			return self.expected
		return f"expected {self.expected!r}, got {self.actual!r}"


def version_check(name: str, min_version: PythonVersion, actual: str | None) -> CheckResult | None:
	if actual is None:
		return None
	try:
		actual_version = PythonVersion.parse(actual)
	except ValueError:
		return CheckResult(name, min_version.short, actual)
	return CheckResult(name, min_version.short, actual_version.short)


def ruff_check(min_version: PythonVersion, actual: str | None) -> CheckResult | None:
	if actual is None:
		return None
	try:
		actual_version = PythonVersion.from_ruff_target(actual)
	except ValueError:
		return CheckResult(
			"tool.ruff target-version",
			min_version.ruff_target,
			actual,
		)
	return CheckResult(
		"tool.ruff target-version",
		min_version.ruff_target,
		actual_version.ruff_target,
	)


def setup_requires_check(min_version: PythonVersion) -> CheckResult | None:
	try:
		setup_requires = read_setup_python_requires()
	except ValueError:
		return None
	setup_min = min_version_from_requires(setup_requires)
	return CheckResult(
		"setup.py python_requires",
		min_version.setup_requires,
		setup_min.setup_requires,
	)


def classifier_check(min_version: PythonVersion, classifiers: list[str] | None) -> CheckResult | None:
	if classifiers is None:
		return None
	classifier_versions = python_classifiers(classifiers)
	if classifier_versions:
		lowest_classifier = classifier_versions[0]
		classifier_expected = min_version.short
		classifier_actual = lowest_classifier.short
		invalid_classifiers = [
			item.short
			for item in classifier_versions
			if item < min_version
		]
		if invalid_classifiers:
			classifier_actual = (
				f"{classifier_actual} (also below minimum: {', '.join(invalid_classifiers)})"
			)
	else:
		classifier_expected = min_version.short
		classifier_actual = "(missing)"
	return CheckResult(
		"project.classifiers (minimum)",
		classifier_expected,
		classifier_actual,
	)


def classifier_gap_check(classifiers: list[str] | None) -> CheckResult | None:
	if classifiers is None:
		return None
	classifier_versions = python_classifiers(classifiers)
	if len(classifier_versions) < 2:
		if len(classifier_versions) == 1:
			version = classifier_versions[0].short
			return CheckResult("project.classifiers (contiguous)", version, version)
		return None

	expected_versions = contiguous_classifier_versions(
		classifier_versions[0],
		classifier_versions[-1],
	)
	expected = classifier_versions_display(expected_versions)
	missing = classifier_gap_versions(classifier_versions)
	if missing:
		actual = (
			f"{classifier_versions_display(classifier_versions)} "
			f"(missing {classifier_versions_display(missing)})"
		)
	else:
		actual = expected
	return CheckResult("project.classifiers (contiguous)", expected, actual)


def collect_checks() -> tuple[PythonVersion, list[CheckResult]]:
	project = read_pyproject()
	requires_python = project["project"]["requires-python"]
	min_version = min_version_from_requires(requires_python)

	checks: list[CheckResult] = [
		CheckResult(
			"project.requires-python (reference)",
			min_version.short,
			min_version.short,
		),
	]

	for check in (
		setup_requires_check(min_version),
		ruff_check(min_version, get_str(project, "tool", "ruff", "target-version")),
		version_check(
			"tool.mypy python_version",
			min_version,
			get_str(project, "tool", "mypy", "python_version"),
		),
		version_check(
			"tool.pylint.main py-version",
			min_version,
			get_str(project, "tool", "pylint", "main", "py-version"),
		),
		classifier_check(min_version, get_str_list(project, "project", "classifiers")),
		classifier_gap_check(get_str_list(project, "project", "classifiers")),
	):
		if check is not None:
			checks.append(check)

	return min_version, checks


def apply_fixes(min_version: PythonVersion) -> None:
	project = read_pyproject()

	try:
		read_setup_python_requires()
	except ValueError:
		pass
	else:
		replace_setup_python_requires(min_version.setup_requires)

	if get_str(project, "tool", "ruff", "target-version") is not None:
		replace_pyproject_line_if_present("target-version", min_version.ruff_target)

	if get_str(project, "tool", "mypy", "python_version") is not None:
		replace_pyproject_line_if_present("python_version", min_version.short)

	if get_str(project, "tool", "pylint", "main", "py-version") is not None:
		replace_pyproject_line_if_present("py-version", min_version.short)

	if get_str_list(project, "project", "classifiers") is not None:
		fix_classifiers(min_version)


def main() -> int:
	parser = argparse.ArgumentParser(
		description=(
			"Check that Python version settings match project.requires-python "
			"in pyproject.toml."
		),
	)
	parser.add_argument(
		"--fix",
		action="store_true",
		help="Update setup.py and pyproject.toml tool sections to match requires-python.",
	)
	args = parser.parse_args()

	min_version, checks = collect_checks()
	failures = [check for check in checks if not check.ok and check.name != "project.requires-python (reference)"]

	if args.fix:
		if failures:
			apply_fixes(min_version)
			_, checks = collect_checks()
			failures = [
				check
				for check in checks
				if not check.ok and check.name != "project.requires-python (reference)"
			]
			if failures:
				print("Updated files, but some checks still fail:", file=sys.stderr)
				for check in failures:
					print(
						f"  {check.name}: {check.values_display}",
						file=sys.stderr,
					)
			else:
				print(
					f"Updated Python version settings to match requires-python ({min_version.short}).",
				)
		else:
			print(f"Python version settings already match requires-python ({min_version.short}).")
		return 1 if failures else 0

	for check in checks:
		status = "ok" if check.ok else "MISMATCH"
		print(
			f"[{color_status(status)}] {check.name}: {check.values_display}",
		)

	if failures:
		print(
			"\nRun with --fix to update setup.py and pyproject.toml from project.requires-python.",
			file=sys.stderr,
		)
		return 1
	return 0


if __name__ == "__main__":
	sys.exit(main())
