Files
CapacityReport/tests/test_database_import_mysql.py
T

233 lines
8.3 KiB
Python

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