#!/usr/bin/env python3
"""Static consistency checks for BCOS MariaDB migrations, views and seeds.

This is not a substitute for a clean MariaDB staging installation. It catches
known packaging defects before database execution: malformed concatenated SQL
tokens, duplicate active table definitions, INSERT/schema mismatches, foreign
key creation-order errors and qualified view-column mismatches.
"""
from __future__ import annotations

from collections import defaultdict
from pathlib import Path
import json
import re
import sys

ROOT = Path(__file__).resolve().parents[1]
MIGRATIONS = ROOT / "database" / "migrations"
VIEWS = ROOT / "database" / "views"
SEEDS = ROOT / "database" / "seeds"

TYPE_PATTERN = r"VARCHAR|CHAR|BIGINT|SMALLINT|TINYINT|MEDIUMINT|INT|INTEGER|DECIMAL|NUMERIC|FLOAT|DOUBLE|TEXT|LONGTEXT|JSON|ENUM|DATETIME|TIMESTAMP|DATE|TIME|BOOLEAN|BLOB"
BAD_TOKENS = ["PRIMARYKEY", "NOTNULL", "TEXTNULL", "BIGINTUNSIGNED", "ONUPDATE", "CURRENT_TIMESTAMPON"]
CONSTRAINT_PREFIXES = {"PRIMARY", "UNIQUE", "KEY", "INDEX", "CONSTRAINT", "FOREIGN", "CHECK", "FULLTEXT"}


def split_statements(sql: str) -> list[str]:
    out: list[str] = []
    buf: list[str] = []
    quote: str | None = None
    line_comment = False
    block_comment = False
    i = 0
    while i < len(sql):
        char = sql[i]
        nxt = sql[i + 1] if i + 1 < len(sql) else ""
        if line_comment:
            if char == "\n":
                line_comment = False
                buf.append(char)
            i += 1
            continue
        if block_comment:
            if char == "*" and nxt == "/":
                block_comment = False
                i += 2
                continue
            i += 1
            continue
        if quote is not None:
            buf.append(char)
            if char == quote:
                if quote == "'" and nxt == "'":
                    buf.append(nxt)
                    i += 2
                    continue
                if i == 0 or sql[i - 1] != "\\":
                    quote = None
            i += 1
            continue
        if char == "-" and nxt == "-" and (i + 2 >= len(sql) or sql[i + 2].isspace()):
            line_comment = True
            i += 2
            continue
        if char == "#":
            line_comment = True
            i += 1
            continue
        if char == "/" and nxt == "*":
            block_comment = True
            i += 2
            continue
        if char in "'\"`":
            quote = char
            buf.append(char)
            i += 1
            continue
        if char == ";":
            statement = "".join(buf).strip()
            if statement:
                out.append(statement)
            buf = []
            i += 1
            continue
        buf.append(char)
        i += 1
    statement = "".join(buf).strip()
    if statement:
        out.append(statement)
    return out


def split_top_level(text: str) -> list[str]:
    out: list[str] = []
    start = 0
    depth = 0
    quote: str | None = None
    for i, char in enumerate(text):
        if quote is not None:
            if char == quote and (i == 0 or text[i - 1] != "\\"):
                quote = None
            continue
        if char in "'\"`":
            quote = char
        elif char == "(":
            depth += 1
        elif char == ")":
            depth -= 1
        elif char == "," and depth == 0:
            out.append(text[start:i].strip())
            start = i + 1
    out.append(text[start:].strip())
    return out


def create_table_blocks(sql: str):
    pattern = re.compile(r"CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?`?([A-Za-z0-9_]+)`?\s*\(", re.I)
    for match in pattern.finditer(sql):
        table = match.group(1).lower()
        start = match.end() - 1
        depth = 0
        quote: str | None = None
        for end in range(start, len(sql)):
            char = sql[end]
            if quote is not None:
                if char == quote and sql[end - 1] != "\\":
                    quote = None
                continue
            if char in "'\"`":
                quote = char
            elif char == "(":
                depth += 1
            elif char == ")":
                depth -= 1
                if depth == 0:
                    yield table, sql[start + 1 : end]
                    break


def build_schemas() -> dict[str, set[str]]:
    schemas: dict[str, set[str]] = {}
    for path in sorted(MIGRATIONS.glob("*.sql")):
        sql = path.read_text(encoding="utf-8")
        for table, body in create_table_blocks(sql):
            if table in schemas:
                continue
            columns: set[str] = set()
            for item in split_top_level(body):
                match = re.match(r"`?([A-Za-z0-9_]+)`?\s+", item)
                if match and match.group(1).upper() not in CONSTRAINT_PREFIXES:
                    columns.add(match.group(1).lower())
            schemas[table] = columns
        for match in re.finditer(r"ALTER\s+TABLE\s+`?([A-Za-z0-9_]+)`?\s+(.*?);", sql, re.I | re.S):
            table = match.group(1).lower()
            body = match.group(2)
            schemas.setdefault(table, set())
            for add in re.finditer(r"ADD\s+COLUMN\s+(?:IF\s+NOT\s+EXISTS\s+)?`?([A-Za-z0-9_]+)`?", body, re.I):
                schemas[table].add(add.group(1).lower())
            for add in re.finditer(rf"ADD\s+(?!COLUMN|KEY|INDEX|CONSTRAINT|PRIMARY|UNIQUE|FOREIGN)`?([A-Za-z0-9_]+)`?\s+(?:{TYPE_PATTERN})", body, re.I):
                schemas[table].add(add.group(1).lower())
    return schemas


def main() -> int:
    errors: list[str] = []
    schemas = build_schemas()

    definitions: dict[str, list[str]] = defaultdict(list)
    for path in sorted(MIGRATIONS.glob("*.sql")):
        sql = path.read_text(encoding="utf-8")
        for bad in BAD_TOKENS:
            if bad in sql:
                errors.append(f"{path.name}: contains malformed token {bad}")
        for table, _ in create_table_blocks(sql):
            definitions[table].append(path.name)
    for table, files in definitions.items():
        if len(files) > 1:
            errors.append(f"duplicate active CREATE TABLE {table}: {', '.join(files)}")

    # InnoDB foreign-key constraint names must be unique within a schema.
    named_constraints: dict[str, list[str]] = defaultdict(list)
    for path in sorted(MIGRATIONS.glob("*.sql")):
        sql = path.read_text(encoding="utf-8")
        for match in re.finditer(r"\bCONSTRAINT\s+`?([A-Za-z0-9_]+)`?\s+FOREIGN\s+KEY", sql, re.I):
            named_constraints[match.group(1).lower()].append(path.name)
    for name, files in named_constraints.items():
        if len(files) > 1:
            errors.append(f"duplicate schema-wide foreign-key constraint {name}: {', '.join(files)}")

    for path in [*sorted(MIGRATIONS.glob("*.sql")), *sorted(SEEDS.glob("*.sql"))]:
        sql = path.read_text(encoding="utf-8")
        for match in re.finditer(r"INSERT\s+(?:IGNORE\s+)?INTO\s+`?([A-Za-z0-9_]+)`?\s*\(([^)]*)\)", sql, re.I | re.S):
            table = match.group(1).lower()
            columns = [x.strip(" `\n\r\t").lower() for x in match.group(2).split(",")]
            if table not in schemas:
                errors.append(f"{path.name}: INSERT targets undefined table {table}")
                continue
            missing = [column for column in columns if column not in schemas[table]]
            if missing:
                errors.append(f"{path.name}: INSERT {table} references missing columns {missing}")

    created = {"bcos_migration_registry"}
    for path in sorted(MIGRATIONS.glob("*.sql")):
        for index, statement in enumerate(split_statements(path.read_text(encoding="utf-8")), 1):
            create = re.match(r"CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?`?([A-Za-z0-9_]+)`?", statement, re.I)
            if create:
                table = create.group(1).lower()
                for ref in re.findall(r"REFERENCES\s+`?([A-Za-z0-9_]+)`?", statement, re.I):
                    target = ref.lower()
                    if target not in created and target != table:
                        errors.append(f"{path.name} statement {index}: {table} references not-yet-created {target}")
                created.add(table)
                continue
            alter = re.match(r"ALTER\s+TABLE\s+`?([A-Za-z0-9_]+)`?", statement, re.I)
            if alter and alter.group(1).lower() not in created:
                errors.append(f"{path.name} statement {index}: ALTER targets missing table {alter.group(1)}")

    available = set(schemas)
    for path in sorted(VIEWS.glob("*.sql")):
        sql = path.read_text(encoding="utf-8")
        for match in re.finditer(r"CREATE\s+(?:OR\s+REPLACE\s+)?VIEW\s+`?([A-Za-z0-9_]+)`?\s+AS\s+(.*?);", sql, re.I | re.S):
            view = match.group(1).lower()
            body = match.group(2)
            refs = {x.lower() for x in re.findall(r"\b(?:FROM|JOIN)\s+`?([A-Za-z0-9_]+)`?", body, re.I)}
            missing = sorted(refs - available)
            if missing:
                errors.append(f"{path.name}: view {view} references missing objects {missing}")
            aliases: dict[str, str] = {}
            for alias_match in re.finditer(r"\b(?:FROM|JOIN)\s+`?([A-Za-z0-9_]+)`?(?:\s+(?:AS\s+)?([A-Za-z0-9_]+))?", body, re.I):
                table = alias_match.group(1).lower()
                alias = (alias_match.group(2) or table).lower()
                if alias.upper() in {"WHERE", "LEFT", "RIGHT", "INNER", "OUTER", "JOIN", "ON", "GROUP", "ORDER", "LIMIT", "UNION", "CROSS"}:
                    alias = table
                aliases[alias] = table
                aliases.setdefault(table, table)
            for col_match in re.finditer(r"\b([A-Za-z_][A-Za-z0-9_]*)\.\s*`?([A-Za-z_][A-Za-z0-9_]*)`?", body):
                alias = col_match.group(1).lower()
                column = col_match.group(2).lower()
                table = aliases.get(alias)
                if table in schemas and column not in schemas[table]:
                    errors.append(f"{path.name}: view {view} references missing {table}.{column}")
            available.add(view)

    result = {
        "status": "FAIL" if errors else "PASS",
        "errors": errors,
        "migration_files": len(list(MIGRATIONS.glob('*.sql'))),
        "view_files": len(list(VIEWS.glob('*.sql'))),
        "seed_files": len(list(SEEDS.glob('*.sql'))),
        "tables": len(schemas),
    }
    print(json.dumps(result, indent=2))
    return 1 if errors else 0


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