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()