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