from clickhouse_driver import Client
from datetime import datetime
import time
import sys
import logging
import MySQLdb

logger = logging.getLogger(__name__)
max_concurrent_queries_for_user = 200
settings = 'SETTINGS max_concurrent_queries_for_user=200,max_memory_usage=20000000000,lock_acquire_timeout=360;'


def get_first_master_host(hosts_file='/etc/hosts'):
    # 集群 master 机器名随集群不同而不同，从 /etc/hosts 解析出所有 master 节点，
    # 取第一台。命名带序号(如 master-xxx-1)时按序号升序取最小；
    # 命名不带序号(如 master-xxx)时也兼容，按主机名字典序兜底排序。
    masters = []
    with open(hosts_file) as f:
        for line in f:
            line = line.strip()
            if not line or line.startswith('#'):
                continue
            tokens = line.split()
            for name in tokens[1:]:
                if name == 'master' or name.startswith('master-'):
                    suffix = name.rsplit('-', 1)[1] if '-' in name else ''
                    index = int(suffix) if suffix.isdigit() else None
                    masters.append((index, name))
                    break
    if not masters:
        raise Exception('No master host found in {}'.format(hosts_file))
    # 有序号的优先按序号升序；无序号的排在其后，按主机名字典序，保证结果确定。
    masters.sort(key=lambda item: (item[0] is None, item[0] if item[0] is not None else 0, item[1]))
    return masters[0][1]

class Migrator(object):
    def __init__(self, job_id, task_id, task_size):
        self.job_id = job_id
        self.task_id = task_id
        self.task_size = task_size
        self.src_client = None
        self.dst_client = None
        self.src_host = ''
        self.src_port = 9000
        self.dst_host = ''
        self.dst_port = 9000
        self.src_user = ''
        self.src_passwd = ''
        self.dst_user = ''
        self.dst_passwd = ''
        self.meta = MySQLdb.connect(get_first_master_host(), 'clickhouse', '*Clickhouse123', 'clickhouse', charset='utf8mb4')

    def get_job_info(self):
        # sql = 'select src_host,src_port,src_user,src_password,dst_host,dst_port,dst_user,dst_password from migrator.jobs where '
        cursor = self.meta.cursor()
        cursor.execute(str.format('select src_host,src_port,src_user,src_password,dst_host,dst_port,dst_user,dst_password from migrator.jobs where job_id={}', self.job_id))
        infos = cursor.fetchall()
        
        self.src_host = infos[0][0]
        self.src_port = infos[0][1]
        self.src_user = infos[0][2]
        self.src_passwd = infos[0][3]
        self.dst_host = infos[0][4]
        self.dst_port = infos[0][5]
        self.dst_user = infos[0][6]
        self.dst_passwd = infos[0][7]

        self.src_client = Client(host=self.src_host, port=self.src_port, user=self.src_user, password=self.src_passwd)
        self.set_config(self.src_client, 'max_concurrent_queries_for_user', max_concurrent_queries_for_user)
        self.dst_client = Client(host=self.dst_host, port=self.dst_port, user=self.dst_user, password=self.dst_passwd)
        for info in infos:
            print(info)

    def open_clickhouse_instance(self, host_, port_, user_, passwd_):
        client = Client(host=host_, port=port_, user=user_, password=passwd_)
        return client
    
    def open_source_db_instance(self, host_, port_, user_, passwd_):
        self.src_client = self.open_clickhouse_instance(host_, port_, user_, passwd_)

    # def open_destination_db_instance(self, host_, port_, user_, passwd_):
    #     self.dst_client = self.open_clickhouse_instance(host_, port_, user_, passwd_)

    def set_config(self, client, key, value):
        query = str.format("SET {}={}", key, value)
        self.execute_statement(client, query)

    def execute_statement(self, client, statement):
        if client is None:
            raise Exception("The connection is not established!")
        if statement is None:
            raise Exception("The statement not specified!")
        return client.execute(statement)

    def get_dbs_from_instance(self, client):
        query = '''SHOW DATABASES'''
        return self.execute_statement(client, query)

    def get_create_db_statement(self, client, db):
        query = str.format('''SHOW CREATE DATABASE '{}' ''', db)
        return self.execute_statement(client, query)
    
    def get_tables_from_db(self, client, db):
        query = str.format('''SHOW TABLES FROM '{}' ''', db)
        return self.execute_statement(client, query)

    def get_create_table_statement(self, client, db, table):
        query = str.format('''SHOW CREATE TABLE '{}.{}' ''', db, table)
        return self.execute_statement(client, query)

    def get_partition_key(self, client, db, table):
        query = str.format('''SELECT partition_key FROM system.tables WHERE database='{}' and name='{}' ''', db, table)
        return self.execute_statement(client, query)

    def get_partiton_rows_and_bytes(self, client, db, table):
        query = str.format('''SELECT distinct(partition), sum(rows), sum(bytes_on_disk) FROM system.parts WHERE database='{}' and table='{}' GROUP BY partition''', db, table)
        return self.execute_statement(client, query)

    def get_databases_from_instance(self, client, filter=''' 'system','INFORMATION_SCHEMA','information_schema','_temporary_and_external_tables' '''):
        query = str.format('''SELECT name FROM system.databases WHERE name NOT IN ({})''', filter)
        return self.execute_statement(self.src_client, query)

    def get_tables_from_database(self, client, db, filter=''):
        query = ''
        if filter == '':
            query = str.format('''SELECT name,partition_key,create_table_query from system.tables WHERE database='{}' ''', db)
        else:
            query = str.format('''SELECT name,partition_key,create_table_query from system.tables WHERE database='{}' AND name NOT IN ({})''', db, filter)
        logger.info("The create query {query}")
        print("The create query:{}", query)
        return self.execute_statement(client, query)

    def get_partitons_from_table(self, client, db, table, condition=""):
        filter_clause = ""
        if len(condition) != 0:
            filter_clause = " AND " + condition
        query = str.format('''SELECT partition,sum(rows),sum(bytes_on_disk) FROM system.parts WHERE active=1 AND database='{}' AND table='{}' {} GROUP BY partition''', db, table, filter_clause)

        return self.execute_statement(client, query)

    def preserve_partition_infos(self, src_db, src_table, dst_db, dst_table, partition_keys, partitions):
        cursor = self.meta.cursor()
        dt = str(datetime.now())
        sql = '''insert into migrator.partitions (job_id,src_db,src_table,dst_db,dst_table,status,create_time,partition_keys,partition_id,rows,bytes) values '''
        db_table = (self.job_id, src_db, src_table, dst_db, dst_table, 'ready', dt, partition_keys)
        values = []
        for partition in partitions:
            values.append(str(db_table + partition))
        sql += ','.join(values)
        try:
            # cursor.execute(str.format('delete from migrator.partitions where job_id={}', self.job_id))
            cursor.execute(sql)
            self.meta.commit()
        except Exception as e:
            self.meta.rollback()
            print("Encounter an error when preserve partitions src_db:{}, dst_db:{}, src_table:{}, dst_table:{}, partition size:{}, will rollback this preserve!",
             src_db, dst_db, src_table, dst_table, len(partitions))
            print(e)
            raise(e)


        # print(values)
            
        # cursor.execute('''show tables from migrator''')
        # for row in cursor.fetchall():
        #     print(row)
    def get_doing_or_failed_partitions(self, src_db, src_table, dst_db, dst_table):
        sql = str.format('''select partition_keys,partition_id,rows,bytes from migrator.partitions where job_id={} and mod(id, {})={} and src_db='{}' and src_table='{}' and dst_db='{}' and dst_table='{}' and (status='doing' or status='failed') ''',
            self.job_id, self.task_size, self.task_id, src_db, src_table, dst_db, dst_table)
        cursor = self.meta.cursor()
        cursor.execute(sql)
        return cursor.fetchall()

    def get_ready_partitions(self, src_db, src_table, dst_db, dst_table):
        sql = str.format('''select partition_keys,partition_id,rows,bytes from migrator.partitions where job_id={} and mod(id, {})={} and src_db='{}' and src_table='{}' and dst_db='{}' and dst_table='{}' and status='ready' ''',
            self.job_id, self.task_size, self.task_id, src_db, src_table, dst_db, dst_table)
        cursor = self.meta.cursor()
        cursor.execute(sql)
        return cursor.fetchall()

    def get_invalid_partitions(self, src_db, src_table, dst_db, dst_table):
        sql = str.format('''select partition_keys,partition_id,rows,bytes from migrator.partitions where job_id={} and mod(id, {})={} and src_db='{}' and src_table='{}' and dst_db='{}' and dst_table='{}' and status='invalid' ''',
            self.job_id, self.task_size, self.task_id, src_db, src_table, dst_db, dst_table)
        cursor = self.meta.cursor()
        cursor.execute(sql)
        return cursor.fetchall()

    def get_done_partitions(self, src_db, src_table, dst_db, dst_table):
        sql = str.format('''select partition_keys,partition_id,rows,bytes from migrator.partitions where job_id={} and mod(id, {})={} and src_db='{}' and src_table='{}' and dst_db='{}' and dst_table='{}' and status='done' ''',
            self.job_id, self.task_size, self.task_id, src_db, src_table, dst_db, dst_table)
        cursor = self.meta.cursor()
        cursor.execute(sql)
        return cursor.fetchall()

    def update_partition_status(self, src_db, src_table, dst_db, dst_table, partition_id, status):
        sql = str.format('''update migrator.partitions set status='{}' where job_id={} and mod(id, {})={} and src_db='{}' and src_table='{}' and dst_db='{}' and dst_table='{}' and partition_id="{}" ''',
            status, self.job_id, self.task_size, self.task_id, src_db, src_table, dst_db, dst_table, partition_id)
        cursor = self.meta.cursor()
        try:
            cursor.execute(sql)
            self.meta.commit()
        except Exception as e:
            self.meta.rollback()
            print('Encounter an error when update partition status src_db:{}, dst_db:{}, src_table:{}, dst_table:{}, partition id:{}, will rollback this update!',
                src_db, dst_db, src_table, dst_table, partition_id)
            print(e)

    def migrate_partition(self, src_db, src_table, dst_db, dst_table, partition_key, partition_value):
        partition_value = partition_value if ('(' in partition_key and ')' in partition_key) else str.format("'{}'", partition_value)
        where_clause = "" if partition_key == '' else str.format('''WHERE {}={} ''', partition_key, partition_value)
        remote_query = str.format('''SELECT * FROM remote('{}:{}', '{}.{}', '{}', '{}') {} {}''', self.src_host, self.src_port, src_db, src_table, self.src_user, self.src_passwd, where_clause, settings)
        print(remote_query)
        # ret = self.execute_statement(self.dst_client, remote_query)
        # print(ret)
        insert_query = str.format('''INSERT INTO {}.{} {}''', dst_db, dst_table, remote_query)
        print(insert_query)
        self.execute_statement(self.dst_client, insert_query)

    def drop_partition(self, client, db, table, partition_value):
        partition_value =  partition_value if ('(' in partition_value and ')' in partition_value) else str.format(''' '{}' ''', partition_value)
        query = str.format('''ALTER TABLE {}.{} DROP PARTITION {} {}''', db, table, partition_value, settings)
        print(query)
        self.execute_statement(client, query)

    def check_partition(self, src_db, src_table, dst_db, dst_table, partition_key, partition_value):
        partition_value = partition_value if ('(' in partition_key and ')' in partition_key) else str.format("'{}'", partition_value)
        where_clause = "" if partition_key == '' else str.format('''WHERE {}={} ''', partition_key, partition_value)
        remote_query = str.format('''SELECT count() FROM remote('{}:{}', '{}.{}', '{}', '{}') {} {}''', self.src_host, self.src_port, src_db, src_table, self.src_user, self.src_passwd, where_clause, settings)
        local_query = str.format('''SELECT count() FROM {}.{} {} {}''', dst_db, dst_table, where_clause, settings)
        remote_ret = self.execute_statement(self.dst_client, remote_query)
        local_ret = self.execute_statement(self.dst_client, local_query)
        print(str.format('''Check result: src->{}.{}.{}:{}, dst->{}.{}.{}:{}''', src_db, src_table, partition_value, remote_ret[0][0], dst_db, dst_table, partition_value, local_ret[0][0]))
        return (remote_ret[0][0], local_ret[0][0])

    def migrate_table_failed_data(self, src_db, src_table, dst_db, dst_table):
        partitions = self.get_doing_or_failed_partitions(src_db, src_table, dst_db, dst_table)
        print(str.format("Total {} failed partitions for table {}.{} to migration", len(partitions), src_db,src_table))
        for partition in partitions:
            try:
                status = 'done'
                start_time = time.time()
                self.drop_partition(self.dst_client, dst_db, dst_table, partition[1])
                self.migrate_partition(src_db, src_table, dst_db, dst_table, partition[0], partition[1])
                self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                #self.check_partition(dst_db, dst_table, partition)
                end_time = time.time()
                avg_speed = partition[3] / (end_time-start_time)
                print("Current speed is " + str(avg_speed))
            except Exception as e:
                status = 'failed'
                print("Encounter an error when migrator partition src_db:{}, src_table:{}, dst_db:{}, dst_table:{}, paritition_id:{}",
                src_db, src_table, dst_db, dst_table, partition[1])
                self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                print(e)
                raise(e)
    
    def migrate_table_data(self, src_db, src_table, dst_db, dst_table):
        partitions = self.get_ready_partitions(src_db, src_table, dst_db, dst_table)
        print(str.format("Total {} partitions for table {}.{} to migration", len(partitions), src_db,src_table))
        for partition in partitions:
            try:
                start_time = time.time()
                status = 'doing'
                # self.drop_partition(self.dst_client, dst_db, dst_table, partition[1])
                self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                ret = self.check_partition(src_db, src_table, dst_db, dst_table, partition[0], partition[1])
                if ret[0] != ret[1]:
                    if ret[1] != 0:
                        print(str.format('''Partition check inconsistent and drop it {}.{}.{}''', dst_db, dst_table, partition[1]))
                        self.drop_partition(self.dst_client, dst_db, dst_table, partition[1])
                    self.migrate_partition(src_db, src_table, dst_db, dst_table, partition[0], partition[1])
                    status = 'done'
                    self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                    #self.check_partition(dst_db, dst_table, partition)
                    end_time = time.time()
                    avg_speed = partition[3] / (end_time-start_time)
                    print("Current speed is " + str(avg_speed))
                else:
                    print(str.format('''Partition check consistent and update it's status to 'checked' '''))
                    status = 'checked'
                    self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
            except Exception as e:
                status = 'failed'
                print("Encounter an error when migrator partition src_db:{}, src_table:{}, dst_db:{}, dst_table:{}, paritition_id:{}",
                src_db, src_table, dst_db, dst_table, partition[1])
                self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                print(e)
                raise(e)

    def check_table_data(self, src_db, src_table, dst_db, dst_table):
        # ready_partitions = self.get_ready_partitions(src_db, src_table, dst_db, dst_table)
        # invalid_partitions = self.get_invalid_partitions(src_db, src_table, dst_db, dst_table)
        # partitions = ready_partitions + invalid_partitions
        partitions = self.get_done_partitions(src_db, src_table, dst_db, dst_table)
        print(str.format("Total {} partitions for table {}.{} to check", len(partitions), src_db,src_table))
        for partition in partitions:
            try:
                start_time = time.time()
                ret = self.check_partition(src_db, src_table, dst_db, dst_table, partition[0], partition[1])
                if ret[0] != ret[1] and ret[1] != 0:
                    print(str.format('''Partition check inconsistent and drop it {}.{}.{}''', dst_db, dst_table, partition[1]))
                    # self.drop_partition(self.dst_client, dst_db, dst_table, partition[1])
                    status = 'failed'
                    self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                elif ret[0] == ret[1]:
                    status = 'checked'
                    print(str.format('''Partition check consistent and update it's status to 'checked' '''))
                    self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                elif ret[1] == 0:
                    status = 'ready'
                    print(str.format('''This partition {}.{}.{} not been migrated, will skip this partition in next check''', dst_db, dst_table, partition[1]))
                    self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                # status = 'passed'
                # self.drop_partition(self.dst_client, dst_db, dst_table, partition[1])
                # self.migrate_partition(src_db, src_table, dst_db, dst_table, partition[0], partition[1])
                # self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                #self.check_partition(dst_db, dst_table, partition)
                end_time = time.time()
                avg_speed = partition[3] / (end_time-start_time)
                print("Current speed is " + str(avg_speed))
            except Exception as e:
                status = 'invalid'
                print("Encounter an error when check partition src_db:{}, src_table:{}, dst_db:{}, dst_table:{}, paritition_id:{}",
                src_db, src_table, dst_db, dst_table, partition[1])
                print(e)
                self.update_partition_status(src_db, src_table, dst_db, dst_table, partition[1], status)
                raise(e)

    def check_db_data(self, src_db, dst_db):
        tables = self.get_tables_from_database(self.dst_client, dst_db)
        for table in tables:
            self.check_table_data(src_db, table[0], dst_db, table[0])

    def migrate_db_data(self, src_db, dst_db):
        tables = self.get_tables_from_database(self.dst_client, dst_db)
        for table in tables:
            self.migrate_table_data(src_db, table[0], dst_db, table[0])

    def migrate_db_failed_data(self, src_db, dst_db):
        tables = self.get_tables_from_database(self.dst_client, dst_db)
        for table in tables:
            self.migrate_table_failed_data(src_db, table[0], dst_db, table[0])

    def migrate_data(self):
        dbs = self.get_databases_from_instance(self.dst_client)
        for db in dbs:
            self.migrate_db_data(db[0], db[0])
        for db in dbs:
            self.migrate_db_failed_data(db[0], db[0])

    def check_data(self):
        dbs = self.get_databases_from_instance(self.dst_client)
        for db in dbs:
            self.check_db_data(db[0], db[0])

    def migrate_tables_schema(self, src_db, dst_db):
        tables = self.get_tables_from_database(self.src_client, src_db)
        for table in tables:
            if table[0].startswith('.inner.'):
                continue
            create_query = ''
            print(table[2])
            if type(table[2]) != str:
                create_query = str(table[2], encoding='Latin-1')
            else:
                create_query = table[2]

            if create_query.startswith('CREATE MATERIALIZED VIEW'):
                print("Will skip this table: " + table[0])
                continue
                
            #print(table[0] + " create query is:" + create_query)
            #create_query = self.execute_statement(self.src_client, str.format('''SHOW CREATE TABLE {}.{} ''', src_db, table[0]))
            #print(str(table[2]).replace('CREATE TABLE', 'CREATE TABLE IF NOT EXISTS', 1))
            #print(create_query)
            #query = str(create_query[0][0], 'UTF-8')
            print(create_query.replace('CREATE TABLE', 'CREATE TABLE IF NOT EXISTS', 1).
                                                                 replace('CREATE VIEW', 'CREATE VIEW IF NOT EXISTS', 1).
                                                                 replace('CREATE MATERIALIZED VIEW', 'CREATE MATERIALIZED VIEW IF NOT EXISTS', 1))
            #print(query)
            self.execute_statement(self.dst_client, create_query.replace('CREATE TABLE', 'CREATE TABLE IF NOT EXISTS', 1).
                                                                 replace('CREATE VIEW', 'CREATE VIEW IF NOT EXISTS', 1).
                                                                 replace('CREATE MATERIALIZED VIEW', 'CREATE MATERIALIZED VIEW IF NOT EXISTS', 1))
            print('Create done')
            #self.execute_statement(self.dst_client, query.replace('CREATE TABLE', 'CREATE TABLE IF NOT EXISTS', 1))
            partitions = self.get_partitons_from_table(self.src_client, src_db, table[0])
            partition_keys = self.get_partition_key(self.src_client, src_db, table[0])[0][0]
            # partition_key = '' if partition_keys is None or len(partition_keys)==0 else partition_keys[0]
            print(str.format("Will preserve {} partitions (include duplicated partitions) for table {}.{}", len(partitions), src_db, table[0]))
            self.preserve_partitions(src_db, table[0], dst_db, table[0], partition_keys, partitions)

            # Not need anymore
            # incremental_parttitions = self.get_incremental_partitions(src_db, table[0], dst_db, table[0])
            # print(str.format("Will preserve {} incremental partitions for table {}.{}", len(incremental_parttitions), src_db, table[0]))
            # self.preserve_partitions(src_db, table[0], dst_db, table[0], partition_keys, incremental_parttitions)
            

    def migrate_database_schema(self):
        dbs = self.get_databases_from_instance(self.src_client)
        for db in dbs:
            query = str.format('''SHOW CREATE DATABASE {}''', db[0])
            # print(query)
            statement = self.execute_statement(self.src_client, query)
            # print(statement)
            print("Will migrate database schema for %s", db[0])
            create_statement = statement[0][0].replace('CREATE DATABASE', 'CREATE DATABASE IF NOT EXISTS', 1)
            self.execute_statement(self.dst_client, create_statement)
            print("Migrate database schema for %s done", db[0])
        return dbs   
            # print(ret)

    def drop_schemas(self):
        dbs = self.get_databases_from_instance(self.dst_client)
        for db in dbs:
            query = str.format('''DROP DATABASE IF EXISTS {}''', db[0])
            print("Will drop database: " + db[0])
            self.execute_statement(self.dst_client, query)
        print("Drop all schema done!")

    def migrate_schema(self):
        print("Will migrator all schemas!")
        dbs = self.migrate_database_schema()
        for db in dbs:
            self.migrate_tables_schema(db[0], db[0])
        print("Migrator all shcema done!")

    # def get_incremental_partitions(self, src_db, src_table, dst_db, dst_table):
    #     sql = str.format('''select max(partition_id),partition_keys from migrator.partitions where job_id={} and src_db='{}' and src_table='{}' and dst_db='{}' and dst_table='{}' ''',
    #         self.job_id, src_db, src_table, dst_db, dst_table)
    #     print(sql)
    #     cursor = self.meta.cursor()
    #     cursor.execute(sql)
    #     last_partition = cursor.fetchone()
    #     last_partition_id = last_partition[0]
    #     partition_keys = last_partition[1]
    #     increment_conditions = ""
    #     if last_partition_id is not None:
    #         increment_conditions = str.format(" partition > {} ", last_partition_id) if ('(' in last_partition_id and ')' in last_partition_id) else str.format(" partition > '{}' ", last_partition_id)
    #     print(str.format('''Conditions for get {}.{} incrmental partition is {}''', src_db, src_table, increment_conditions))
    #     return self.get_partitons_from_table(self.src_client, src_db, src_table, increment_conditions)


    def preserve_partitions(self, src_db, src_table, dst_db, dst_table, partition_keys, partitions):
        if len(partitions) == 0:
            return
        cursor = self.meta.cursor()
        dt = str(datetime.now())
        try:
            for partition in partitions:
                values = str((self.job_id, src_db, src_table, dst_db, dst_table, 'ready', dt, partition_keys) + partition)
                sql = str.format('''insert into migrator.partitions (job_id,src_db,src_table,dst_db,dst_table,status,create_time,partition_keys,partition_id,rows,bytes) values{}''' \
                    ''' on duplicate key update status='ready',rows={},bytes={},update_time='{}' ''',
                    values, partition[1], partition[2], dt)
                cursor.execute(sql)
            self.meta.commit()
        except Exception as e:
            self.meta.rollback()
            print('Execute sql error:' + sql)
            print(e)

    def preserve_incremental_partitions(self, src_db, src_table, dst_db, dst_table, partition_increments):
        partition_keys = self.get_partition_key(self.src_client, src_db, src_table)[0][0]
        self.preserve_partitions(src_db, src_table, dst_db, dst_table, partition_keys, partition_increments)
        
    def migate_materialized_view(self):
        query = '''SELECT create_table_query FROM system.tables where engine='MaterializedView' AND database!='system' '''
        mvs = self.execute_statement(self.src_client, query)
        for mv in mvs:
            print("Will migrate materialized view " + mv[0])
            self.execute_statement(self.dst_client, mv[0].replace('CREATE MATERIALIZED VIEW', 'CREATE MATERIALIZED VIEW IF NOT EXISTS', 1))
            print("Create materialized view done!")
        print("Migrate all materialized view done!")


# def get_dbs_from_instance(host_, port_, user_, passwd_):
#     pass

if __name__ == '__main__':
    if len(sys.argv) != 5:
        print('Error! wrong parameters number usage: python ck_migrator.py hostname taskID taskSize runType')
        exit(1)
    job_id = int(sys.argv[1].split('-')[2])
    task_id = int(sys.argv[2])
    task_size = int(sys.argv[3])
    run_type = sys.argv[4]
    migrator = Migrator(job_id, task_id, task_size)
    run_info = str.format("Will start job with job_id:{}, task_id:{}", job_id, task_id)
    print(run_info)
    # client = migrator.open_clickhouse_instance("127.0.0.1", 9000, "default", "")
    # ret = migrator.get_partiton_rows_and_bytes(client, "default", "store_sales_part")
    # print(ret)
    # migrator.open_source_db_instance("127.0.0.1", 9000, "default", "")
    # migrator.open_destination_db_instance("yq01-bdg-jarvis07.yq01.baidu.com", 9000, "default", "")
    migrator.get_job_info()
    if sys.argv[4] == 'migration':
        migrator.migrate_schema()
        migrator.migrate_data()
    elif sys.argv[4] == 'migration_schema':
        migrator.migrate_schema()
    elif sys.argv[4] == 'migration_data':
        migrator.migrate_data()
    elif sys.argv[4] == 'check':
        migrator.check_data()
    elif sys.argv[4] == 'drop_schema':
        migrator.drop_schemas()
    elif sys.argv[4] == 'migration_mv':
        migrator.migate_materialized_view()
    else:
        print("Unknown run type:" + sys.argv[4])
        exit(1)
    # partitions = migrator.get_partitons_from_table(migrator.src_client, "default", "store_sales_part")
    # print(partitions)
    # migrator.get_job_info()
    # migrator.preserve_partition_infos('default', 'store_sales_part', 'default', 'store_sales_part', partitions)
    # migrator.migrate_tables_schema('default')
    #migrator.migrate_schema()
    # migrator.migrate_partition('default','store_sales_part','default','store_sales_part','')
    # partitions = self.get_ready_or_failed_partitions(src_db, src_table, dst_db, dst_table)
    # migrator.migrator_data()
    # partitions_inc = migrator.get_incremental_partitions('default', 'store_sales_part', 'default', 'store_sales_part')
    # migrator.preserve_incremental_partitions('default', 'store_sales_part', 'default', 'store_sales_part', partitions_inc)


