fix: 修复扇区推断及按CGI更新导入
This commit is contained in:
+106
-14
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user