#!/usr/bin/python3


# parse arguments before importing the modules because networkx has a very slow
# startup time
import argparse
parser = argparse.ArgumentParser(description="Parse the database schema"
        " specified in SQL_SCHEMA and generate the PHP file OUTPUT"
        " containing one class for each table in SCHEMA.",
        formatter_class=argparse.ArgumentDefaultsHelpFormatter)

parser.add_argument("schema_files", nargs="+", metavar="SQL_FILE",
        help="path to one or more SQL files containing DDL statements;"
        " statements other than CREATE TABLE are ignored")
parser.add_argument("-o", "--output", dest="php_outfile", metavar="PHP_FILE",
        required=True, help="the generated PHP file is written to this path; if the file"
        " already exists, it is overwritten without further notice")
parser.add_argument("--ns-rad", dest="ns_rad", metavar="NAMESPACE",
        default="\\Cbikt\\Rad",
        help="absolute PHP namespace in which the Rad toolkit resides")
parser.add_argument("--ns-gen", dest="ns_generated", metavar="NAMESPACE",
        default="\\Cbikt\\Rad\\Generated\\Schema",
        help="absolute PHP namespace in which the generated classes will reside")
parser.add_argument("--sql-out", dest="sql_outfile", metavar="PATH",
        help="TODO")
parser.add_argument("--pickle-out", dest="pickle_outfile", metavar="PATH",
        help="TODO")

args = parser.parse_args()

import collections
import datetime
import io
import networkx as nx
import os.path
import pickle
import sqlparse
import sys
import traceback

import rad.console
from rad.parser import init_parser, peek, consume
from rad.sql import *

SERIALIZATION_FORMAT_VERSION = 5
T = sqlparse.tokens.Token

########################################
############ SCHEMA PARSER  ############
########################################
def parse_create_table(stmt, source_file):
    # grammar constructs
    def parse_value_list(type_hints=None, required=True, default=None,
            with_annotations=False):
        def parse_value(hint=None):
            values.append(parse_sql_value(hint))
            if with_annotations:
                if peek(T.Punctuation, ":"):
                    consume()
                    values.append(parse_sql_value(type_hint=T.Literal.String))
                else:
                    values.append(NullValue())

        if not required and not peek(T.Punctuation, "("):
            return default
        consume(T.Punctuation, "(")
        values = []
        if type_hints is not None:
            for hint in type_hints:
                parse_value(hint)
        else:
            while True:
                parse_value()
                if peek(T.Punctuation, ")"):
                    break
                consume(T.Punctuation, ",")
        consume(T.Punctuation, ")")

        return values

    def parse_type():
        tname = consume(T.Name, icmp=True).value.upper()
        if tname == "INT":
            args = parse_value_list(type_hints=[T.Literal.Number.Integer],
                    required=False, default=[IntegerValue(11)])
        elif tname == "TINYINT":
            args = parse_value_list(type_hints=[T.Literal.Number.Integer],
                    required=False, default=[IntegerValue(4)])
        elif tname == "BIGINT":
            args = parse_value_list(type_hints=[T.Literal.Number.Integer],
                    required=False, default=[IntegerValue(20)])
        elif tname == "VARCHAR" or tname == "CHAR":
            args = parse_value_list(type_hints=[T.Literal.Number.Integer],
                    required=True)
        elif tname == "ENUM":
            args = parse_value_list(required=True, with_annotations=True)
        elif tname in ("DATETIME", "DATE", "TEXT", "MEDIUMTEXT", "LONGTEXT", "BLOB", "MEDIUMBLOB", "LONGBLOB", "TIME", "JSON"):
            args = [] # has no arguments
        else:
            raise ValueError("Unknown type '{}'.".format(tname))

        return ColumnType(tname, args)

    def parse_column_definition():
        name = get_name(consume(T.Name))
        ctype = parse_type()
        column = Column(name, ctype)

        set_pk = False
        foreign_key = None
        while True:
            if peek(T.Keyword, "NULL"):
                consume()
                column.flags |= Column.FLAG_NULL
            elif peek(T.Keyword, "NOT NULL"):
                consume()
                column.flags |= Column.FLAG_NOTNULL
            elif peek(T.Keyword, "DEFAULT"):
                consume()
                column.default_value = parse_sql_value(allow_null=True)

                #nasty fix for on update
                if peek(T.Keyword, "ON"):
                    consume()
                    if peek(T.Keyword, "UPDATE"):
                        consume()
                        if peek(T.Keyword, "CURRENT_TIMESTAMP"):
                            consume()
            elif peek(T.Keyword, "AUTO_INCREMENT", icmp=True):
                consume()
                column.flags |= Column.FLAG_AUTO
            elif peek(T.Keyword, "UNIQUE"):
                consume()
                if peek(T.Keyword, "KEY"):
                    consume()
                table.add_unique_key(frozenset([column.name]))
            elif peek(T.Keyword, "PRIMARY"):
                consume()
                consume(T.Keyword, "KEY")
                set_pk = True
            elif peek(T.Keyword, "KEY"):
                consume() # discard
            elif peek(T.Keyword, "COMMENT"):
                consume()
                consume(T.Literal.String) # discard
            elif peek(T.Keyword, "REFERENCES"):
                foreign_key = parse_references((column.name,))
            else:
                break

        table.add_column(column)

        if set_pk:
            table.set_primary_key(frozenset([column.name]))
        if foreign_key is not None:
            table.add_foreign_key(foreign_key)

    def parse_column_list(allow_asc_desc=False):
        cols = []
        consume(T.Punctuation, "(")
        while True:
            cols.append(get_name(consume(T.Name)))
            if (allow_asc_desc and (peek(T.Keyword, "ASC")
                or peek(T.Keyword, "DESC"))):
                consume() # consume and ignore ASC or DESC suffix
            if not peek(T.Punctuation, ","):
                break
            consume()
        consume(T.Punctuation, ")")
        return tuple(cols)

    def parse_key(name_allowed=True):
        if name_allowed and peek(T.Name):
            name = get_name(consume())
        else:
            name = None
        cols = parse_column_list(allow_asc_desc=True)
        if peek(T.Keyword, "COMMENT"):
            consume()
            consume(T.Literal.String) # discard trailing comment
        return (name, cols)

    def parse_references(columns_left):
        def parse_restrict_option():
            if peek(T.Keyword, "RESTRICT"):
                consume()
                return "RESTRICT"
            elif peek(T.Keyword, "CASCADE"):
                consume()
                return "CASCADE"
            elif peek(T.Keyword, "SET"):
                consume()
                consume(T.Keyword, "NULL")
                return "SET NULL"
            elif peek(T.Keyword, "NO"):
                consume()
                consume(T.Name, "ACTION")
                return "NO ACTION"
            else:
                pprint(consume())
                raise ValueError("Invalid restrict action.")

        consume(T.Keyword, "REFERENCES")
        (tbl, cols) = parse_key()
        if tbl is None:
            raise ValueError("Foreign key must specify a table name.")
        on_update = None
        on_delete = None

        while True:
            if not peek(T.Keyword, "ON"):
                break
            consume()
            if peek(T.Keyword, "UPDATE"):
                consume()
                if on_update is not None:
                    raise ValueError("ON UPDATE specified twice.")
                on_update = parse_restrict_option()
            elif peek(T.Keyword, "DELETE"):
                consume()
                if on_delete is not None:
                    raise ValueError("ON DELETE specified twice.")
                on_delete = parse_restrict_option()
            else:
                consume(T.Keyword) # trigger error

        return ForeignKey(tbl, dict(list(zip(columns_left, cols))),
                on_update, on_delete)

    def parse_constraint():
        if peek(T.Keyword, "PRIMARY"):
            consume(T.Keyword, "PRIMARY")
            consume(T.Keyword, "KEY")
            if table.primary_key is not None:
                raise ValueError("Primary key specified twice for table {}."
                        .format(table.name))
            (_, pk) = parse_key(name_allowed=False)
            table.set_primary_key(frozenset(pk))
        elif peek(T.Keyword, "UNIQUE"):
            consume()
            if peek(T.Keyword, "KEY") or peek(T.Keyword, "INDEX"):
                consume() # discard
            (_, uk) = parse_key()
            table.add_unique_key(frozenset(uk))
        elif peek(T.Keyword, "FOREIGN"):
            consume()
            if peek(T.Keyword, "KEY"):
                consume() # discard
            (_, cols) = parse_key()
            table.add_foreign_key(parse_references(cols))
        elif peek(T.Keyword, "INDEX"):
            # ignore INDEX clauses
            consume()
            parse_key()
        else:
            return False
        return True

    # main
    init_parser(stmt)
    consume(T.Keyword.DDL, "CREATE")
    consume(T.Keyword, "TABLE")
    if peek(T.Keyword, "IF"):
        consume()
        consume(T.Keyword, "NOT")
        consume(T.Keyword, "EXISTS")
    name = get_name(consume(T.Name))
    table = Table(name, source_file)

    consume(T.Punctuation, "(")
    while True:
        if peek(T.Name):
            parse_column_definition()
        elif peek(T.Keyword, "KEY"):
            consume()
            parse_key() # discard
        elif peek(T.Keyword, "CONSTRAINT"):
            consume()
            consume(T.Name) # discard name
            if not parse_constraint():
                raise ValueError("Unknown constraint.")
        else:
            # will do nothing if we're not at the beginning of a constraint
            parse_constraint()

        if not peek(T.Punctuation, ","):
            break
        consume()
    consume(T.Punctuation, ")")

    auto_seen = False
    for col in list(table.columns.values()):
        # allow at most one AUTO column per table
        if col.flags & Column.FLAG_AUTO:
            if not auto_seen:
                auto_seen = True
            else:
                raise ValueError("Table '{}' has more than one AUTO column."
                        .format(table.name))

        # validate NULL/NOT NULL constraints of all columns
        if (col.flags & Column.FLAG_NULL and col.flags & Column.FLAG_NOTNULL):
            raise ValueError("Column '{}' cannot be both NULL and NOT NULL for table '{}'.".format(col.name, table.name))
        if (not (col.flags & Column.FLAG_NULL or col.flags & Column.FLAG_NOTNULL)):
            col.flags |= Column.FLAG_NULL # default constraint is NOT NULL
        if (col.flags & Column.FLAG_NOTNULL
                and isinstance(col.default_value, NullValue)):
            raise ValueError("Column '{}' of table '{}' cannot be declared NOT NULL and have default value NULL.".format(col.name, table.name))

    if table.primary_key is None:
        rad.console.warn()
        print(("Table {} has no primary key.".format(table.name)))
        table.set_primary_key(frozenset())

    return table

def parse_schema_files(paths):
    # parse source DDLs
    schema_unsorted = {}

    for path in paths:
        with open(path, "r") as f:
            for stmt in sqlparse.parsestream(f, encoding="utf-8"):
                if stmt.get_type() != "CREATE":
                    continue
                table = parse_create_table(stmt, path)
                #print("Parsed definition of table {}.".format(table.name))
                if table.name in schema_unsorted:
                    raise ValueError("{}: Cannot redeclare table '{}' previously declared in {}."
                            .format(path, table.name,
                                schema_unsorted[table.name].source_file))
                schema_unsorted[table.name] = table

    # TODO: AUTO column may not have a default value
    # TODO: validate foreign keys
    # TODO: forbid AUTO->AUTO foreign keys

    # topologically sort tables
    schema_sorted = collections.OrderedDict()
    G = nx.DiGraph()
    G.add_nodes_from(list(schema_unsorted.keys()))
    for table in list(schema_unsorted.values()):
        for fk in table.foreign_keys:
            G.add_edge(table.name, fk.table)
    for name in reversed(list(nx.topological_sort(G))):
        schema_sorted[name] = schema_unsorted[name]

    return schema_sorted

# parse schema:
try:
    schema = parse_schema_files(args.schema_files)
    schema_classes = {alias: args.ns_generated + "\\" + table.get_php_class()
                      for (alias, table) in list(schema.items())}
except:
    rad.console.error()
    traceback.print_exc(file=sys.stdout)
    sys.exit(3)

# generate files:
print(("Generating file {}...".format(args.php_outfile)))
with io.open(args.php_outfile, "w", encoding="utf-8") as f:
    f.write("<?php\n")
    f.write("// file generated by genschema.py on {}\n"
            .format(datetime.datetime.now()))
    f.write("// input file(s): {}\n".format(", ".join(args.schema_files)))
    f.write("// output file: {}\n".format(args.php_outfile))
    f.write("// complete command line: {}\n".format(" ".join(sys.argv)))
    f.write("\n")
    f.write("namespace {};\n".format(args.ns_generated.lstrip("\\")))
    f.write("use {} as R;\n".format(args.ns_rad))
    for table in list(schema.values()):
        f.write(table.to_php(schema_classes))
        f.write("\n\n")

if args.sql_outfile is not None:
    print(("Generating file {}...".format(args.sql_outfile)))
    with io.open(args.sql_outfile, "w", encoding="utf-8") as f:
        f.write("-- file generated by genschema.py on {}\n"
                .format(datetime.datetime.now()))
        f.write("-- input file(s): {}\n".format(", ".join(args.schema_files)))
        f.write("-- output file: {}\n".format(args.sql_outfile))
        f.write("-- complete command line: {}\n".format(" ".join(sys.argv)))
        f.write("\n")

        for table in list(schema.values()):
            f.write(table.to_sql())
            f.write("\n\n")

if args.pickle_outfile is not None:
    print(("Generating file {}...".format(args.pickle_outfile)))
    with open(args.pickle_outfile, "wb") as f:
        pickle.dump((schema, schema_classes, SERIALIZATION_FORMAT_VERSION), f)
