Files
huheng-research/src/shanghai_housing/import_templates.py
2026-06-23 18:59:42 +08:00

307 lines
9.2 KiB
Python

from __future__ import annotations
import csv
import json
import re
from dataclasses import dataclass
from datetime import date
from pathlib import Path
from typing import Any, Mapping
from urllib.parse import urlparse
PROJECT_ROOT = Path(__file__).resolve().parents[2]
IMPORT_TEMPLATE_SPEC_PATH = PROJECT_ROOT / "config" / "import_templates.json"
@dataclass(frozen=True)
class ImportValidationError:
row: int
column: str
message: str
@dataclass(frozen=True)
class ImportValidationResult:
template_name: str
row_count: int
valid: bool
errors: list[ImportValidationError]
class ImportTemplateError(ValueError):
pass
def load_import_template_specs(
path: Path = IMPORT_TEMPLATE_SPEC_PATH,
) -> dict[str, dict[str, Any]]:
with path.open("r", encoding="utf-8") as handle:
return json.load(handle)
def list_template_names(path: Path = IMPORT_TEMPLATE_SPEC_PATH) -> list[str]:
return sorted(load_import_template_specs(path).keys())
def get_template_spec(
template_name: str,
path: Path = IMPORT_TEMPLATE_SPEC_PATH,
) -> dict[str, Any]:
specs = load_import_template_specs(path)
try:
return specs[template_name]
except KeyError as exc:
valid_names = ", ".join(sorted(specs.keys()))
raise ImportTemplateError(
f"unknown template '{template_name}', expected one of: {valid_names}"
) from exc
def validate_csv_file(
template_name: str,
file_path: Path,
spec_path: Path = IMPORT_TEMPLATE_SPEC_PATH,
) -> ImportValidationResult:
rows = read_csv_rows(file_path)
return validate_import_rows(template_name, rows, spec_path=spec_path)
def validate_tabular_file(
template_name: str,
file_path: Path,
spec_path: Path = IMPORT_TEMPLATE_SPEC_PATH,
) -> ImportValidationResult:
suffix = file_path.suffix.lower()
if suffix == ".csv":
rows = read_csv_rows(file_path)
elif suffix in {".xlsx", ".xlsm"}:
rows = read_excel_rows(file_path)
else:
raise ImportTemplateError("file must be a .csv, .xlsx, or .xlsm document")
return validate_import_rows(template_name, rows, spec_path=spec_path)
def load_tabular_rows(file_path: Path) -> list[dict[str, str]]:
suffix = file_path.suffix.lower()
if suffix == ".csv":
return read_csv_rows(file_path)
if suffix in {".xlsx", ".xlsm"}:
return read_excel_rows(file_path)
raise ImportTemplateError("file must be a .csv, .xlsx, or .xlsm document")
def read_csv_rows(file_path: Path) -> list[dict[str, str]]:
with file_path.open("r", encoding="utf-8-sig", newline="") as handle:
reader = csv.DictReader(handle)
return [dict(row) for row in reader]
def read_excel_rows(file_path: Path) -> list[dict[str, str]]:
try:
from openpyxl import load_workbook
except ImportError as exc:
raise ImportTemplateError(
"Excel validation requires openpyxl; install project dependencies first"
) from exc
workbook = load_workbook(file_path, read_only=True, data_only=True)
worksheet = workbook.active
rows = worksheet.iter_rows(values_only=True)
try:
headers = [str(value).strip() if value is not None else "" for value in next(rows)]
except StopIteration:
return []
records: list[dict[str, str]] = []
for values in rows:
if values is None or all(value is None or str(value).strip() == "" for value in values):
continue
record = {
header: "" if value is None else str(value).strip()
for header, value in zip(headers, values)
if header
}
records.append(record)
return records
def validate_import_rows(
template_name: str,
rows: list[Mapping[str, Any]],
spec_path: Path = IMPORT_TEMPLATE_SPEC_PATH,
) -> ImportValidationResult:
spec = get_template_spec(template_name, spec_path)
columns = spec.get("columns", [])
expected_columns = [column["name"] for column in columns]
required_columns = [
column["name"] for column in columns if bool(column.get("required", False))
]
errors: list[ImportValidationError] = []
seen_columns: set[str] = set()
if rows:
seen_columns = set(rows[0].keys())
missing_headers = [column for column in expected_columns if column not in seen_columns]
unexpected_headers = sorted(seen_columns - set(expected_columns))
for column in missing_headers:
errors.append(ImportValidationError(0, column, "missing required header"))
for column in unexpected_headers:
errors.append(ImportValidationError(0, column, "unexpected header"))
if not rows:
errors.append(ImportValidationError(0, "*", "file has no data rows"))
for index, row in enumerate(rows, start=2):
for column in columns:
name = column["name"]
value = normalize_value(row.get(name))
if name in required_columns and value == "":
errors.append(ImportValidationError(index, name, "value is required"))
continue
if value == "":
continue
errors.extend(validate_value(index, name, value, column))
return ImportValidationResult(
template_name=template_name,
row_count=len(rows),
valid=not errors,
errors=errors,
)
def normalize_value(value: Any) -> str:
if value is None:
return ""
return str(value).strip()
def validate_value(
row_index: int,
column_name: str,
value: str,
column: Mapping[str, Any],
) -> list[ImportValidationError]:
errors: list[ImportValidationError] = []
if has_sensitive_text(value):
return [
ImportValidationError(
row_index,
column_name,
"value appears to contain credentials or a sensitive local path",
)
]
column_type = column["type"]
if column_type == "text":
return errors
if column_type == "month":
if not is_valid_month(value):
errors.append(ImportValidationError(row_index, column_name, "must be YYYY-MM"))
return errors
if column_type == "date":
if not is_valid_date(value):
errors.append(ImportValidationError(row_index, column_name, "must be YYYY-MM-DD"))
return errors
if column_type == "url":
if not is_valid_url(value):
errors.append(ImportValidationError(row_index, column_name, "must be http(s) URL"))
return errors
if column_type == "enum":
allowed_values = column.get("allowed_values", [])
if value not in allowed_values:
allowed = ", ".join(allowed_values)
errors.append(
ImportValidationError(row_index, column_name, f"must be one of: {allowed}")
)
return errors
if column_type == "integer":
parsed = parse_integer(value)
if parsed is None:
errors.append(ImportValidationError(row_index, column_name, "must be an integer"))
return errors
errors.extend(validate_numeric_range(row_index, column_name, parsed, column))
return errors
if column_type == "decimal":
parsed = parse_decimal(value)
if parsed is None:
errors.append(ImportValidationError(row_index, column_name, "must be a number"))
return errors
errors.extend(validate_numeric_range(row_index, column_name, parsed, column))
return errors
errors.append(
ImportValidationError(row_index, column_name, f"unknown type '{column_type}'")
)
return errors
def validate_numeric_range(
row_index: int,
column_name: str,
value: float,
column: Mapping[str, Any],
) -> list[ImportValidationError]:
errors: list[ImportValidationError] = []
minimum = column.get("min")
maximum = column.get("max")
if minimum is not None and value < float(minimum):
errors.append(
ImportValidationError(row_index, column_name, f"must be >= {minimum}")
)
if maximum is not None and value > float(maximum):
errors.append(
ImportValidationError(row_index, column_name, f"must be <= {maximum}")
)
return errors
def parse_integer(value: str) -> int | None:
if not re.fullmatch(r"-?\d+", value):
return None
return int(value)
def parse_decimal(value: str) -> float | None:
try:
return float(value)
except ValueError:
return None
def is_valid_month(value: str) -> bool:
if not re.fullmatch(r"\d{4}-\d{2}", value):
return False
month = int(value[5:])
return 1 <= month <= 12
def is_valid_date(value: str) -> bool:
if not re.fullmatch(r"\d{4}-\d{2}-\d{2}", value):
return False
try:
date.fromisoformat(value)
except ValueError:
return False
return True
def is_valid_url(value: str) -> bool:
parsed = urlparse(value)
return parsed.scheme in {"http", "https"} and bool(parsed.netloc)
def has_sensitive_text(value: str) -> bool:
lower = value.lower()
return (
"password" in lower
or "secret" in lower
or "postgres://" in lower
or "database_url" in lower
or "root_" in lower
or value.startswith("/Users/")
or value.startswith("/private/")
)