From b6cf0e61e72d65b3305bf0408f52488e1a49e22a Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 26 Aug 2026 17:08:19 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E6=89=87=E5=8C=BA?= =?UTF-8?q?=E6=8E=A8=E6=96=AD=E5=8F=8A=E6=8C=89CGI=E6=9B=B4=E6=96=B0?= =?UTF-8?q?=E5=AF=BC=E5=85=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- CellData.sql | 4 +- app/api/routers/database.py | 11 ++ app/database.py | 120 ++++++++++++-- db_init/celldata.sector.sql | 2 +- tests/test_database_import.py | 33 ++++ tests/test_database_import_mysql.py | 232 ++++++++++++++++++++++++++++ 6 files changed, 385 insertions(+), 17 deletions(-) create mode 100644 tests/test_database_import.py create mode 100644 tests/test_database_import_mysql.py diff --git a/CellData.sql b/CellData.sql index 2066034..125f0ce 100644 --- a/CellData.sql +++ b/CellData.sql @@ -43,10 +43,10 @@ UPDATE _sector_infer SET -- 2.2 制式/频段:按特征库匹配;未命中回落 cellinfo 原制式 UPDATE _sector_infer t JOIN sector_band_ref r - ON r.`网络` = t.`网络` + ON r.`网络` COLLATE utf8mb4_unicode_ci = t.`网络` COLLATE utf8mb4_unicode_ci AND CAST(NULLIF(t.`频点`,'') AS DECIMAL(10,2)) >= r.`频点下限` AND CAST(NULLIF(t.`频点`,'') AS DECIMAL(10,2)) < r.`频点上限` - AND (r.`PLMN` IS NULL OR r.`PLMN` = t.PLMN) + AND (r.`PLMN` IS NULL OR r.`PLMN` COLLATE utf8mb4_unicode_ci = t.PLMN COLLATE utf8mb4_unicode_ci) SET t.`制式` = r.`制式`, t.`频段` = r.`频段`; UPDATE _sector_infer SET `制式` = ci_zs, `频段` = ci_zs WHERE `制式` IS NULL; diff --git a/app/api/routers/database.py b/app/api/routers/database.py index e8c6a84..f8f0027 100644 --- a/app/api/routers/database.py +++ b/app/api/routers/database.py @@ -297,6 +297,17 @@ def import_table_csv( if extra: parts.append("多余字段: " + ", ".join(extra)) raise HTTPException(status_code=400, detail="CSV 字段与模板不一致,导入失败。" + ";".join(parts)) + if database_source == "cell_data" and table_name.lower() == "sector": + stats = db.upsert_csv(str(tmp_path), table_name, "CGI", encoding=encoding) + return { + "success": True, + "message": ( + f"导入成功,共 {stats['imported_rows']} 行:" + f"新增 {stats['inserted_rows']} 行,更新 {stats['updated_rows']} 行," + f"清理历史重复 {stats['removed_duplicate_rows']} 行" + ), + **stats, + } imported = db.import_csv(str(tmp_path), table_name, encoding=encoding) except HTTPException: raise diff --git a/app/database.py b/app/database.py index 1205dc0..511031e 100644 --- a/app/database.py +++ b/app/database.py @@ -22,6 +22,44 @@ def detect_csv_encoding(file_path: str) -> str: raise UnicodeError("CSV 编码无法识别,仅支持 UTF-8 或 GBK 编码") +def _read_csv_rows(file_path: str, encoding: str) -> tuple[List[str], List[Tuple]]: + with open(file_path, "r", encoding=encoding, newline="") as handle: + reader = csv.reader(handle) + try: + header = next(reader) + except StopIteration: + return [], [] + columns = [str(name).strip() for name in header] + width = len(columns) + data: List[Tuple] = [] + for row in reader: + if not any(str(cell).strip() for cell in row): + continue + cells = list(row[:width]) + [""] * (width - len(row)) + data.append(tuple(cells)) + return columns, data + + +def _deduplicate_rows_by_key( + columns: List[str], + data: List[Tuple], + key_column: str, +) -> tuple[List[Tuple], int]: + if key_column not in columns: + raise ValueError(f"CSV 缺少业务键字段: {key_column}") + + key_index = columns.index(key_column) + rows_by_key: dict[str, Tuple] = {} + for row_number, row in enumerate(data, start=2): + key = str(row[key_index] or "").strip() + if not key: + raise ValueError(f"CSV 第 {row_number} 行 {key_column} 为空") + normalized = list(row) + normalized[key_index] = key + rows_by_key[key] = tuple(normalized) + return list(rows_by_key.values()), len(data) - len(rows_by_key) + + class DatabaseManager: """数据库管理器""" @@ -252,24 +290,78 @@ class DatabaseManager: def import_csv(self, file_path: str, table_name: str, encoding: str | None = None) -> int: """按 CSV 表头列追加导入(列须与表字段一致,由调用方校验)。返回导入行数。""" csv_encoding = encoding or detect_csv_encoding(file_path) - with open(file_path, "r", encoding=csv_encoding, newline="") as handle: - reader = csv.reader(handle) - try: - header = next(reader) - except StopIteration: - return 0 - columns = [str(name).strip() for name in header] - width = len(columns) - data: List[Tuple] = [] - for row in reader: - if not any(str(cell).strip() for cell in row): - continue - cells = list(row[:width]) + [""] * (width - len(row)) - data.append(tuple(cells)) + columns, data = _read_csv_rows(file_path, csv_encoding) if not data: return 0 return self.bulk_insert(table_name, columns, data) + def upsert_csv( + self, + file_path: str, + table_name: str, + key_column: str, + encoding: str | None = None, + batch_size: int = 1000, + ) -> Dict[str, int]: + """按业务键覆盖导入;上传行获胜,并清理该键已有的重复行。""" + csv_encoding = encoding or detect_csv_encoding(file_path) + columns, raw_data = _read_csv_rows(file_path, csv_encoding) + data, input_duplicate_rows = _deduplicate_rows_by_key(columns, raw_data, key_column) + if not data: + return { + "imported_rows": 0, + "inserted_rows": 0, + "updated_rows": 0, + "removed_duplicate_rows": 0, + "input_duplicate_rows": input_duplicate_rows, + } + + key_index = columns.index(key_column) + keys = [str(row[key_index]) for row in data] + placeholders = ", ".join(["%s"] * len(columns)) + column_names = ", ".join(f"`{column}`" for column in columns) + insert_sql = f"INSERT INTO `{table_name}` ({column_names}) VALUES ({placeholders})" + existing_keys: set[str] = set() + removed_rows = 0 + + with self.get_fast_connection() as connection: + try: + with connection.cursor() as cursor: + for start in range(0, len(keys), batch_size): + batch = keys[start:start + batch_size] + marks = ", ".join(["%s"] * len(batch)) + cursor.execute( + f"SELECT DISTINCT `{key_column}` FROM `{table_name}` " + f"WHERE `{key_column}` IN ({marks})", + batch, + ) + existing_keys.update(str(row[0]) for row in cursor.fetchall()) + + for start in range(0, len(keys), batch_size): + batch = keys[start:start + batch_size] + marks = ", ".join(["%s"] * len(batch)) + cursor.execute( + f"DELETE FROM `{table_name}` WHERE `{key_column}` IN ({marks})", + batch, + ) + removed_rows += max(cursor.rowcount, 0) + + for start in range(0, len(data), batch_size): + cursor.executemany(insert_sql, data[start:start + batch_size]) + connection.commit() + except Exception: + connection.rollback() + raise + + updated_rows = len(existing_keys) + return { + "imported_rows": len(data), + "inserted_rows": len(data) - updated_rows, + "updated_rows": updated_rows, + "removed_duplicate_rows": max(removed_rows - updated_rows, 0), + "input_duplicate_rows": input_duplicate_rows, + } + def truncate_table(self, table_name: str) -> bool: """清空表""" with self.get_connection() as conn: diff --git a/db_init/celldata.sector.sql b/db_init/celldata.sector.sql index a32baec..08cec1e 100644 --- a/db_init/celldata.sector.sql +++ b/db_init/celldata.sector.sql @@ -10,5 +10,5 @@ CREATE TABLE IF NOT EXISTS `sector` ( `带宽` varchar(20) DEFAULT NULL, `站型` varchar(50) DEFAULT NULL, `网络` varchar(20) DEFAULT NULL, - KEY `CGI` (`CGI`) + UNIQUE KEY `uq_sector_cgi` (`CGI`) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci; diff --git a/tests/test_database_import.py b/tests/test_database_import.py new file mode 100644 index 0000000..0e61dda --- /dev/null +++ b/tests/test_database_import.py @@ -0,0 +1,33 @@ +import unittest + +from app.database import _deduplicate_rows_by_key + + +class SectorImportTests(unittest.TestCase): + def test_duplicate_cgi_uses_last_uploaded_row(self): + columns = ["CGI", "扇区", "物理站"] + rows = [ + ("460-00-1-1", "旧扇区", "旧物理站"), + ("460-00-2-1", "新增扇区", "新增物理站"), + ("460-00-1-1", "新扇区", "新物理站"), + ] + + result, duplicate_count = _deduplicate_rows_by_key(columns, rows, "CGI") + + self.assertEqual(duplicate_count, 1) + self.assertEqual(result, [ + ("460-00-1-1", "新扇区", "新物理站"), + ("460-00-2-1", "新增扇区", "新增物理站"), + ]) + + def test_blank_cgi_is_rejected(self): + with self.assertRaisesRegex(ValueError, "第 2 行 CGI 为空"): + _deduplicate_rows_by_key(["CGI", "扇区"], [(" ", "扇区1")], "CGI") + + def test_missing_cgi_column_is_rejected(self): + with self.assertRaisesRegex(ValueError, "缺少业务键字段: CGI"): + _deduplicate_rows_by_key(["扇区"], [("扇区1",)], "CGI") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_database_import_mysql.py b/tests/test_database_import_mysql.py new file mode 100644 index 0000000..ad52ea3 --- /dev/null +++ b/tests/test_database_import_mysql.py @@ -0,0 +1,232 @@ +import csv +import os +import tempfile +import unittest +from pathlib import Path + +import pymysql +from pymysql.constants import CLIENT + +from app.config import AppConfig, MySQLConfig +from app.database import DatabaseManager + + +MYSQL_HOST = os.environ.get("CAPACITYREPORT_TEST_MYSQL_HOST") +MYSQL_PORT = int(os.environ.get("CAPACITYREPORT_TEST_MYSQL_PORT", "3306")) +MYSQL_USER = os.environ.get("CAPACITYREPORT_TEST_MYSQL_USER", "root") +MYSQL_PASSWORD = os.environ.get("CAPACITYREPORT_TEST_MYSQL_PASSWORD", "") +TEST_DATABASE = "capacityreport_sector_import_test" + + +@unittest.skipUnless(MYSQL_HOST, "未配置 CapacityReport 测试 MySQL") +class SectorImportMySQLTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.admin_connection = pymysql.connect( + host=MYSQL_HOST, + port=MYSQL_PORT, + user=MYSQL_USER, + password=MYSQL_PASSWORD, + charset="utf8mb4", + autocommit=True, + ) + with cls.admin_connection.cursor() as cursor: + cursor.execute(f"DROP DATABASE IF EXISTS `{TEST_DATABASE}`") + cursor.execute( + f"CREATE DATABASE `{TEST_DATABASE}` " + "CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci" + ) + config = MySQLConfig( + host=MYSQL_HOST, + port=MYSQL_PORT, + user=MYSQL_USER, + passwd=MYSQL_PASSWORD, + dbname=TEST_DATABASE, + ) + cls.manager = DatabaseManager(AppConfig(mysql=config), config) + + @classmethod + def tearDownClass(cls): + with cls.admin_connection.cursor() as cursor: + cursor.execute(f"DROP DATABASE IF EXISTS `{TEST_DATABASE}`") + cls.admin_connection.close() + + def setUp(self): + self.connection = pymysql.connect( + host=MYSQL_HOST, + port=MYSQL_PORT, + user=MYSQL_USER, + password=MYSQL_PASSWORD, + database=TEST_DATABASE, + charset="utf8mb4", + autocommit=True, + ) + + def tearDown(self): + self.connection.close() + + def _write_csv(self, columns, rows): + handle = tempfile.NamedTemporaryFile( + mode="w", + encoding="utf-8-sig", + newline="", + suffix=".csv", + delete=False, + ) + try: + writer = csv.writer(handle) + writer.writerow(columns) + writer.writerows(rows) + return Path(handle.name) + finally: + handle.close() + + def test_upsert_updates_inserts_and_removes_historical_duplicates(self): + with self.connection.cursor() as cursor: + cursor.execute("DROP TABLE IF EXISTS `sector_upsert_test`") + cursor.execute( + "CREATE TABLE `sector_upsert_test` (" + "`CGI` varchar(120), `扇区` varchar(200), `物理站` varchar(200)) " + "CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci" + ) + cursor.executemany( + "INSERT INTO `sector_upsert_test` VALUES (%s, %s, %s)", + [ + ("460-00-1-1", "旧扇区1", "旧站1"), + ("460-00-1-1", "旧扇区2", "旧站2"), + ("460-00-9-9", "保留扇区", "保留站"), + ], + ) + + csv_path = self._write_csv( + ["CGI", "扇区", "物理站"], + [ + ("460-00-1-1", "更新扇区", "更新站"), + ("460-00-2-2", "新增扇区", "新增站"), + ], + ) + try: + stats = self.manager.upsert_csv( + str(csv_path), "sector_upsert_test", "CGI" + ) + finally: + csv_path.unlink(missing_ok=True) + + self.assertEqual( + stats, + { + "imported_rows": 2, + "inserted_rows": 1, + "updated_rows": 1, + "removed_duplicate_rows": 1, + "input_duplicate_rows": 0, + }, + ) + with self.connection.cursor() as cursor: + cursor.execute( + "SELECT `CGI`, `扇区`, `物理站` FROM `sector_upsert_test` ORDER BY `CGI`" + ) + self.assertEqual( + cursor.fetchall(), + ( + ("460-00-1-1", "更新扇区", "更新站"), + ("460-00-2-2", "新增扇区", "新增站"), + ("460-00-9-9", "保留扇区", "保留站"), + ), + ) + + def test_upsert_rolls_back_deletion_when_insert_fails(self): + with self.connection.cursor() as cursor: + cursor.execute("DROP TABLE IF EXISTS `sector_rollback_test`") + cursor.execute( + "CREATE TABLE `sector_rollback_test` (" + "`CGI` varchar(120), `扇区编号` int NOT NULL)" + ) + cursor.execute( + "INSERT INTO `sector_rollback_test` VALUES (%s, %s)", + ("460-00-1-1", 7), + ) + + csv_path = self._write_csv( + ["CGI", "扇区编号"], [("460-00-1-1", "not-an-integer")] + ) + try: + with self.assertRaises((pymysql.DataError, pymysql.IntegrityError)): + self.manager.upsert_csv( + str(csv_path), "sector_rollback_test", "CGI" + ) + finally: + csv_path.unlink(missing_ok=True) + + with self.connection.cursor() as cursor: + cursor.execute("SELECT `CGI`, `扇区编号` FROM `sector_rollback_test`") + self.assertEqual(cursor.fetchall(), (("460-00-1-1", 7),)) + + def test_celldata_script_handles_mixed_collations(self): + with self.connection.cursor() as cursor: + cursor.execute("DROP TABLE IF EXISTS `cellinfo`") + cursor.execute("DROP TABLE IF EXISTS `sector`") + cursor.execute("DROP TABLE IF EXISTS `sector_band_ref`") + cursor.execute( + "CREATE TABLE `cellinfo` (" + "`CGI` varchar(120), `PLMN` varchar(100), `小区名称` varchar(200), " + "`频点` varchar(50), `带宽` varchar(20), `制式` varchar(50), " + "`网络` varchar(20)) CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci" + ) + cursor.execute( + "CREATE TABLE `sector` (" + "`CGI` varchar(120), `扇区` varchar(200), `物理站` varchar(200), " + "`制式` varchar(50), `频段` varchar(50), `带宽` varchar(20), " + "`站型` varchar(50), `网络` varchar(20), UNIQUE KEY (`CGI`)) " + "CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci" + ) + cursor.execute( + "CREATE TABLE `sector_band_ref` (" + "`网络` varchar(8), `频点下限` decimal(10,2), `频点上限` decimal(10,2), " + "`PLMN` varchar(16), `制式` varchar(20), `频段` varchar(20)) " + "CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci" + ) + cursor.execute( + "INSERT INTO `cellinfo` VALUES " + "('460-00-122116-32','460-00','江门恩平敬老院F-ZLH-102'," + "'1909.4','20','TDD','4G')" + ) + cursor.execute( + "INSERT INTO `sector_band_ref` VALUES " + "('4G',1880,1920,NULL,'TDD','F频')" + ) + + script = (Path(__file__).parents[1] / "CellData.sql").read_text( + encoding="utf-8-sig" + ) + script_connection = pymysql.connect( + host=MYSQL_HOST, + port=MYSQL_PORT, + user=MYSQL_USER, + password=MYSQL_PASSWORD, + database=TEST_DATABASE, + charset="utf8mb4", + autocommit=True, + client_flag=CLIENT.MULTI_STATEMENTS, + ) + try: + with script_connection.cursor() as cursor: + cursor.execute(script) + while cursor.nextset(): + pass + finally: + script_connection.close() + + with self.connection.cursor() as cursor: + cursor.execute( + "SELECT `扇区`, `物理站`, `制式`, `频段` FROM `sector` " + "WHERE `CGI` = '460-00-122116-32'" + ) + self.assertEqual( + cursor.fetchone(), + ("江门恩平敬老院2", "江门恩平敬老院", "TDD", "F频"), + ) + + +if __name__ == "__main__": + unittest.main()