diff --git a/app/api/routers/database.py b/app/api/routers/database.py index 0779230..e8c6a84 100644 --- a/app/api/routers/database.py +++ b/app/api/routers/database.py @@ -11,7 +11,7 @@ from starlette.background import BackgroundTask from app import state from app.config import CACHE_DIR -from app.database import DatabaseManager +from app.database import DatabaseManager, detect_csv_encoding from app.services.platform import make_client from app.utils.files import remove_file_safely from app.warehouse import make_cell_data_warehouse, make_warehouse @@ -255,8 +255,8 @@ def download_table_template( ) -def _read_csv_header(path: Path) -> list[str]: - with path.open("r", encoding="utf-8-sig", newline="") as handle: +def _read_csv_header(path: Path, encoding: str) -> list[str]: + with path.open("r", encoding=encoding, newline="") as handle: for row in csv.reader(handle): return [str(cell).strip() for cell in row] return [] @@ -284,7 +284,8 @@ def import_table_csv( tmp_path = CACHE_DIR / f"import_{timestamp}.csv" tmp_path.write_bytes(file.file.read()) try: - header = _read_csv_header(tmp_path) + encoding = detect_csv_encoding(str(tmp_path)) + header = _read_csv_header(tmp_path, encoding) if not header: raise HTTPException(status_code=400, detail="CSV 文件为空或缺少表头") missing = [name for name in columns if name not in header] @@ -296,7 +297,7 @@ def import_table_csv( if extra: parts.append("多余字段: " + ", ".join(extra)) raise HTTPException(status_code=400, detail="CSV 字段与模板不一致,导入失败。" + ";".join(parts)) - imported = db.import_csv(str(tmp_path), table_name) + imported = db.import_csv(str(tmp_path), table_name, encoding=encoding) except HTTPException: raise except Exception as exc: diff --git a/app/database.py b/app/database.py index 9ab9777..1205dc0 100644 --- a/app/database.py +++ b/app/database.py @@ -9,6 +9,19 @@ from typing import Any, Dict, List, Optional, Tuple from app.config import AppConfig, MySQLConfig +def detect_csv_encoding(file_path: str) -> str: + """识别数据管理导入 CSV 的 UTF-8 或 GBK 系列编码。""" + for encoding in ("utf-8-sig", "gb18030"): + try: + with open(file_path, "r", encoding=encoding, errors="strict") as handle: + while handle.read(1024 * 1024): + pass + return encoding + except UnicodeDecodeError: + continue + raise UnicodeError("CSV 编码无法识别,仅支持 UTF-8 或 GBK 编码") + + class DatabaseManager: """数据库管理器""" @@ -236,9 +249,10 @@ class DatabaseManager: conn.commit() return cursor.rowcount - def import_csv(self, file_path: str, table_name: str) -> int: + def import_csv(self, file_path: str, table_name: str, encoding: str | None = None) -> int: """按 CSV 表头列追加导入(列须与表字段一致,由调用方校验)。返回导入行数。""" - with open(file_path, "r", encoding="utf-8-sig", newline="") as handle: + 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) diff --git a/app/warehouse.py b/app/warehouse.py index 858583a..fa54ae8 100644 --- a/app/warehouse.py +++ b/app/warehouse.py @@ -182,16 +182,31 @@ class MetrixWarehouse: raise RuntimeError(str(result)) return int(result.get("affected_rows", 0)) if isinstance(result, dict) else 0 - def import_csv(self, file_path: str, table_name: str) -> int: + def import_csv(self, file_path: str, table_name: str, encoding: str | None = None) -> int: """走 Metrix 平台导入 API 追加导入到已存在的表(字段映射由平台按列名处理)。""" - job_id = self.client.import_csv( - self.conn_id, table_name, Path(file_path), mode="append", - database=self.database, create_table=False, - ) - job = self.client.wait_job(job_id) - if job.get("status") != "success": - raise RuntimeError(job.get("error_code") or job.get("status") or "导入失败") - return int(job.get("row_count") or 0) + source_path = Path(file_path) + upload_path = source_path + converted_path: Path | None = None + if encoding and encoding.lower() not in {"utf-8", "utf-8-sig"}: + converted_path = source_path.with_name(f"{source_path.stem}_utf8{source_path.suffix}") + with source_path.open("r", encoding=encoding, newline="") as source: + with converted_path.open("w", encoding="utf-8-sig", newline="") as target: + while chunk := source.read(1024 * 1024): + target.write(chunk) + upload_path = converted_path + + try: + job_id = self.client.import_csv( + self.conn_id, table_name, upload_path, mode="append", + database=self.database, create_table=False, + ) + job = self.client.wait_job(job_id) + if job.get("status") != "success": + raise RuntimeError(job.get("error_code") or job.get("status") or "导入失败") + return int(job.get("row_count") or 0) + finally: + if converted_path is not None: + converted_path.unlink(missing_ok=True) def execute_sql(self, sql: str) -> Tuple[bool, Any]: try: