307 lines
9.2 KiB
Python
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/")
|
|
)
|