
import collections
import sqlparse
import operator

import rad.console
from rad.parser import peek, consume
from rad.php import to_php_identifier, PHPString, PHPArray, PHPDict, PHPLiteral

T = sqlparse.tokens.Token

########################################
######### COMMON SQL FUNCTIONS #########
########################################
def get_name(tok):
    """ unescape an SQL name, e.g. convert `abc` to abc """
    if tok is None:
        return None
    v = tok.value
    if v[0] == "`" and v[-1] == "`":
        return v[1:-1]
    else:
        return v

def sql_build_string(s):
    return "'{}'".format(s.replace("'", "''"))

def parse_sql_value(type_hint=None, allow_null=False):
    if allow_null and peek(T.Keyword, "NULL"):
        consume()
        return NullValue()

    if type_hint is not None:
        tok = consume(type_hint)
    else:
        tok = consume()

    if tok.ttype in T.Literal.Number.Integer:
        return IntegerValue(int(tok.value))
    elif tok.ttype in T.Literal.String:
        return StringValue(sql_parse_string(tok))
    elif tok.ttype in T.Keyword and tok.value == "CURRENT_TIMESTAMP":
        return DateTimeValue("CURRENT_TIMESTAMP")
    elif tok.ttype in T.Keyword and tok.value == "CURRENT_DATE":
        return DateTimeValue("CURRENT_DATE")
    else:
        rad.console.warn()
        print(("Unrecognized value of type {}: {}."
                .format(tok.ttype, tok)))
        return LiteralValue(tok.value)
    # TODO: date and time default values

def sql_parse_string(token):
    val = token.value[1:-1]
    if token.ttype in T.Literal.String.Single:
        return val.replace("''", "'")
    elif token.ttype in T.Literal.String.Symbol:
        return val.replace('\\"', '"')
    else:
        raise ValueError("Unknown string token type {}".format(token.ttype))

########################################
########## COMMON SQL CLASSES ##########
########################################
class Table:
    def __init__(self, name, source_file=None):
        self.name = name
        self.columns = collections.OrderedDict()
        self.primary_key = None
        self.unique_keys = set()
        self.foreign_keys = set()
        self.source_file = source_file

    def add_column(self, column):
        if column.name in self.columns:
            raise ValueError("Column '{}' defined twice for table '{}'."
                    .format(column.name, self.name))
        self.columns[column.name] = column

    def has_column(self, name):
        return name in self.columns

    def check_columns(self, columns, description):
        for col in columns:
            if not self.has_column(col):
                raise ValueError(
                        "Unknown column '{}' encountered while {} for table '{}'."
                        .format(col, description, self.name))

    def set_primary_key(self, pk):
        if self.primary_key is not None:
            raise ValueError("Primary key specified twice for table '{}'."
                    .format(self.name))
        self.check_columns(pk, "adding a primary key")
        for col in pk:
            if self.columns[col].flags & Column.FLAG_NULL:
                raise ValueError("Primary key cannot contain NULL column '{}' for table '{}'.".format(col, self.name))
            self.columns[col].flags |= Column.FLAG_NOTNULL
        self.primary_key = pk

    def add_unique_key(self, uk):
        self.check_columns(uk, "adding a unique key")
        self.unique_keys.add(uk)

    def add_foreign_key(self, fk):
        self.check_columns(list(fk.columns.keys()), "adding a foreign key")
        self.foreign_keys.add(fk)

    def foreign_keys_sorted(self):
        return self.foreign_keys
        print(self.foreign_keys)
        if(len(self.foreign_keys)==0):
            return set()
        #return self.foreign_keys
        #print(list(self.foreign_keys)[0])
        return sorted(self.foreign_keys, key=operator.attrgetter('columns'))
        #return {k: self.foreign_keys[k] for k in sorted(self.foreign_keys)}sorted(self.foreign_keys, key=lambda fk: fk.columns)

    def is_unique(self, columns):
        assert len(columns) > 0
        if self.primary_key <= columns:
            return True
        for uk in self.unique_keys:
            if uk <= columns:
                return True
        return False

    def get_php_class(self):
        return "Table_" + to_php_identifier(self.name).capitalize()

    def to_sql(self, add_drop=False):
        lines = []

        # columns:
        for col in list(self.columns.values()):
            lines.append(col.to_sql())
        # primary key:
        if self.primary_key is not None and len(self.primary_key) > 0:
            lines.append("PRIMARY KEY (`" + "`, `".join(self.primary_key) + "`)")
        # unique key
        for uk in self.unique_keys:
            lines.append("UNIQUE KEY (`" + "`, `".join(uk) + "`)")
        # foreign keys
        for fk in self.foreign_keys_sorted():
            lines.append(fk.to_sql())
        if add_drop:
            sql = "DROP TABLE IF EXISTS `{}`;\n".format(self.name)
        else:
            sql = ""
        sql += "CREATE TABLE `" + self.name + "` (\n  "
        sql += ",\n  ".join(lines)
        sql += "\n) ENGINE=InnoDB DEFAULT CHARSET=utf8;"

        return sql

    def to_php(self, schema_classes):
        cls_name = self.get_php_class()

        lines = []
        lines.append("class {} extends R\\Table {{".format(cls_name))
        lines.append("    private static $instance = null;");
        lines.append("    public static function getInstance() {")
        lines.append("        if (self::$instance === null) {")
        lines.append("            self::$instance = new {}();".format(cls_name))
        lines.append("        }")
        lines.append("        return self::$instance;")
        lines.append("    }")
        lines.append("    private $foreignKeys;")
        lines.append("    protected function __construct() {")
        lines.append("       parent::__construct({}, {});".format(
            PHPString(self.name),
            PHPArray([col.to_php() for col in list(self.columns.values())],
                multiline=True)))
        lines.append("       $this->foreignKeys = {};".format(
            PHPArray([fk.to_php(schema_classes) for fk in self.foreign_keys_sorted()])))
        lines.append("    }")
        lines.append("    public function getForeignKeys() {")
        lines.append("        return $this->foreignKeys;")
        lines.append("    }")
        lines.append("    public function getPrimaryKey() {")

        pk_php = PHPArray([PHPString(col) for col in sorted(self.primary_key)])
        lines.append("        return {};".format(pk_php))
        lines.append("    }")

        for col in list(self.columns.values()):
            if col.flags & Column.FLAG_AUTO:
                auto_column = PHPLiteral(col.name)
                break
        else:
            auto_column = "null"
        lines.append("    public function getAutoColumn() {")
        lines.append("        return {};".format(auto_column))
        lines.append("    }")
        lines.append("}")

        return "\n".join(lines)

class Column:
    FLAG_NULL = 1
    FLAG_NOTNULL = 2
    FLAG_AUTO = 4

    def __init__(self, name, ctype):
        self.name = name
        self.type = ctype
        self.flags = 0
        self.default_value = None
    def to_sql(self):
        sql = []
        sql.append("`" + self.name + "`")
        sql.append(self.type.to_sql())
        if self.flags & Column.FLAG_NOTNULL:
            sql.append("NOT NULL")
        if self.flags & Column.FLAG_NULL:
            sql.append("NULL")
        if self.default_value is not None:
            sql.append("DEFAULT")
            sql.append(self.default_value.to_sql())
        if self.flags & Column.FLAG_AUTO:
            sql.append("AUTO_INCREMENT")
        return " ".join(sql)
    def to_php(self):
        if self.default_value is not None:
            default_php = self.default_value.to_php()
        else:
            default_php = "null"

        flags_php = []
        if self.flags & Column.FLAG_AUTO:
            flags_php.append("R\\Column::FLAG_AUTO")
        if self.flags & Column.FLAG_NULL:
            flags_php.append("R\\Column::FLAG_NULLABLE")
        return "new R\\Column({}, {}, {}, {})".format(
                PHPString(self.name), self.type.to_php(), default_php,
                "|".join(flags_php) if len(flags_php) > 0 else "0")


class ForeignKey:
    def __init__(self, table, columns, on_update, on_delete):
        self.table = table
        self.columns = columns
        self.on_update = on_update
        self.on_delete = on_delete
    def __repr__(self):
        return repr((self.table, self.columns, self.on_update, self.on_delete))
    def to_sql(self):
        sql = []
        (cols_from, cols_to) = list(zip(*list(self.columns.items())))
        sql.append("CONSTRAINT FOREIGN KEY (`{}`) REFERENCES `{}` (`{}`)"
                .format("`, `".join(cols_from), self.table, "`, `".join(cols_to)))
        if self.on_update is not None:
            sql.append("ON UPDATE " + self.on_update)
        if self.on_delete is not None:
            sql.append("ON DELETE " + self.on_delete)
        return " ".join(sql)
    def to_php(self, schema_classes):
        colmap = PHPDict()
        for (c1, c2) in list(self.columns.items()):
            colmap[c1] = PHPString(c2)
        return "new R\\ForeignKey({}, {}, {})".format(PHPString(self.table),
                PHPString(schema_classes[self.table]), colmap)

class NullValue:
    def get_value(self):
        return None
    def to_sql(self):
        return "NULL"
    def to_php(self):
        return "new R\\NullValue()"

class IntegerValue:
    def __init__(self, value):
        assert isinstance(value, int)
        self.value = value
    def get_value(self):
        return self.value
    def to_sql(self):
        return str(self.value)
    def to_php(self):
        return "new R\\ConstantValue({})".format(self.value)

class StringValue:
    def __init__(self, value):
        assert isinstance(value, str)
        self.value = value
    def get_value(self):
        return self.value
    def to_sql(self):
        return sql_build_string(self.value)
    def to_php(self):
        return "new R\\ConstantValue({})".format(PHPString(self.value))

class DateTimeValue:
    # TODO: work around MySQL not supporting CURRENT_DATE as a default value
    def __init__(self, name):
        assert name in ("CURRENT_TIMESTAMP", "CURRENT_DATE"), "Unknown DateTime type."
        self.name = name
    def get_value(self):
        return self
    def to_sql(self):
        return self.name
    def to_php(self):
        return "new R\\DateTimeValue(R\\DateTimeValue::{})".format(self.name)

class LiteralValue:
    def __init__(self, literal):
        self.literal = literal
    def to_sql(self):
        return self.literal
    def get_value(self):
        return self.literal
    def to_php(self):
        rad.console.warn()
        print(("Using unmodified SQL literal as PHP code: {}"
                .format(self.literal)))
        return "new R\\ConstantValue({})".format(self.literal)

class ColumnType:
    def __init__(self, name, args):
        self.name = name
        self.args = args
    def to_sql(self):
        if len(self.args) > 0:
            return "{}({})".format(self.name,
                    ",".join([val.to_sql() for val in self.args]))
        else:
            return self.name
    def to_php(self):
        args_php = [PHPLiteral(arg.get_value()) for arg in self.args]
        return "new R\\{}Type({})".format(self.name.capitalize(),
                ", ".join(map(str, args_php)))
