"""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