#!/usr/bin/env python3
"""
PostgreSQL to MySQL converter
Converts PostgreSQL database schema and data to MySQL format
"""

import argparse
import logging
import sys
from typing import Dict, List, Tuple

import psycopg2
import pymysql

# Configure logging
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)

# Data type mapping from PostgreSQL to MySQL
DATA_TYPE_MAPPING = {
    # Numeric types
    'smallint': 'SMALLINT',
    'integer': 'INT',
    'bigint': 'BIGINT',
    'decimal': 'DECIMAL',
    'numeric': 'DECIMAL',
    'real': 'FLOAT',
    'double precision': 'DOUBLE',
    'serial': 'INT AUTO_INCREMENT',
    'bigserial': 'BIGINT AUTO_INCREMENT',
    
    # String types
    'character varying': 'VARCHAR',
    'varchar': 'VARCHAR',
    'character': 'CHAR',
    'char': 'CHAR',
    'text': 'TEXT',
    
    # Date/time types
    'date': 'DATE',
    'time': 'TIME',
    'timestamp': 'TIMESTAMP',
    'timestamptz': 'TIMESTAMP',
    'interval': 'VARCHAR(255)',  # MySQL doesn't have interval type
    
    # Boolean type
    'boolean': 'BOOLEAN',
    
    # Binary types
    'bytea': 'BLOB',
    
    # JSON types
    'json': 'JSON',
    'jsonb': 'JSON',
    
    # UUID type
    'uuid': 'VARCHAR(36)',  # MySQL doesn't have native UUID type
}

class PostgreSQLToMySQLConverter:
    def __init__(self, psql_conn, mysql_conn):
        self.psql_conn = psql_conn
        self.mysql_conn = mysql_conn
        self.psql_cursor = psql_conn.cursor()
        self.mysql_cursor = mysql_conn.cursor()
    
    def get_tables(self) -> List[str]:
        """Get list of tables in PostgreSQL database"""
        query = """
        SELECT table_name 
        FROM information_schema.tables 
        WHERE table_schema = 'public' 
        AND table_type = 'BASE TABLE'
        """
        self.psql_cursor.execute(query)
        return [row[0] for row in self.psql_cursor.fetchall()]
    
    def get_table_schema(self, table_name: str) -> List[Dict]:
        """Get table schema from PostgreSQL"""
        query = """
        SELECT 
            column_name, 
            data_type, 
            character_maximum_length, 
            is_nullable, 
            column_default
        FROM information_schema.columns 
        WHERE table_schema = 'public' 
        AND table_name = %s
        ORDER BY ordinal_position
        """
        self.psql_cursor.execute(query, (table_name,))
        columns = []
        for row in self.psql_cursor.fetchall():
            column_name, data_type, max_length, is_nullable, default = row
            columns.append({
                'name': column_name,
                'data_type': data_type,
                'max_length': max_length,
                'is_nullable': is_nullable == 'YES',
                'default': default
            })
        return columns
    
    def convert_data_type(self, postgres_type: str, max_length: int) -> str:
        """Convert PostgreSQL data type to MySQL data type"""
        mysql_type = DATA_TYPE_MAPPING.get(postgres_type.lower(), 'VARCHAR(255)')
        
        # Handle types that need length specification
        if mysql_type in ['VARCHAR', 'CHAR'] and max_length:
            return f"{mysql_type}({max_length})"
        elif mysql_type == 'DECIMAL':
            return "DECIMAL(10,2)"  # Default precision
        else:
            return mysql_type
    
    def generate_create_table_sql(self, table_name: str, columns: List[Dict]) -> str:
        """Generate CREATE TABLE statement for MySQL"""
        column_defs = []
        primary_keys = []
        
        for column in columns:
            col_name = column['name']
            col_type = self.convert_data_type(column['data_type'], column['max_length'])
            nullable = 'NULL' if column['is_nullable'] else 'NOT NULL'
            
            # Handle default values
            default = ''
            if column['default']:
                # Handle serial types (auto increment)
                if 'serial' in column['data_type'].lower():
                    pass  # Already handled in data type mapping
                elif column['data_type'] == 'boolean':
                    if column['default'] == 'true':
                        default = 'DEFAULT TRUE'
                    elif column['default'] == 'false':
                        default = 'DEFAULT FALSE'
                elif 'nextval' in column['default']:
                    pass  # Skip PostgreSQL sequence syntax
                elif 'now()' in column['default'] or 'CURRENT_TIMESTAMP' in column['default']:
                    default = 'DEFAULT CURRENT_TIMESTAMP'
                else:
                    # Handle other date/time functions
                    if 'created_at' in column['name'] or 'updated_at' in column['name'] or 'marked_at' in column['name']:
                        default = 'DEFAULT CURRENT_TIMESTAMP'
                    else:
                        default = f"DEFAULT {column['default']}"
            
            # Check for primary key (simplified - assumes serial types are primary keys)
            if 'serial' in column['data_type'].lower():
                primary_keys.append(col_name)
            
            column_def = f"`{col_name}` {col_type} {nullable} {default}".strip()
            column_defs.append(column_def)
        
        # Add primary key constraint if found
        if primary_keys:
            pk_def = f"PRIMARY KEY ({', '.join([f'`{pk}`' for pk in primary_keys])})"
            column_defs.append(pk_def)
        
        create_sql = f"CREATE TABLE IF NOT EXISTS `{table_name}` (\n"
        create_sql += "  " + ",\n  ".join(column_defs)
        create_sql += "\n);"
        
        return create_sql
    
    def get_table_data(self, table_name: str, columns: List[Dict]) -> List[Tuple]:
        """Get data from PostgreSQL table"""
        column_names = [col['name'] for col in columns]
        query = f"SELECT {', '.join(column_names)} FROM {table_name}"
        self.psql_cursor.execute(query)
        return self.psql_cursor.fetchall()
    
    def generate_insert_sql(self, table_name: str, columns: List[Dict], data: List[Tuple]) -> List[str]:
        """Generate INSERT statements for MySQL"""
        if not data:
            return []
        
        column_names = [f"`{col['name']}`" for col in columns]
        column_list = ", ".join(column_names)
        
        insert_statements = []
        batch_size = 1000  # Insert in batches
        
        for i in range(0, len(data), batch_size):
            batch = data[i:i+batch_size]
            values_list = []
            
            for row in batch:
                values = []
                for j, value in enumerate(row):
                    col_type = columns[j]['data_type']
                    if value is None:
                        values.append('NULL')
                    elif col_type in ['text', 'character', 'character varying', 'varchar', 'char', 'date', 'time', 'timestamp', 'timestamptz', 'json', 'jsonb', 'uuid']:
                        # Escape single quotes
                        if isinstance(value, str):
                            escaped_value = value.replace("'", "''")
                            values.append(f"'{escaped_value}'")
                        else:
                            values.append(f"'{value}'")
                    elif col_type == 'boolean':
                        values.append('1' if value else '0')
                    else:
                        values.append(str(value))
                values_str = "(" + ", ".join(values) + ")"
                values_list.append(values_str)
            
            values_batch = ", ".join(values_list)
            insert_sql = f"INSERT INTO `{table_name}` ({column_list}) VALUES {values_batch};"
            insert_statements.append(insert_sql)
        
        return insert_statements
    
    def convert_table(self, table_name: str):
        """Convert a single table from PostgreSQL to MySQL"""
        logger.info(f"Converting table: {table_name}")
        
        # Get table schema
        columns = self.get_table_schema(table_name)
        
        # Generate CREATE TABLE statement
        create_sql = self.generate_create_table_sql(table_name, columns)
        logger.debug(f"CREATE TABLE SQL: {create_sql}")
        
        # Execute CREATE TABLE in MySQL
        try:
            self.mysql_cursor.execute(create_sql)
            self.mysql_conn.commit()
            logger.info(f"Created table: {table_name}")
        except Exception as e:
            logger.error(f"Error creating table {table_name}: {e}")
            self.mysql_conn.rollback()
            return False
        
        # Get table data
        data = self.get_table_data(table_name, columns)
        logger.info(f"Found {len(data)} rows in {table_name}")
        
        # Generate and execute INSERT statements
        if data:
            insert_statements = self.generate_insert_sql(table_name, columns, data)
            for i, insert_sql in enumerate(insert_statements):
                try:
                    self.mysql_cursor.execute(insert_sql)
                    self.mysql_conn.commit()
                    logger.debug(f"Inserted batch {i+1}/{len(insert_statements)} for {table_name}")
                except Exception as e:
                    logger.error(f"Error inserting data into {table_name}: {e}")
                    self.mysql_conn.rollback()
                    return False
        
        logger.info(f"Successfully converted table: {table_name}")
        return True
    
    def convert_all_tables(self):
        """Convert all tables from PostgreSQL to MySQL"""
        tables = self.get_tables()
        logger.info(f"Found {len(tables)} tables to convert")
        
        success_count = 0
        failure_count = 0
        
        for table in tables:
            if self.convert_table(table):
                success_count += 1
            else:
                failure_count += 1
        
        logger.info(f"Conversion completed: {success_count} succeeded, {failure_count} failed")
        return success_count, failure_count

def get_postgresql_connection(host, port, database, user, password):
    """Get PostgreSQL connection"""
    try:
        conn = psycopg2.connect(
            host=host,
            port=port,
            database=database,
            user=user,
            password=password
        )
        logger.info("Connected to PostgreSQL database")
        return conn
    except Exception as e:
        logger.error(f"Error connecting to PostgreSQL: {e}")
        sys.exit(1)

def get_mysql_connection(host, port, database, user, password):
    """Get MySQL connection"""
    try:
        conn = pymysql.connect(
            host=host,
            port=port,
            database=database,
            user=user,
            password=password,
            charset='utf8mb4',
            cursorclass=pymysql.cursors.DictCursor
        )
        logger.info("Connected to MySQL database")
        return conn
    except Exception as e:
        logger.error(f"Error connecting to MySQL: {e}")
        sys.exit(1)

def main():
    """Main function"""
    parser = argparse.ArgumentParser(description='Convert PostgreSQL database to MySQL')
    
    # PostgreSQL connection parameters
    parser.add_argument('--psql-host', default='localhost', help='PostgreSQL host')
    parser.add_argument('--psql-port', type=int, default=5432, help='PostgreSQL port')
    parser.add_argument('--psql-db', required=True, help='PostgreSQL database name')
    parser.add_argument('--psql-user', required=True, help='PostgreSQL username')
    parser.add_argument('--psql-password', required=True, help='PostgreSQL password')
    
    # MySQL connection parameters
    parser.add_argument('--mysql-host', default='localhost', help='MySQL host')
    parser.add_argument('--mysql-port', type=int, default=3306, help='MySQL port')
    parser.add_argument('--mysql-db', required=True, help='MySQL database name')
    parser.add_argument('--mysql-user', required=True, help='MySQL username')
    parser.add_argument('--mysql-password', required=True, help='MySQL password')
    
    # Log level
    parser.add_argument('--verbose', action='store_true', help='Enable verbose logging')
    
    args = parser.parse_args()
    
    # Set log level
    if args.verbose:
        logger.setLevel(logging.DEBUG)
    
    # Get database connections
    psql_conn = get_postgresql_connection(
        args.psql_host, args.psql_port, args.psql_db, args.psql_user, args.psql_password
    )
    
    mysql_conn = get_mysql_connection(
        args.mysql_host, args.mysql_port, args.mysql_db, args.mysql_user, args.mysql_password
    )
    
    # Create converter instance
    converter = PostgreSQLToMySQLConverter(psql_conn, mysql_conn)
    
    # Convert all tables
    success_count, failure_count = converter.convert_all_tables()
    
    # Close connections
    psql_conn.close()
    mysql_conn.close()
    
    logger.info("Conversion process completed")
    return 0 if failure_count == 0 else 1

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

使用说明

python psql2mysql.py --psql-host localhost --psql-port 5432 --psql-db postgres_db --psql-user postgres_user --psql-password postgres_password --mysql-host localhost --mysql-port 3306 --mysql-db mysql_db --mysql-user mysql_user --mysql-password mysql_password

参数说明

- --psql-host : PostgreSQL主机地址(默认:localhost)
- --psql-port : PostgreSQL端口(默认:5432)
- --psql-db : PostgreSQL数据库名称(必需)
- --psql-user : PostgreSQL用户名(必需)
- --psql-password : PostgreSQL密码(必需)
- --mysql-host : MySQL主机地址(默认:localhost)
- --mysql-port : MySQL端口(默认:3306)
- --mysql-db : MySQL数据库名称(必需)
- --mysql-user : MySQL用户名(必需)
- --mysql-password : MySQL密码(必需)
- --verbose : 启用详细日志记录

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐