#!/usr/bin/python3

import sqlite3
import re
import os
from typing import List, Tuple, Dict
import argparse


def read_symbols_from_db(db_path: str) -> Dict[Tuple[str, str], List[str]]:
    """
    读取数据库中 symbols 表的 lib_name、deb_name 和 symbol 数据。
    按空格分隔 symbol 的内容，将每个分隔后的值作为单独的元素。
    返回数据结构：{(lib_name, deb_name): [symbol_element_list]}。
    """
    conn = sqlite3.connect(db_path)
    cursor = conn.cursor()

    # 正则模式：匹配文件名中的版本号
    version_pattern = re.compile(r"(.*?)(?:\.so(?:\.\d+)*)$")

    # 读取 symbols 表的数据
    query = "SELECT lib_name, deb_name, symbol FROM symbols;"
    cursor.execute(query)
    rows = cursor.fetchall()

    # 组织成字典结构，按空格分隔 symbol
    data = {}
    for lib_name, deb_name, symbol in rows:
        match = version_pattern.match(os.path.basename(lib_name))
        if match:
            normalized_lib_name = match.group(1) # 提取主名称部分
        else:
            normalized_lib_name = lib_name # 若无法匹配，使用原名

        key = (normalized_lib_name, deb_name)
        if key not in data:
            data[key] = []
        # 按空格分隔 symbol 并添加到列表
        data[key].extend(symbol.split())

    conn.close()
    return data

def compare_symbols(
    db1_data: Dict[Tuple[str, str], List[str]],
    db2_data: Dict[Tuple[str, str], List[str]]
) -> List[Tuple[str, str, str, str]]:
    """
    对比两个数据库的 symbols 数据。
    返回差异数据：[(deb_name, lib_name, symbol, type)]。
    """
    results = []

    all_keys = set(db1_data.keys()).union(db2_data.keys())
    for key in all_keys:
        lib_name, deb_name = key
        symbols1 = db1_data.get(key, [])
        symbols2 = db2_data.get(key, [])

        # 新增和删除的 symbols
        added = [symbol for symbol in symbols2 if symbol not in symbols1]
        removed = [symbol for symbol in symbols1 if symbol not in symbols2]

        # 记录新增和删除的数据
        # results.extend([(deb_name, lib_name, symbol, "Added") for symbol in added])
        # results.extend([(deb_name, lib_name, symbol, "Removed") for symbol in removed])

        # 将符号列表转为字符串，并添加到结果中
        if added:
            results.append((deb_name, lib_name, " ".join(added), "Added"))
        if removed:
            results.append((deb_name, lib_name, " ".join(removed), "Removed"))

    return results

def save_differences_to_db(db_path: str, differences: List[Tuple[str, str, str, str]]):
    """
    将对比结果保存到数据库中的 symbol_differences 表。
    """
    conn = sqlite3.connect(db_path)
    cursor = conn.cursor()

    # 创建 symbol_differences 表（如果不存在）
    cursor.execute("""
    CREATE TABLE IF NOT EXISTS symbol_differences (
        id INTEGER PRIMARY KEY AUTOINCREMENT,
        deb_name TEXT NOT NULL,
        lib_name TEXT NOT NULL,
        symbol TEXT NOT NULL,
        type TEXT NOT NULL
    );
    """)

    # 插入差异数据
    cursor.executemany("""
    INSERT INTO symbol_differences (deb_name, lib_name, symbol, type)
    VALUES (?, ?, ?, ?);
    """, differences)

    # 提交更改并关闭连接
    conn.commit()
    conn.close()


def parse_args():
    parse = argparse.ArgumentParser()
    parse.add_argument('-base',type=str)
    parse.add_argument('-update',type=str)
    parse.add_argument('-out', type=str)
    return  parse.parse_args()

parse_args = parse_args()

def main():
    # 数据库文件路径
    db1_path = parse_args.base
    db2_path = parse_args.update
    output_db_path = parse_args.out
 
    # db1_path = "/home/wangweiran/桌面/新建文件夹/2303-u2-990升级前.db"
    # db2_path = "/home/wangweiran/桌面/新建文件夹/2303-u2-990升级后.db"

    # output_db_path = "./2303-u2-990.db"

    # 读取两个数据库的 symbols 表数据
    db1_data = read_symbols_from_db(db1_path)
    db2_data = read_symbols_from_db(db2_path)

    # 对比两个数据库中的 symbols
    differences = compare_symbols(db1_data, db2_data)

    # 将差异存入指定的数据库
    save_differences_to_db(output_db_path, differences)
    print(f"对比结果已保存到数据库：{output_db_path}")

if __name__ == "__main__":
    main()