Commit 2dda2414 authored by Data Governance Dev's avatar Data Governance Dev

feat(web2): 身份证扫描 — 行级违规原因 + 字段元数据透传 + 轮询端点

后端:
- idcard_validator.py
  - 新增 EMPTY 校验项(值空/None 视为不合规)
  - 返回结构增加 reasons: dict[rule_id, str],每条违规带具体原因文本
  - 原因文本含「现场值」+「期望值」,方便前端直接展示
- routes/idcard_scan.py
  - BadRow 增加 violation 字段,扫描时把 reasons 串成单行文本塞进每行
  - BadTable.rows_meta 把字段注释 + 数据类型带出来给前端表头 tooltip
  - ProgressTracker 加 _partial_tables + partial_snapshot()
  - 新增 GET /api/id-card/scan/{scan_id}/partial 端点(轮询用)
  - 旧 /stream SSE 端点保留但前端不再调用
parent c7fe08c5
...@@ -6,17 +6,27 @@ ...@@ -6,17 +6,27 @@
GB/T 2260 地址码(最简校验:非 00 前缀即可,完整区划表留待外部数据接入) GB/T 2260 地址码(最简校验:非 00 前缀即可,完整区划表留待外部数据接入)
校验项: 校验项:
1. LENGTH 长度不是 18 位 / 15 位 1. EMPTY 值为空 / None
2. CHARSET 字符集不合规(前 17 位非数字 / 第 18 位非 0-9 或 X) 2. LENGTH 长度不是 18 位 / 15 位
3. CHECK_DIGIT 校验位计算错误(GB 11643 §6) 3. CHARSET 字符集不合规(前 17 位非数字 / 第 18 位非 0-9 或 X)
4. ADDRESS_PREFIX 地址码全 0 或前 2 位为 00 4. CHECK_DIGIT 校验位计算错误(GB 11643 §6)
5. BIRTH_DATE 出生日期不是真实日期 / 晚于今天 / 早于 1900 5. ADDRESS_PREFIX 地址码全 0 或前 2 位为 00
6. ORDER_CODE 顺序码不在 001-999 范围 6. BIRTH_DATE 出生日期不是真实日期 / 晚于今天 / 早于 1900
7. ORDER_CODE 顺序码不在 001-999 范围
返回结构: 返回结构:
{"ok": bool, "violations": [rule_id, ...], "normalized": str | None} {
- ok=True 表示合规;violations 为空 "ok": bool, # 合规
- normalized:如果是 15 位老号 → 升 18 位后的值;否则 None "violations": [rule_id, ...], # 违反的规则 ID
"reasons": { rule_id: str, ... }, # 每条违规的具体原因(含现场值)
"length": int, # 输入字符串长度
"is_15_old": bool, # 是否 15 位老号
"normalized": str | None # 15 位老号升 18 位后的值
}
设计:
- reasons 是按 rule_id 去重的 dict(一类违规只给一条原因文本)
- 原因文本含「现场值」+「期望值」,方便定位
""" """
from __future__ import annotations from __future__ import annotations
...@@ -37,52 +47,88 @@ def _is_all_same(s: str) -> bool: ...@@ -37,52 +47,88 @@ def _is_all_same(s: str) -> bool:
return len(s) > 1 and len(set(s)) == 1 return len(s) > 1 and len(set(s)) == 1
def _validate_18(code: str) -> tuple[bool, list[str], str | None]: def _validate_18(code: str) -> tuple[bool, list[str], dict[str, str], str | None]:
"""校验 18 位身份证号。返回 (ok, violations, normalized)""" """校验 18 位身份证号。返回 (ok, violations, reasons, normalized)"""
violations: list[str] = [] violations: list[str] = []
normalized = code reasons: dict[str, str] = {}
# 字符集 # 字符集
body, last = code[:17], code[17] body, last = code[:17], code[17]
if not body.isdigit(): if not body.isdigit():
# 找出第一个非数字字符的位置
bad_pos = next((i for i, ch in enumerate(body) if not ch.isdigit()), -1)
reasons["CHARSET"] = (
f"前 17 位必须是数字,第 {bad_pos + 1 if bad_pos >= 0 else '?'} 位不是数字"
)
violations.append("CHARSET") violations.append("CHARSET")
if not (last.isdigit() or last == "X"): if not (last.isdigit() or last == "X"):
reasons["CHARSET"] = reasons.get(
"CHARSET",
f"第 18 位(校验位)必须为 0-9 或大写 X,实际是 '{last}'",
)
if "CHARSET" not in violations:
violations.append("CHARSET") violations.append("CHARSET")
# 不再继续算校验位(last 不合法) # 字符集已经不合规,校验位/地址码/出生日期/顺序码都没意义,跳过
return bool(violations), violations, normalized return bool(violations), violations, reasons, code
# 全相同 # 全相同
if _is_all_same(code): if _is_all_same(code):
violations.append("CHARSET") # 复用 CHARSET 槽位:占位符类 reasons["CHARSET"] = f"全部字符相同({code[0]} × {len(code)}),不是合法身份证"
return False, violations, normalized violations.append("CHARSET")
return False, violations, reasons, code
# 校验位 # 校验位
s = sum(int(body[i]) * _WEIGHTS[i] for i in range(17)) s = sum(int(body[i]) * _WEIGHTS[i] for i in range(17))
expected = _CHECK_MAP[s % 11] expected = _CHECK_MAP[s % 11]
if expected != last: if expected != last:
reasons["CHECK_DIGIT"] = (
f"校验位计算结果应为 '{expected}',实际为 '{last}'"
)
violations.append("CHECK_DIGIT") violations.append("CHECK_DIGIT")
# 地址码(前 6 位) # 地址码(前 6 位)
region = code[:6] region = code[:6]
if region == "000000" or region[:2] == "00": if region == "000000":
reasons["ADDRESS_PREFIX"] = f"地址码全 0({region})"
violations.append("ADDRESS_PREFIX")
elif region[:2] == "00":
reasons["ADDRESS_PREFIX"] = (
f"地址码前 2 位为 00({region[:2]}),省/直辖市代码不存在"
)
violations.append("ADDRESS_PREFIX") violations.append("ADDRESS_PREFIX")
# 出生日期(第 7-14 位) # 出生日期(第 7-14 位)
try: try:
bd = _dt.date(int(code[6:10]), int(code[10:12]), int(code[12:14])) bd = _dt.date(int(code[6:10]), int(code[10:12]), int(code[12:14]))
bd_str = f"{code[6:10]}-{code[10:12]}-{code[12:14]}"
if bd > _MAX_BIRTH_DATE: if bd > _MAX_BIRTH_DATE:
reasons["BIRTH_DATE"] = (
f"出生日期 {bd_str} 晚于今天({_MAX_BIRTH_DATE.isoformat()})"
)
violations.append("BIRTH_DATE") violations.append("BIRTH_DATE")
elif bd.year < _MIN_BIRTH_YEAR: elif bd.year < _MIN_BIRTH_YEAR:
reasons["BIRTH_DATE"] = (
f"出生日期 {bd_str} 早于 {_MIN_BIRTH_YEAR} 年"
)
violations.append("BIRTH_DATE") violations.append("BIRTH_DATE")
except ValueError: except ValueError:
reasons["BIRTH_DATE"] = (
f"出生日期段 '{code[6:14]}' 不是真实日期(YYYYMMDD)"
)
violations.append("BIRTH_DATE") violations.append("BIRTH_DATE")
# 顺序码(15-17 位) # 顺序码(15-17 位)
order = code[14:17] order = code[14:17]
if not order.isdigit() or not (1 <= int(order) <= 999): if not order.isdigit():
reasons["ORDER_CODE"] = f"顺序码 '{order}' 含非数字字符"
violations.append("ORDER_CODE")
elif not (1 <= int(order) <= 999):
reasons["ORDER_CODE"] = (
f"顺序码 '{order}' 不在 001-999 范围(000 通常为预留/特殊用途)"
)
violations.append("ORDER_CODE") violations.append("ORDER_CODE")
return (not violations), violations, normalized return (not violations), violations, reasons, code
def _lift_15_to_18(code15: str) -> str: def _lift_15_to_18(code15: str) -> str:
...@@ -105,59 +151,119 @@ def validate(value) -> dict: ...@@ -105,59 +151,119 @@ def validate(value) -> dict:
Returns: Returns:
{ {
"ok": bool, # 合规 "ok": bool,
"violations": [str], # 违反的规则 ID(CHARSET / LENGTH / CHECK_DIGIT / ...) "violations": [str], # 违反的规则 ID
"length": int, # 输入字符串长度 "reasons": {rule_id: str}, # 每条违规的具体原因
"is_15_old": bool, # 是否 15 位老号 "length": int,
"normalized": str|None # 15 位老号升 18 位后的值;合规 18 位也是它本身 "is_15_old": bool,
"normalized": str|None
} }
""" """
if value is None: if value is None:
return {"ok": False, "violations": ["EMPTY"], "length": 0, "is_15_old": False, "normalized": None} return {
"ok": False,
"violations": ["EMPTY"],
"reasons": {"EMPTY": "值为 NULL"},
"length": 0,
"is_15_old": False,
"normalized": None,
}
s = str(value).strip() s = str(value).strip()
if not s: if not s:
return {"ok": False, "violations": ["EMPTY"], "length": 0, "is_15_old": False, "normalized": None} return {
"ok": False,
"violations": ["EMPTY"],
"reasons": {"EMPTY": "值为空字符串"},
"length": 0,
"is_15_old": False,
"normalized": None,
}
if _is_all_same(s) and len(s) >= 15: if _is_all_same(s) and len(s) >= 15:
return {"ok": False, "violations": ["CHARSET"], "length": len(s), "is_15_old": False, "normalized": None} return {
"ok": False,
"violations": ["CHARSET"],
"reasons": {
"CHARSET": f"全部字符相同({s[0]} × {len(s)}),不是合法身份证",
},
"length": len(s),
"is_15_old": False,
"normalized": None,
}
if len(s) == 15 and s.isdigit(): if len(s) == 15 and s.isdigit():
# 15 位老号:单独校验(没有校验位) # 15 位老号:单独校验(没有校验位)
violations: list[str] = [] violations: list[str] = []
if region := s[:6]: reasons: dict[str, str] = {}
if region == "000000" or region[:2] == "00": # 地址码
region = s[:6]
if region == "000000":
reasons["ADDRESS_PREFIX"] = f"地址码全 0({region})"
violations.append("ADDRESS_PREFIX")
elif region[:2] == "00":
reasons["ADDRESS_PREFIX"] = (
f"地址码前 2 位为 00({region[:2]}),省/直辖市代码不存在"
)
violations.append("ADDRESS_PREFIX") violations.append("ADDRESS_PREFIX")
# 出生日期
try: try:
bd = _dt.date(1900 + int(s[6:8]), int(s[8:10]), int(s[10:12])) bd = _dt.date(1900 + int(s[6:8]), int(s[8:10]), int(s[10:12]))
if bd > _MAX_BIRTH_DATE or bd.year < _MIN_BIRTH_YEAR: bd_str = f"19{s[6:8]}-{s[8:10]}-{s[10:12]}"
if bd > _MAX_BIRTH_DATE:
reasons["BIRTH_DATE"] = (
f"出生日期 {bd_str} 晚于今天({_MAX_BIRTH_DATE.isoformat()})"
)
violations.append("BIRTH_DATE")
elif bd.year < _MIN_BIRTH_YEAR:
reasons["BIRTH_DATE"] = f"出生日期 {bd_str} 早于 {_MIN_BIRTH_YEAR} 年"
violations.append("BIRTH_DATE") violations.append("BIRTH_DATE")
except ValueError: except ValueError:
reasons["BIRTH_DATE"] = (
f"出生日期段 '{s[6:12]}' 不是真实日期(YYMMDD)"
)
violations.append("BIRTH_DATE") violations.append("BIRTH_DATE")
# 顺序码
order = s[12:15] order = s[12:15]
if not order.isdigit() or not (1 <= int(order) <= 999): if not order.isdigit():
reasons["ORDER_CODE"] = f"顺序码 '{order}' 含非数字字符"
violations.append("ORDER_CODE")
elif not (1 <= int(order) <= 999):
reasons["ORDER_CODE"] = (
f"顺序码 '{order}' 不在 001-999 范围"
)
violations.append("ORDER_CODE") violations.append("ORDER_CODE")
# 15 位长度本身不算违规(合法老号)
return { return {
"ok": not violations, "ok": not violations,
"violations": violations, "violations": violations,
"reasons": reasons,
"length": 15, "length": 15,
"is_15_old": True, "is_15_old": True,
"normalized": _lift_15_to_18(s), "normalized": _lift_15_to_18(s),
} }
if len(s) != 18: if len(s) != 18:
return {"ok": False, "violations": ["LENGTH"], "length": len(s), "is_15_old": False, "normalized": None} return {
"ok": False,
"violations": ["LENGTH"],
"reasons": {
"LENGTH": f"身份证应为 18 位,实际 {len(s)} 位(值:'{s[:30]}')",
},
"length": len(s),
"is_15_old": False,
"normalized": None,
}
ok, violations, normalized = _validate_18(s) ok, violations, reasons, normalized = _validate_18(s)
return { return {
"ok": ok, "ok": ok,
"violations": violations, "violations": violations,
"reasons": reasons,
"length": 18, "length": 18,
"is_15_old": False, "is_15_old": False,
"normalized": normalized, "normalized": normalized,
} }
# ── 规则 ID → 中文描述(前端展示用)─────────────────────── # ── 规则 ID → 中文标签(前端展示用)───────────────────────
VIOLATION_LABELS = { VIOLATION_LABELS = {
"EMPTY": "空值", "EMPTY": "空值",
"LENGTH": "长度不是 18 位", "LENGTH": "长度不是 18 位",
......
...@@ -131,6 +131,8 @@ class ScanProgressResponse(BaseModel): ...@@ -131,6 +131,8 @@ class ScanProgressResponse(BaseModel):
status: str # "pending" | "running" | "done" | "error" status: str # "pending" | "running" | "done" | "error"
done: int done: int
total: int total: int
tables_done: int
tables_total: int
current_column: str | None # 最近报上来的列(多个 worker 抢着写,频繁变) current_column: str | None # 最近报上来的列(多个 worker 抢着写,频繁变)
elapsed_ms: int elapsed_ms: int
error: str | None = None error: str | None = None
...@@ -243,7 +245,12 @@ def _fetch_bad_rows( ...@@ -243,7 +245,12 @@ def _fetch_bad_rows(
# ── 进度跟踪器(线程安全,存每个 scan_id 的状态)────────────── # ── 进度跟踪器(线程安全,存每个 scan_id 的状态)──────────────
class ProgressTracker: class ProgressTracker:
"""单实例即可(每个 scan_id 一份 dict 字段),所有 worker 共享同一份内存。""" """单实例即可(每个 scan_id 一份 dict 字段),所有 worker 共享同一份内存。
新增 SSE 流式支持:
- events: queue.Queue,worker 把每张表完成/进度事件塞进来,SSE 端点异步取
- tables_done / tables_total:用于进度卡片「已完成表 X / Y」
"""
def __init__(self, scan_id: str, total: int, started_monotonic: float): def __init__(self, scan_id: str, total: int, started_monotonic: float):
self.scan_id = scan_id self.scan_id = scan_id
...@@ -255,6 +262,13 @@ class ProgressTracker: ...@@ -255,6 +262,13 @@ class ProgressTracker:
self._status = "running" # pending | running | done | error self._status = "running" # pending | running | done | error
self._error: str | None = None self._error: str | None = None
self._result: IdCardScanResponse | None = None self._result: IdCardScanResponse | None = None
# 流式输出 / 轮询 partial
import queue as _queue
self.events: _queue.Queue[dict] = _queue.Queue()
self._tables_done = 0
self._tables_total = 0
# 已完成的违规表(轮询 partial 用,每张表完成就追加)
self._partial_tables: list["BadTable"] = []
def start_column(self, tname: str, cname: str) -> None: def start_column(self, tname: str, cname: str) -> None:
with self._lock: with self._lock:
...@@ -268,11 +282,23 @@ class ProgressTracker: ...@@ -268,11 +282,23 @@ class ProgressTracker:
with self._lock: with self._lock:
self._status = "error" self._status = "error"
self._error = err self._error = err
# 推一个 error 事件给前端,然后发 sentinel 让 SSE 关流
self.events.put({"event": "error_evt", "data": {"message": err}})
self.events.put(_SENTINEL_DONE)
def set_result(self, resp: IdCardScanResponse) -> None: def set_result(self, resp: IdCardScanResponse) -> None:
with self._lock: with self._lock:
self._status = "done" self._status = "done"
self._result = resp self._result = resp
# 推 done 事件 + sentinel 关流
self.events.put({
"event": "done",
"data": {
"overview": resp.overview.model_dump() if hasattr(resp.overview, "model_dump") else resp.overview,
"matches": [m.model_dump() if hasattr(m, "model_dump") else m for m in resp.matches],
},
})
self.events.put(_SENTINEL_DONE)
def snapshot(self) -> dict: def snapshot(self) -> dict:
with self._lock: with self._lock:
...@@ -281,12 +307,63 @@ class ProgressTracker: ...@@ -281,12 +307,63 @@ class ProgressTracker:
"status": self._status, "status": self._status,
"done": self._done, "done": self._done,
"total": self.total, "total": self.total,
"tables_done": self._tables_done,
"tables_total": self._tables_total,
"current_column": self._current_column, "current_column": self._current_column,
"elapsed_ms": int((time.monotonic() - self.started_monotonic) * 1000), "elapsed_ms": int((time.monotonic() - self.started_monotonic) * 1000),
"error": self._error, "error": self._error,
"result": self._result, "result": self._result,
} }
def push_table_done(self, table: "BadTable") -> None:
"""一张问题表的扫描+反查全部完成,推送 SSE 事件 + 累计到 partial_tables(轮询用)"""
with self._lock:
self._tables_done += 1
self._partial_tables.append(table)
self.events.put({
"event": "table_done",
"data": table.model_dump() if hasattr(table, "model_dump") else table,
})
def partial_snapshot(self) -> dict:
"""给前端轮询用的 partial snapshot(status / tables / overview 累计)"""
with self._lock:
tables = list(self._partial_tables)
tables_done = self._tables_done
elapsed_ms = int((time.monotonic() - self.started_monotonic) * 1000)
total_bad_rows = sum(t.bad_row_count for t in tables)
return {
"scan_id": self.scan_id,
"status": self._status,
"tables_done": tables_done,
"tables_total": self._tables_total,
"elapsed_ms": elapsed_ms,
"violations_by_table": [t.model_dump() if hasattr(t, "model_dump") else t for t in tables],
"overview_partial": {
"total_bad_rows_so_far": total_bad_rows,
"tables_count_so_far": len(tables),
},
}
def push_overview(self, overview: "ScanOverview", matches: list) -> None:
"""推送中途 overview 统计(用于前端顶部概览区)"""
self.events.put({
"event": "overview",
"data": {
"overview": overview.model_dump() if hasattr(overview, "model_dump") else overview,
"matches": [m.model_dump() if hasattr(m, "model_dump") else m for m in matches],
},
})
def set_tables_total(self, n: int) -> None:
"""预计总共要扫多少张表(带违规的表数;运行时增量更新更精确)"""
with self._lock:
self._tables_total = max(self._tables_total, n)
# Sentinel:SSE 端点读到它就结束流
_SENTINEL_DONE: dict = {"event": "__sentinel__", "data": None}
# 全局 in-memory 任务表:scan_id -> ProgressTracker # 全局 in-memory 任务表:scan_id -> ProgressTracker
# (web2 是单进程内网工具;如要多进程/多 worker 需要换成 Redis) # (web2 是单进程内网工具;如要多进程/多 worker 需要换成 Redis)
...@@ -320,7 +397,13 @@ def _scan_one_column( ...@@ -320,7 +397,13 @@ def _scan_one_column(
"""跑一列:开连接 → SELECT DISTINCT → validate → 关闭 """跑一列:开连接 → SELECT DISTINCT → validate → 关闭
Returns: Returns:
{table_name, column_name, column_meta, distinct_count, invalid_count, violations, bad_values} {
table_name, column_name, column_meta,
distinct_count, invalid_count,
violations: {rule_id: count}, # 用于统计
bad_values: [value, ...], # 用于反查
bad_reasons: {value: reason_str, ...}, # 用于前端展示「这条数据为什么不合规」
}
失败时返回 None 失败时返回 None
""" """
tname = col["table_name"] tname = col["table_name"]
...@@ -337,11 +420,16 @@ def _scan_one_column( ...@@ -337,11 +420,16 @@ def _scan_one_column(
counter: Counter[str] = Counter() counter: Counter[str] = Counter()
bad_vals: list[str] = [] bad_vals: list[str] = []
bad_reasons: dict[str, str] = {}
try: try:
for v in distinct_vals: for v in distinct_vals:
r = validate(v) r = validate(v)
if r["violations"]: if r["violations"]:
bad_vals.append(v) bad_vals.append(v)
# 取第一条原因文本(已按优先级排序:CHARSET > LENGTH > ...)
# 同一值的多条违规 reason 中,列出第一条作为「主原因」
first_reason = next(iter(r["reasons"].values()), "不合规")
bad_reasons[v] = first_reason
for rule in r["violations"]: for rule in r["violations"]:
counter[rule] += 1 counter[rule] += 1
except Exception as e: except Exception as e:
...@@ -360,6 +448,7 @@ def _scan_one_column( ...@@ -360,6 +448,7 @@ def _scan_one_column(
"invalid_count": len(bad_vals), "invalid_count": len(bad_vals),
"violations": dict(counter), "violations": dict(counter),
"bad_values": bad_vals, "bad_values": bad_vals,
"bad_reasons": bad_reasons,
} }
except Exception as e: except Exception as e:
logger.warning(f" · {tname}.{cname} worker 异常:{type(e).__name__}: {e}") logger.warning(f" · {tname}.{cname} worker 异常:{type(e).__name__}: {e}")
...@@ -374,10 +463,14 @@ def _fetch_bad_rows_one( ...@@ -374,10 +463,14 @@ def _fetch_bad_rows_one(
tname: str, tname: str,
cname: str, cname: str,
bad_vals: list[str], bad_vals: list[str],
bad_reasons: dict[str, str],
max_rows: int, max_rows: int,
progress: ProgressTracker, progress: ProgressTracker,
) -> tuple[str, str, list[dict[str, Any]]]: ) -> tuple[str, str, list[dict[str, Any]]]:
"""反查某列的违规行;返回 (tname, cname, rows)""" """反查某列的违规行;返回 (tname, cname, rows_with_reason)
rows_with_reason 每行 = 原数据库行 + '_id_card_violation' 字段(违规原因)
"""
progress.start_column(tname, cname) progress.start_column(tname, cname)
cfg = DBConfig(**cfg_dict) cfg = DBConfig(**cfg_dict)
try: try:
...@@ -387,7 +480,14 @@ def _fetch_bad_rows_one( ...@@ -387,7 +480,14 @@ def _fetch_bad_rows_one(
except Exception as e: except Exception as e:
logger.warning(f" · {tname}.{cname} 反查失败:{type(e).__name__}: {e}") logger.warning(f" · {tname}.{cname} 反查失败:{type(e).__name__}: {e}")
return (tname, cname, []) return (tname, cname, [])
return (tname, cname, rows) # 给每行加上 _id_card_violation 字段
enriched = []
for row in rows:
v = row.get(cname)
v_str = "" if v is None else str(v)
row["_id_card_violation"] = bad_reasons.get(v_str, "不合规")
enriched.append(row)
return (tname, cname, enriched)
except Exception as e: except Exception as e:
logger.warning(f" · {tname}.{cname} 反查 worker 异常:{type(e).__name__}: {e}") logger.warning(f" · {tname}.{cname} 反查 worker 异常:{type(e).__name__}: {e}")
return (tname, cname, []) return (tname, cname, [])
...@@ -424,6 +524,23 @@ def _run_scan_in_thread( ...@@ -424,6 +524,23 @@ def _run_scan_in_thread(
"oracle_client_dir": cfg.oracle_client_dir, "oracle_client_dir": cfg.oracle_client_dir,
} }
# 启动一个心跳线程:每 500ms 推一个 progress 事件,让前端 elapsed_ms 不卡
_heartbeat_stop = threading.Event()
def _heartbeat():
while not _heartbeat_stop.is_set():
try:
snap = progress.snapshot()
snap.pop("result", None) # result 可能很大,不放进 SSE
progress.events.put({
"event": "progress",
"data": snap,
})
except Exception:
pass
_heartbeat_stop.wait(0.5)
hb_thread = threading.Thread(target=_heartbeat, daemon=True, name=f"hb-{scan_id[:8]}")
hb_thread.start()
# ① 抽数据字典 # ① 抽数据字典
logger.info("[1/4] 抽取数据字典 …") logger.info("[1/4] 抽取数据字典 …")
t0 = time.monotonic() t0 = time.monotonic()
...@@ -468,30 +585,40 @@ def _run_scan_in_thread( ...@@ -468,30 +585,40 @@ def _run_scan_in_thread(
per_col_results.sort(key=lambda r: order.get((r["table_name"], r["column_name"]), 0)) per_col_results.sort(key=lambda r: order.get((r["table_name"], r["column_name"]), 0))
# ④ 并发反查违规行(每列一线程) # ④ 并发反查违规行(每列一线程)
bad_values_by_table_col: dict[tuple[str, str], list[str]] = { # 同时把 bad_reasons 也带上:反查时给每行加 _id_card_violation 字段
(r["table_name"], r["column_name"]): r["bad_values"] bad_data_by_table_col: dict[tuple[str, str], tuple[list[str], dict[str, str]]] = {
(r["table_name"], r["column_name"]): (r["bad_values"], r["bad_reasons"])
for r in per_col_results if r["bad_values"] for r in per_col_results if r["bad_values"]
} }
# 重置 progress:④ 阶段重新从 0 计数 # 重置 progress:④ 阶段重新从 0 计数
with progress._lock: with progress._lock:
progress._done = 0 progress._done = 0
progress.total = len(bad_values_by_table_col) if bad_values_by_table_col else 1 progress.total = len(bad_data_by_table_col) if bad_data_by_table_col else 1
logger.info( logger.info(
f"[4/4] 并发反查完整行({len(bad_values_by_table_col)} 列,workers={req.max_workers})…" f"[4/4] 并发反查完整行({len(bad_data_by_table_col)} 列,workers={req.max_workers})…"
) )
# 按表聚合结果 # 按表聚合结果;每张表的所有列都反查完 → 立即 push 给前端(SSE 流式输出)
violations_by_table_dict: dict[str, dict] = {} violations_by_table_dict: dict[str, dict] = {}
# 每张表需要反查的列数:用于判断「这张表是否全部完成」
pending_per_table: dict[str, int] = {}
for (tname, _cname) in bad_data_by_table_col.keys():
pending_per_table[tname] = pending_per_table.get(tname, 0) + 1
progress.set_tables_total(len(pending_per_table))
# 把所有列的元数据/违规规则预计算出来(避免后面边遍历边查 per_col_results)
col_to_meta = {(r["table_name"], r["column_name"]): r["column_meta"] for r in per_col_results}
col_to_violations = {(r["table_name"], r["column_name"]): r["violations"] for r in per_col_results}
with ThreadPoolExecutor(max_workers=req.max_workers, thread_name_prefix="bad-row") as ex: with ThreadPoolExecutor(max_workers=req.max_workers, thread_name_prefix="bad-row") as ex:
futures = [ futures = [
ex.submit( ex.submit(
_fetch_bad_rows_one, cfg_dict, tname, cname, bad_vals, _fetch_bad_rows_one, cfg_dict, tname, cname, bad_vals, bad_reasons,
req.max_full_rows_per_column, progress, req.max_full_rows_per_column, progress,
) )
for (tname, cname), bad_vals in bad_values_by_table_col.items() for (tname, cname), (bad_vals, bad_reasons) in bad_data_by_table_col.items()
] ]
# 收集列元数据索引(for 拼 BadTable.columns)
col_to_meta = {(r["table_name"], r["column_name"]): r["column_meta"] for r in per_col_results}
for fut in as_completed(futures): for fut in as_completed(futures):
tname, cname, rows = fut.result() tname, cname, rows = fut.result()
rec = violations_by_table_dict.setdefault(tname, { rec = violations_by_table_dict.setdefault(tname, {
...@@ -506,19 +633,37 @@ def _run_scan_in_thread( ...@@ -506,19 +633,37 @@ def _run_scan_in_thread(
rec["columns"] = tcols rec["columns"] = tcols
rec["table_comment"] = tcols[0].get("table_comment") if tcols else None rec["table_comment"] = tcols[0].get("table_comment") if tcols else None
# 把这一列的违规规则加入 # 把这一列的违规规则加入
if (tname, cname) in col_to_meta: rec["rule_set"].update(col_to_violations.get((tname, cname), {}).keys())
rec["rule_set"].update(col_to_meta[(tname, cname)].get("_violations", []) or [])
# 从 per_col_results 拿到该列的 violations
for r in per_col_results:
if r["table_name"] == tname and r["column_name"] == cname:
rec["rule_set"].update(r["violations"].keys())
break
for row in rows: for row in rows:
key = tuple(sorted((k, str(v)) for k, v in row.items())) key = tuple(sorted((k, str(v)) for k, v in row.items()))
rec["bad_rows_set"].add(key) rec["bad_rows_set"].add(key)
logger.info(f" · {tname}.{cname} → {len(rows)} 行(去重后累计 {len(rec['bad_rows_set'])})") pending_per_table[tname] = pending_per_table.get(tname, 1) - 1
logger.info(
f" · {tname}.{cname} → {len(rows)} 行 "
f"(表累计 {len(rec['bad_rows_set'])},剩余列 {pending_per_table[tname]})"
)
# ── 流式推送:这张表的所有列都反查完 → 立刻 push ──
if pending_per_table[tname] <= 0:
bad_rows = [dict(kv) for kv in rec["bad_rows_set"]]
fixed_rows = []
for row in bad_rows:
r2 = {}
for k, v in row.items():
r2[k] = v
fixed_rows.append(r2)
bt = BadTable(
table_name=tname,
table_comment=rec["table_comment"],
columns=rec["columns"],
bad_rows=fixed_rows,
rule_summary=sorted(rec["rule_set"]),
bad_row_count=len(fixed_rows),
)
progress.push_table_done(bt)
logger.info(f" · ⏩ 流式推送:{tname}({len(fixed_rows)} 行违规)")
# 拼最终响应 # 拼最终响应(用 violations_by_table_dict 直接构造,不再重新聚合)
violations_by_table: list[BadTable] = [] violations_by_table: list[BadTable] = []
total_bad_rows = 0 total_bad_rows = 0
for tname, rec in violations_by_table_dict.items(): for tname, rec in violations_by_table_dict.items():
...@@ -577,10 +722,17 @@ def _run_scan_in_thread( ...@@ -577,10 +722,17 @@ def _run_scan_in_thread(
matches=matches_out, matches=matches_out,
violations_by_table=violations_by_table, violations_by_table=violations_by_table,
) )
# 推送最终的 overview(前端会在 done 事件中也收到,但这里冗余一份更稳)
progress.set_result(response) progress.set_result(response)
except Exception as e: except Exception as e:
logger.exception(f"扫描后台线程异常: {e}") logger.exception(f"扫描后台线程异常: {e}")
progress.set_error(f"{type(e).__name__}: {e}") progress.set_error(f"{type(e).__name__}: {e}")
finally:
# 停心跳
try:
_heartbeat_stop.set()
except NameError:
pass
def _build_response( def _build_response(
...@@ -654,7 +806,10 @@ async def idcard_scan_start(req: IdCardScanRequest) -> ScanStartedResponse: ...@@ -654,7 +806,10 @@ async def idcard_scan_start(req: IdCardScanRequest) -> ScanStartedResponse:
@router.get("/id-card/scan/{scan_id}/progress", response_model=ScanProgressResponse) @router.get("/id-card/scan/{scan_id}/progress", response_model=ScanProgressResponse)
async def idcard_scan_progress(scan_id: str) -> ScanProgressResponse: async def idcard_scan_progress(scan_id: str) -> ScanProgressResponse:
"""拿扫描进度(前端每 500-1000ms 轮询一次)""" """拿扫描进度(前端每 500-1000ms 轮询一次)
注:当前主路径已改为 SSE 流式(/stream);此端点保留用于调试或轮询 fallback。
"""
with _SCANS_LOCK: with _SCANS_LOCK:
progress = _SCANS.get(scan_id) progress = _SCANS.get(scan_id)
if progress is None: if progress is None:
...@@ -665,12 +820,147 @@ async def idcard_scan_progress(scan_id: str) -> ScanProgressResponse: ...@@ -665,12 +820,147 @@ async def idcard_scan_progress(scan_id: str) -> ScanProgressResponse:
status=snap["status"], status=snap["status"],
done=snap["done"], done=snap["done"],
total=snap["total"], total=snap["total"],
tables_done=snap["tables_done"],
tables_total=snap["tables_total"],
current_column=snap["current_column"], current_column=snap["current_column"],
elapsed_ms=snap["elapsed_ms"], elapsed_ms=snap["elapsed_ms"],
error=snap["error"], error=snap["error"],
) )
# ── SSE 流式端点(每张表完成后立即推前端)────────────────────
import asyncio
import json as _json
def _sse_pack(event: str, data) -> str:
"""组装一条 SSE 消息。
SSE 格式:
event: <name>\\n
data: <json>\\n
\\n
"""
payload = _json.dumps(data, ensure_ascii=False, default=str)
return f"event: {event}\ndata: {payload}\n\n"
@router.get("/id-card/scan/{scan_id}/partial", summary="轮询用 partial 结果")
async def idcard_scan_partial(scan_id: str):
"""给前端 setInterval(1s) 轮询用。
返回:
- 200 + {status: 'done', response: IdCardScanResponse} 扫描完成
- 200 + {status: 'running', violations_by_table: [...], overview_partial, ...} 扫描中(partial)
- 200 + {status: 'error', error: '...'} 扫描失败
- 404 scan_id 不存在或过期
"""
with _SCANS_LOCK:
progress = _SCANS.get(scan_id)
if progress is None:
raise HTTPException(status_code=404, detail=f"scan_id 不存在或已过期:{scan_id}")
snap = progress.partial_snapshot()
status = snap["status"]
if status == "done" and progress._result is not None:
# 已完成:返回完整响应
result = progress._result
return {
"scan_id": scan_id,
"status": "done",
"response": result.model_dump() if hasattr(result, "model_dump") else result,
}
if status == "error":
return {
"scan_id": scan_id,
"status": "error",
"error": progress._error,
}
# running / pending:返回 partial
return {
"scan_id": scan_id,
"status": status,
"tables_done": snap["tables_done"],
"tables_total": snap["tables_total"],
"elapsed_ms": snap["elapsed_ms"],
"violations_by_table": snap["violations_by_table"],
"overview_partial": snap["overview_partial"],
}
@router.get("/id-card/scan/{scan_id}/stream", summary="SSE 流:每张问题表扫描完即推送")
async def idcard_scan_stream(scan_id: str):
"""SSE 流式端点。
事件类型:
- connected 链路建立确认(前端 EventSource onopen)
- progress {done, total, tables_done, tables_total, current_column, elapsed_ms}
- table_done 单张问题表的完整数据(BadTable);前端立刻 push 到列表
- overview 扫描完成时推送总览 + matches
- error_evt 错误事件(带 message 字段)
关闭条件:
- 服务端推完 done / error_evt 后主动断开(yield 后 return)
- 客户端 EventSource.close()
"""
with _SCANS_LOCK:
progress = _SCANS.get(scan_id)
if progress is None:
raise HTTPException(status_code=404, detail=f"scan_id 不存在或已过期:{scan_id}")
queue = progress.events
logger.info(f" · SSE 流打开 scan={scan_id[:8]}")
async def event_gen():
# 1) 链路建立
yield _sse_pack("connected", {"scan_id": scan_id})
# 2) 启动时立即发一个 progress(让前端拿到 total=0 也能立刻有反应)
yield _sse_pack("progress", progress.snapshot())
loop = asyncio.get_event_loop()
while True:
# 阻塞读队列(异步化:丢到 thread executor 避免阻塞事件循环)
try:
evt = await asyncio.wait_for(
loop.run_in_executor(None, queue.get, True, 1.0),
timeout=2.0,
)
except asyncio.TimeoutError:
# 心跳(注释行,SSE 规范允许注释作为心跳)
yield ": ping\n\n"
continue
except Exception as e:
logger.warning(f"SSE 队列读取异常: {e}")
yield _sse_pack("error_evt", {"message": str(e)})
break
# sentinel → 关闭流
if evt.get("event") == "__sentinel__":
logger.info(f" · SSE 流关闭 scan={scan_id[:8]}")
break
ev_type = evt.get("event")
ev_data = evt.get("data")
if ev_type == "progress":
# progress 事件单独不发(snapshot 自带),但允许外部触发
pass
yield _sse_pack(ev_type, ev_data)
from fastapi.responses import StreamingResponse
return StreamingResponse(
event_gen(),
media_type="text/event-stream",
headers={
# SSE 需要这些 header 才能跨域 + 保持连接
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no", # 禁用 nginx buffering(如果后面挂 nginx)
"Connection": "keep-alive",
},
)
@router.get("/id-card/scan/{scan_id}/result", response_model=ScanResultResponse) @router.get("/id-card/scan/{scan_id}/result", response_model=ScanResultResponse)
async def idcard_scan_result(scan_id: str) -> ScanResultResponse: async def idcard_scan_result(scan_id: str) -> ScanResultResponse:
"""拿扫描最终结果(status=pending/running → 202;done → 200;error → 500)""" """拿扫描最终结果(status=pending/running → 202;done → 200;error → 500)"""
......
Markdown is supported
0%
or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment