Files
Cloud-Tour-to-Libo/app/data_platform/csv_service.py
T

676 lines
24 KiB
Python

"""CSV preview, partial-safe batch import, and export for project data tables."""
from __future__ import annotations
import csv
import io
import json
import re
import uuid
from datetime import date, datetime, timezone
from decimal import Decimal
from typing import Any
from fastapi import HTTPException
from psycopg import sql
from psycopg.types.json import Jsonb
from app.data_platform.record_service import (
_database_or_404,
_table_or_404,
_validated_payload,
)
from app.data_platform.registry import FieldDefinition, TableDefinition
from app.db import get_conn
MAX_CSV_BYTES = 50 * 1024 * 1024
MAX_IMPORT_ROWS = 50_000
MAX_EXPORT_ROWS = 200_000
PREVIEW_ROWS = 5
MAX_REPORTED_ERRORS = 100
_SYSTEM_COLUMN_ALIASES = {
"id",
"记录 id",
"记录id",
"tenant_id",
"租户 id",
"租户id",
"project_id",
"项目 id",
"项目id",
"created_at",
"创建时间",
"updated_at",
"更新时间",
"deleted_at",
"deleted_by",
}
def _decode_csv(content: bytes) -> tuple[str, str]:
if not content:
raise HTTPException(400, "CSV 文件为空")
if len(content) > MAX_CSV_BYTES:
raise HTTPException(413, "CSV 文件不能超过 50 MB")
for encoding in ("utf-8-sig", "utf-8", "gb18030"):
try:
return content.decode(encoding), encoding
except UnicodeDecodeError:
continue
raise HTTPException(400, "CSV 编码无法识别,请使用 UTF-8 或 GB18030")
def _normalize_header(value: str | None) -> str:
return str(value or "").replace("\ufeff", "").strip()
def _csv_reader(text: str) -> csv.DictReader:
sample = text[:8192]
try:
dialect = csv.Sniffer().sniff(sample, delimiters=",;\t|")
except csv.Error:
dialect = csv.excel
return csv.DictReader(io.StringIO(text, newline=""), dialect=dialect)
def _safe_value(value: Any) -> Any:
if isinstance(value, Jsonb):
return _safe_value(value.obj)
if isinstance(value, Decimal):
return float(value)
if isinstance(value, uuid.UUID):
return str(value)
if isinstance(value, (date, datetime)):
return value.isoformat()
if isinstance(value, dict):
return {key: _safe_value(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_safe_value(item) for item in value]
return value
def _csv_value(value: Any) -> str:
value = _safe_value(value)
if value is None:
return ""
if isinstance(value, bool):
return "true" if value else "false"
if isinstance(value, (dict, list)):
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
return str(value)
def _signature_value(value: Any) -> tuple[str, Any]:
"""Return a stable, hashable value used for conservative exact-row deduplication."""
value = _safe_value(value)
if value is None:
return ("null", None)
if isinstance(value, bool):
return ("boolean", value)
if isinstance(value, (int, float, Decimal)):
try:
return ("number", format(Decimal(str(value)).normalize(), "f"))
except Exception:
return ("number", str(value))
if isinstance(value, str):
return ("text", value.strip().casefold())
if isinstance(value, (dict, list)):
return (
"json",
json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")),
)
return (type(value).__name__, str(value))
def _payload_signature(
payload: dict[str, Any],
field_codes: list[str],
) -> tuple[tuple[str, tuple[str, Any]], ...]:
return tuple(
(field_code, _signature_value(payload.get(field_code)))
for field_code in field_codes
)
def _strict_csv_value(field: FieldDefinition, raw_value: str) -> Any:
value = raw_value.strip()
if not value:
return None
if field.data_type == "boolean":
lowered = value.lower()
if lowered in {"true", "1", "yes", "y", "是", "启用"}:
return True
if lowered in {"false", "0", "no", "n", "否", "禁用"}:
return False
raise HTTPException(400, f"{field.label}必须是 true/false、1/0 或 是/否")
if field.data_type == "integer" or "INTEGER" in field.sql_type.upper():
try:
decimal_value = Decimal(value)
except Exception as exc:
raise HTTPException(400, f"{field.label}必须是整数") from exc
if decimal_value != decimal_value.to_integral_value():
raise HTTPException(400, f"{field.label}必须是整数")
return int(decimal_value)
if field.data_type == "number":
try:
return Decimal(value)
except Exception as exc:
raise HTTPException(400, f"{field.label}必须是数字") from exc
if field.data_type == "json":
try:
return json.loads(value)
except json.JSONDecodeError as exc:
raise HTTPException(400, f"{field.label}不是合法 JSON") from exc
if field.data_type == "date":
try:
date.fromisoformat(value)
except ValueError as exc:
raise HTTPException(400, f"{field.label}必须使用 YYYY-MM-DD 格式") from exc
if field.data_type == "datetime":
try:
datetime.fromisoformat(value.replace("Z", "+00:00"))
except ValueError as exc:
raise HTTPException(400, f"{field.label}不是合法日期时间") from exc
if field.sql_type.upper().startswith("UUID"):
try:
uuid.UUID(value)
except ValueError as exc:
raise HTTPException(400, f"{field.label}不是合法 UUID") from exc
varchar_match = re.match(r"^(?:VAR)?CHAR(?:ACTER VARYING)?\((\d+)\)", field.sql_type.upper())
if varchar_match and len(value) > int(varchar_match.group(1)):
raise HTTPException(
400,
f"{field.label}不能超过 {varchar_match.group(1)} 个字符",
)
if field.options and value not in field.options:
raise HTTPException(400, f"{field.label}可选值为:{', '.join(field.options)}")
return value
def _mapping_for_headers(
table: TableDefinition,
headers: list[str],
) -> tuple[list[dict[str, Any]], dict[str, FieldDefinition], list[str]]:
code_lookup = {field.code.lower(): field for field in table.fields}
label_lookup = {field.label: field for field in table.fields}
mappings: list[dict[str, Any]] = []
mapped_headers: dict[str, FieldDefinition] = {}
schema_errors: list[str] = []
mapped_codes: set[str] = set()
if not headers or not any(headers):
return [], {}, ["CSV 缺少表头"]
duplicate_headers = sorted({header for header in headers if headers.count(header) > 1})
if duplicate_headers:
schema_errors.append(f"CSV 存在重复列:{', '.join(duplicate_headers)}")
for header in headers:
normalized = header.strip()
field = code_lookup.get(normalized.lower()) or label_lookup.get(normalized)
if field:
if field.code in mapped_codes:
mappings.append(
{
"source_header": header,
"field_code": field.code,
"field_label": field.label,
"status": "duplicate",
}
)
schema_errors.append(f"字段“{field.label}”在 CSV 中重复映射")
continue
mapped_codes.add(field.code)
mapped_headers[header] = field
mappings.append(
{
"source_header": header,
"field_code": field.code,
"field_label": field.label,
"status": "mapped",
}
)
continue
if normalized.lower() in _SYSTEM_COLUMN_ALIASES:
mappings.append(
{
"source_header": header,
"field_code": None,
"field_label": "平台自动生成",
"status": "ignored",
}
)
continue
mappings.append(
{
"source_header": header,
"field_code": None,
"field_label": None,
"status": "unknown",
}
)
schema_errors.append(f"无法识别 CSV 列“{header}”")
missing_required = [
field.label
for field in table.fields
if field.required and field.code not in mapped_codes
]
if missing_required:
schema_errors.append(f"缺少必填列:{', '.join(missing_required)}")
if not mapped_headers:
schema_errors.append("CSV 没有可导入的业务字段")
return mappings, mapped_headers, schema_errors
def _parse_csv(
table: TableDefinition,
content: bytes,
file_name: str,
) -> tuple[dict[str, Any], list[dict[str, Any]], list[str]]:
text, encoding = _decode_csv(content)
reader = _csv_reader(text)
headers = [_normalize_header(header) for header in (reader.fieldnames or [])]
mappings, mapped_headers, schema_errors = _mapping_for_headers(table, headers)
rows: list[dict[str, Any]] = []
errors: list[dict[str, Any]] = []
error_count = 0
total_rows = 0
for physical_row, source_row in enumerate(reader, start=2):
normalized_source = {
_normalize_header(key): value
for key, value in source_row.items()
if key is not None
}
if None in source_row and source_row[None]:
error_count += 1
if len(errors) < MAX_REPORTED_ERRORS:
errors.append(
{
"row": physical_row,
"field": "",
"kind": "invalid",
"message": "该行列数多于表头,请检查分隔符或引号",
}
)
continue
if not any(str(value or "").strip() for value in normalized_source.values()):
continue
total_rows += 1
if total_rows > MAX_IMPORT_ROWS:
raise HTTPException(400, f"单次最多导入 {MAX_IMPORT_ROWS:,} 条记录")
raw_payload: dict[str, Any] = {}
row_error: str | None = None
row_field = ""
for header, field in mapped_headers.items():
raw_value = str(normalized_source.get(header) or "")
try:
coerced = _strict_csv_value(field, raw_value)
except HTTPException as exc:
row_error = str(exc.detail)
row_field = field.label
break
if coerced is not None:
raw_payload[field.code] = coerced
if row_error is None:
try:
payload = _validated_payload(table, raw_payload, create=True)
except HTTPException as exc:
row_error = str(exc.detail)
payload = {}
else:
payload = {}
if row_error:
error_count += 1
if len(errors) < MAX_REPORTED_ERRORS:
errors.append(
{
"row": physical_row,
"field": row_field,
"kind": "invalid",
"message": row_error,
}
)
continue
if not payload:
error_count += 1
if len(errors) < MAX_REPORTED_ERRORS:
errors.append(
{
"row": physical_row,
"field": "",
"kind": "invalid",
"message": "该行没有可导入的业务字段",
}
)
continue
rows.append({"row": physical_row, "payload": payload})
if total_rows == 0:
schema_errors.append("CSV 没有可导入的数据行")
dedupe_fields = [
str(item["field_code"])
for item in mappings
if item["status"] == "mapped" and item.get("field_code")
]
preview = {
"file_name": file_name,
"encoding": encoding,
"total_rows": total_rows,
"valid_rows": len(rows),
"invalid_rows": error_count,
"duplicate_rows": 0,
"file_duplicate_rows": 0,
"existing_duplicate_rows": 0,
"skipped_rows": error_count,
"columns": mappings,
"schema_errors": schema_errors,
"errors": errors,
"errors_truncated": error_count > len(errors),
"preview_rows": [],
"can_import": not schema_errors and bool(rows),
}
return preview, rows, dedupe_fields
async def _prepare_csv_import(
project_id: str,
table: TableDefinition,
database: dict[str, Any],
content: bytes,
file_name: str,
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
preview, candidates, dedupe_fields = _parse_csv(table, content, file_name)
if preview["schema_errors"] or not candidates or not dedupe_fields:
preview["can_import"] = False
return preview, []
schema_name = str(database["schema_name"])
existing_signatures: set[tuple[tuple[str, tuple[str, Any]], ...]] = set()
async with get_conn() as conn:
async with conn.cursor() as cur:
await cur.execute(
sql.SQL("SELECT {} FROM {}.{} WHERE deleted_at IS NULL").format(
sql.SQL(", ").join(sql.Identifier(code) for code in dedupe_fields),
sql.Identifier(schema_name),
sql.Identifier(table.code),
)
)
for existing_row in await cur.fetchall():
existing_signatures.add(
_payload_signature(dict(existing_row), dedupe_fields)
)
valid_candidates: list[dict[str, Any]] = []
seen_file_signatures: dict[
tuple[tuple[str, tuple[str, Any]], ...],
int,
] = {}
duplicate_errors: list[dict[str, Any]] = []
file_duplicate_rows = 0
existing_duplicate_rows = 0
for candidate in candidates:
signature = _payload_signature(candidate["payload"], dedupe_fields)
physical_row = int(candidate["row"])
if signature in existing_signatures:
existing_duplicate_rows += 1
duplicate_errors.append(
{
"row": physical_row,
"field": "",
"kind": "duplicate",
"message": "与当前数据表中的正式数据重复,已跳过",
}
)
continue
if signature in seen_file_signatures:
file_duplicate_rows += 1
duplicate_errors.append(
{
"row": physical_row,
"field": "",
"kind": "duplicate",
"message": (
f"与 CSV 第 {seen_file_signatures[signature]} 行重复,已跳过"
),
}
)
continue
seen_file_signatures[signature] = physical_row
valid_candidates.append(candidate)
duplicate_rows = file_duplicate_rows + existing_duplicate_rows
remaining_error_slots = max(0, MAX_REPORTED_ERRORS - len(preview["errors"]))
preview["errors"].extend(duplicate_errors[:remaining_error_slots])
preview["errors"].sort(key=lambda item: int(item["row"]))
preview.update(
{
"valid_rows": len(valid_candidates),
"duplicate_rows": duplicate_rows,
"file_duplicate_rows": file_duplicate_rows,
"existing_duplicate_rows": existing_duplicate_rows,
"skipped_rows": int(preview["invalid_rows"]) + duplicate_rows,
"errors_truncated": (
int(preview["invalid_rows"]) + duplicate_rows
> len(preview["errors"])
),
"preview_rows": [
{
key: _safe_value(value)
for key, value in candidate["payload"].items()
}
for candidate in valid_candidates[:PREVIEW_ROWS]
],
"can_import": bool(valid_candidates),
}
)
return preview, valid_candidates
async def preview_csv_import(
project_id: str,
table_code: str,
content: bytes,
file_name: str,
) -> dict[str, Any]:
table = await _table_or_404(project_id, table_code)
if not table.allow_create:
raise HTTPException(403, "该表不允许导入新增记录")
database = await _database_or_404(project_id)
preview, _rows = await _prepare_csv_import(
project_id,
table,
database,
content,
file_name,
)
return preview
async def import_csv_records(
project_id: str,
table_code: str,
content: bytes,
file_name: str,
actor: str,
) -> dict[str, Any]:
table = await _table_or_404(project_id, table_code)
if not table.allow_create:
raise HTTPException(403, "该表不允许导入新增记录")
database = await _database_or_404(project_id)
preview, rows = await _prepare_csv_import(
project_id,
table,
database,
content,
file_name,
)
if not preview["can_import"]:
first_schema_error = (preview["schema_errors"] or [None])[0]
first_row_error = (preview["errors"] or [None])[0]
if first_schema_error:
raise HTTPException(400, f"CSV 校验失败:{first_schema_error}")
if first_row_error and not preview["duplicate_rows"]:
raise HTTPException(
400,
f"CSV 校验失败:第 {first_row_error['row']} 行 {first_row_error['message']}",
)
if preview["duplicate_rows"]:
raise HTTPException(400, "没有可导入的新数据,文件中的记录均已存在")
raise HTTPException(400, "CSV 没有可导入的有效数据")
schema_name = str(database["schema_name"])
tenant_id = str(database["tenant_id"])
imported_count = 0
failed_errors: list[dict[str, Any]] = []
insert_queries: dict[tuple[str, ...], sql.Composed] = {}
async with get_conn() as conn:
async with conn.cursor() as cur:
audit_query = sql.SQL(
"""
INSERT INTO {}.data_change_logs (
tenant_id, project_id, table_code, record_id, operation,
before_data, after_data, actor
) VALUES (%s, %s, %s, %s, 'import', NULL, %s, %s)
"""
).format(sql.Identifier(schema_name))
for candidate in rows:
payload = candidate["payload"]
payload_keys = tuple(payload.keys())
if payload_keys not in insert_queries:
columns = ["id", "tenant_id", "project_id", *payload_keys]
insert_queries[payload_keys] = sql.SQL(
"INSERT INTO {}.{} ({}) VALUES ({})"
).format(
sql.Identifier(schema_name),
sql.Identifier(table.code),
sql.SQL(", ").join(
sql.Identifier(column) for column in columns
),
sql.SQL(", ").join(sql.Placeholder() for _ in columns),
)
record_id = uuid.uuid4()
try:
async with conn.transaction():
await cur.execute(
insert_queries[payload_keys],
(
record_id,
tenant_id,
project_id,
*(payload[key] for key in payload_keys),
),
)
after_data = {
"id": str(record_id),
**{
key: _safe_value(value)
for key, value in payload.items()
},
}
await cur.execute(
audit_query,
(
tenant_id,
project_id,
table_code,
record_id,
Jsonb(after_data),
actor,
),
)
imported_count += 1
except Exception:
if len(failed_errors) < MAX_REPORTED_ERRORS:
failed_errors.append(
{
"row": int(candidate["row"]),
"field": "",
"kind": "failed",
"message": "数据库约束校验失败,已跳过",
}
)
await conn.commit()
failed_count = len(failed_errors)
skipped_count = int(preview["skipped_rows"]) + failed_count
return {
"ok": True,
"file_name": file_name,
"imported_count": imported_count,
"invalid_count": int(preview["invalid_rows"]),
"duplicate_count": int(preview["duplicate_rows"]),
"failed_count": failed_count,
"skipped_count": skipped_count,
"errors": failed_errors,
"table_code": table.code,
}
async def export_csv_records(
project_id: str,
table_code: str,
search: str | None = None,
) -> tuple[bytes, str, int]:
table = await _table_or_404(project_id, table_code)
database = await _database_or_404(project_id)
schema_name = str(database["schema_name"])
columns = ["id", *(field.code for field in table.fields), "created_at", "updated_at"]
headers = ["记录 ID", *(field.label for field in table.fields), "创建时间", "更新时间"]
where_parts: list[sql.Composable] = [sql.SQL("deleted_at IS NULL")]
params: list[Any] = []
searchable = [field for field in table.fields if field.searchable]
if search and searchable:
pattern = f"%{search.strip()}%"
where_parts.append(
sql.SQL("(")
+ sql.SQL(" OR ").join(
sql.SQL("{}::text ILIKE %s").format(sql.Identifier(field.code))
for field in searchable
)
+ sql.SQL(")")
)
params.extend(pattern for _ in searchable)
where_clause = sql.SQL(" AND ").join(where_parts)
async with get_conn() as conn:
async with conn.cursor() as cur:
await cur.execute(
sql.SQL("SELECT count(*) AS count FROM {}.{} WHERE {}").format(
sql.Identifier(schema_name),
sql.Identifier(table.code),
where_clause,
),
params,
)
total = int((await cur.fetchone())["count"])
if total > MAX_EXPORT_ROWS:
raise HTTPException(
400,
f"当前数据量为 {total:,} 条,单次最多导出 {MAX_EXPORT_ROWS:,} 条,请先搜索筛选",
)
await cur.execute(
sql.SQL("SELECT {} FROM {}.{} WHERE {} ORDER BY created_at ASC").format(
sql.SQL(", ").join(sql.Identifier(column) for column in columns),
sql.Identifier(schema_name),
sql.Identifier(table.code),
where_clause,
),
params,
)
rows = await cur.fetchall()
output = io.StringIO(newline="")
writer = csv.writer(output, lineterminator="\n")
writer.writerow(headers)
for row in rows:
writer.writerow([_csv_value(row[column]) for column in columns])
content = ("\ufeff" + output.getvalue()).encode("utf-8")
timestamp = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S")
filename = f"{table.code}_{timestamp}.csv"
return content, filename, total