Commit 63b3b770 authored by Data Governance Dev's avatar Data Governance Dev

feat(web3): 查询改分页拉取 + 中断按钮 + 后端 session 隔离

把原来 /api/queries/run 单次同步全表扫描改成 3 个端点:
  POST /api/queries/start    COUNT + 编译规则 + 建 session(UUID)
  POST /api/queries/page     按页号拉一页,跑规则,只返回不合规的行
  POST /api/queries/cancel   摘掉 session(best-effort 中断)

session_manager.py 用模块级 dict + Lock 存 QuerySession,TTL 30min 懒清理。
每页新开 DBConnection → fetchall(避开 cursor 生命周期管理),~50ms 重连可接受。

dialect 兼容:
  mysql    : LIMIT <n> OFFSET <m>
  dameng   : OFFSET <m> ROWS FETCH NEXT <n> ROWS ONLY (SQL:2008)
  oracle   : ROWNUM 三层嵌套(兼容 11g/10g;T4 任务实测 11g 服务端用 OFFSET/FETCH NEXT 报 ORA-00933)
  postgres / sqlserver : 同 dameng

前端友好的关键点:
  - /start 一次返回 fields(含 rule_list + field_comment),前端 resultFields 不再变
  - 后端给每行注入 __row_index = (page_no-1)*page_size + idx,给前端做 Vue row-key
  - /page 失败不摘 session,让前端可以重试或走 /cancel

测试:新增 web3/tests/test_queries.py 15 条(start / page / cancel / paginate_sql / 并发 / TTL),
老的 test_run_query_only_bad_rows.py 和 test_rule_types_e2e.py 适配新端点(删 max_rows 那条,
end-to-end 改成 /start + /page 聚合)。73 条全过。
parent 9263f8d7
...@@ -121,6 +121,54 @@ def quote_ident(name: str, db_type: str) -> str: ...@@ -121,6 +121,54 @@ def quote_ident(name: str, db_type: str) -> str:
return f'"{name}"' return f'"{name}"'
# ── 分页 SQL 助手 ────────────────────────────────────────
def paginate_sql(base_sql: str, db_type: str, offset: int, limit: int) -> str:
"""在 base_sql 后追加 LIMIT/OFFSET 子句,屏蔽方言差异。
dialect:
- mysql : ``LIMIT <limit> OFFSET <offset>``
- dameng : ``OFFSET <offset> ROWS FETCH NEXT <limit> ROWS ONLY`` (SQL:2008)
- oracle : **ROWNUM 三层嵌套** —— 兼容 Oracle 11g/10g;
12c+ 也照样能用(row_limiting_clause 不需要时统一走 ROWNUM 简化分支)
- postgres : 同 dameng
- sqlserver: 同 dameng
不支持的 dialect 直接抛 ``ValueError``,让上层早 fail。
注意:``ORDER BY`` 必加,否则 OFFSET 语义在不同实现间不一致
(MySQL/Postgres 默认有顺序,但 Oracle / 达梦没保证)。调用方负责保证
``base_sql`` 已经 ORDER BY。
Oracle 分页模板(11g 兼容):
SELECT * FROM (
SELECT t.*, ROWNUM rn FROM (<base_sql>) t
WHERE ROWNUM <= :end_row
) WHERE rn > :start_row
其中 start_row = offset, end_row = offset + limit
ROWNUM 在内层 WHERE 算(结果集已包含 ORDER BY 顺序),外层只过滤 rn。
注:内层用 ``t.*`` 而非 ``*``,避免 Oracle 对 LONG RAW 列报错;
业务表一般不含这种列,无影响。
"""
dt = db_type.lower()
if dt == "mysql":
return f"{base_sql} LIMIT {int(limit)} OFFSET {int(offset)}"
if dt == "oracle":
offset_i = int(offset)
limit_i = int(limit)
end_row = offset_i + limit_i
# 嵌套子查询:内层执行 ORDER BY,中层 ROWNUM 编号,外层过滤
# ``t.*`` 而非 ``*`` 防 Oracle LONG RAW 列(业务表基本没这列)
return (
f"SELECT * FROM ("
f"SELECT t.*, ROWNUM rn FROM ({base_sql}) t "
f"WHERE ROWNUM <= {end_row}"
f") WHERE rn > {offset_i}"
)
if dt in ("dameng", "postgres", "postgresql", "sqlserver"):
return f"{base_sql} OFFSET {int(offset)} ROWS FETCH NEXT {int(limit)} ROWS ONLY"
raise ValueError(f"Unsupported db_type for pagination: {db_type!r}")
def quote_value(value: Any) -> str: def quote_value(value: Any) -> str:
"""把 Python 值转成 SQL 字面量。仅用于已知安全的字典/枚举值。""" """把 Python 值转成 SQL 字面量。仅用于已知安全的字典/枚举值。"""
if value is None: if value is None:
......
"""查询 session 管理器(web3 自包含版)
为每次 ``POST /api/queries/start`` 在内存里建一个 ``QuerySession``,记录任务配置 +
DB 连接 + 分页参数的快照。前端按页号去 ``POST /api/queries/page`` 拉一页结果,可并发
拉多页,并主动 ``POST /api/queries/cancel`` 终止。
生命周期:
创建 ``/start`` 跑完 COUNT 后建 session 入 ``_REGISTRY``,UUID 作 token 返回
取消 ``/cancel`` 从 ``_REGISTRY`` pop 掉,后续 ``/page`` 立刻 409 短路
TTL 每次 ``/start`` 入口懒清理 ``>30min`` 的 session(无后台线程,简单即正义)
隔离粒度:**per-client** 而非 **per-user**。
当前无 auth 模块 → 任何人拿到 ``session_id`` 都能拉/取消对应 session;
docstring 写明这个 caveat,将来加 auth 时必须把 session_id 绑到 user_id 上。
进程假设:
单进程 uvicorn(无 ``--workers``)→ ``_REGISTRY`` 模块级 dict + ``Lock`` 足够。
改成多 worker / 跨进程时必须换 Redis / 共享 dict(这里先不实现)。
"""
from __future__ import annotations
import threading
import time
from dataclasses import dataclass, field
from typing import Optional
from web3.backend.core.db_adapter import DBConfig
# ── 模块级存储 ────────────────────────────────────────────
# 改 per-user 时记得把 get/pop 都接受 user_id 校验。
_REGISTRY_LOCK = threading.Lock()
_REGISTRY: dict[str, "QuerySession"] = {}
_TTL_SECONDS = 30 * 60
# ── 数据类 ────────────────────────────────────────────────
@dataclass
class CompiledFieldRules:
"""一个字段下的规则快照(被 QuerySession 引用)
为避免 session 把 DB session 持有太久,规则在 ``/start`` 时同步编译好缓存。
字段大小写已归一为小写(与列名对齐)。
"""
field_key: str
# 每条规则:(rule_type, regex_or_None, code_or_None, desc)
rules: list[tuple[str, Optional[str], Optional[str], str]] = field(default_factory=list)
@dataclass
class QuerySession:
"""一次查询的所有上下文。
注意:
``cancelled`` 由 ``/cancel`` 置 True 并从 ``_REGISTRY`` 摘掉;
但 ``/page`` 在并发飞过来时仍可能拿到 in-flight 副本,所以 ``/page`` 入口要
二次检查 ``_REGISTRY`` 里是否还在(不在就 404),以及 session 上的 cancelled。
"""
session_id: str
task_id: int
db_config: DBConfig
db_type: str
page_size: int
total_rows: int
total_pages: int
select_sql_base: str # SELECT cols FROM table ORDER BY first_field
compiled: list[CompiledFieldRules] = field(default_factory=list)
col_names: list[str] = field(default_factory=list)
field_list_out: list[dict] = field(default_factory=list) # 给前端表头用
cancelled: bool = False
pages_scanned: int = 0
bad_rows_total: int = 0
created_at: float = field(default_factory=time.time)
lock: threading.Lock = field(default_factory=threading.Lock)
def mark_page_done(self, bad_delta: int) -> None:
with self.lock:
self.pages_scanned += 1
self.bad_rows_total += bad_delta
# ── CRUD ──────────────────────────────────────────────────
def put(session: QuerySession) -> None:
"""插入 session,入口处顺便清一遍过期 session"""
with _REGISTRY_LOCK:
_purge_expired_locked()
_REGISTRY[session.session_id] = session
def get(session_id: str) -> Optional[QuerySession]:
"""按 session_id 取 session;已被 /cancel 摘掉就返 None"""
with _REGISTRY_LOCK:
return _REGISTRY.get(session_id)
def pop(session_id: str) -> Optional[QuerySession]:
"""原子地取走 session(被 /cancel 用);返回被取走的对象"""
with _REGISTRY_LOCK:
return _REGISTRY.pop(session_id, None)
def size() -> int:
with _REGISTRY_LOCK:
return len(_REGISTRY)
# ── 内部 ──────────────────────────────────────────────────
def _purge_expired_locked() -> None:
"""把超过 TTL 的 session 删掉。调用方必须已持有 _REGISTRY_LOCK。
Lazy cleanup —— 没有后台线程;只在 /start 入口触发一次。
单进程 uvicorn 这个开销可忽略(session 数个位数)。
"""
cutoff = time.time() - _TTL_SECONDS
expired = [sid for sid, s in _REGISTRY.items() if s.created_at < cutoff]
for sid in expired:
_REGISTRY.pop(sid, None)
\ No newline at end of file
"""web3 任务查询 + 校验 API """web3 任务查询 + 校验 API(分页版)
端点: 端点:
POST /api/queries/run 按任务配置:连目标库 → 读表 → 跑规则 → 返回结果 POST /api/queries/start 创建查询 session:COUNT → 编译规则 → 入 session_manager
POST /api/queries/page 按页号拉一页,跑规则,只返回不命中行的(每页独立连接)
POST /api/queries/cancel 摘掉 session(best-effort 中断)
实现: 设计:
- 复用 web3.backend.core.db_adapter 建连 - 复用 web3.backend.core.db_adapter 建连(每页新连接,简单)
- 复用 routers/db._normalize_db_type 做中文 label → 内部小写 - 复用 routers/db._normalize_db_type 做中文 label → 内部小写
- 列名按 task.field_list 显式 SELECT(不 SELECT * 防止暴露未配置列) - 列名按 task.field_list 显式 SELECT(不 SELECT * 防止暴露未配置列)
SELECT 的列 = 「有规则的字段」∪「show_default 的字段」,
后者没有规则也要拉,否则结果表这一列全空
- 列名小写归一(达梦默认大写,MySQL 视 collation)→ 与 field_key 对齐
- 校验引擎:rule_runner.run_rule() 按 rule_type dispatch - 校验引擎:rule_runner.run_rule() 按 rule_type dispatch
- regex : re.compile(regex).search(value),空 regex 回退到 `^.+$`(非空校验) - regex : re.compile(regex).search(value),空 regex 回退到 `^.+$`(非空校验)
- number/date : 用户提供的 def check(value) -> bool,安全沙箱 exec - number/date : 用户提供的 def check(value) -> bool,安全沙箱 exec
...@@ -19,26 +18,37 @@ ...@@ -19,26 +18,37 @@
且所有未通过的规则都会进 issues(备注列列全原因),不止第一条 且所有未通过的规则都会进 issues(备注列列全原因),不止第一条
- 行级:任一字段不合规 → 整行不合规 → 返回;全部字段都合规的行直接丢弃 - 行级:任一字段不合规 → 整行不合规 → 返回;全部字段都合规的行直接丢弃
扫描范围:**全表**,SQL 不带 LIMIT。 分页策略:
- 走 DBConnection.iter_rows 流式逐行读(MySQL 用 SSDictCursor), - SQL 不带 LIMIT,由 paginate_sql(base, db_type, offset, limit) 拼方言
整张表不会进内存;内存只由「命中的不合规行」决定 - ORDER BY 必加,否则 OFFSET 语义不一致(MySQL 有默认顺序,达梦/Oracle 无)
- max_rows 只截断**返回**给前端的行数,不截断扫描: - 每页新开 DBConnection → fetchall(不用 iter_rows,页内是有限结果集)
扫描继续跑完,scanned / bad_rows 始终是全表的真实数字, - 每页 ~50ms 重连开销 vs 长连接 cursor 生命周期管理复杂度,选前者
被截断时 truncated=True,前端据此提示「还有更多,请导出」
Session 隔离:
- per-client 而非 per-user(当前无 auth,详见 session_manager.py 注释)
- UUID 作 token,TTL 30min(懒清理)
""" """
from __future__ import annotations from __future__ import annotations
import re import re
import time import time
import uuid
from typing import Any, Optional from typing import Any, Optional
from fastapi import APIRouter, Depends from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field as PydField from pydantic import BaseModel, Field as PydField
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from web3.backend._logging import get_logger from web3.backend._logging import get_logger
from web3.backend.core.db_adapter import DBConfig, DBConnection, quote_ident from web3.backend.core.db_adapter import DBConfig, DBConnection, paginate_sql, quote_ident
from web3.backend.core.rule_runner import RuleRunError, run_rule from web3.backend.core.rule_runner import RuleRunError, run_rule
from web3.backend.core.session_manager import (
CompiledFieldRules,
QuerySession,
get as sm_get,
pop as sm_pop,
put as sm_put,
)
from web3.backend.db.database import get_session from web3.backend.db.database import get_session
from web3.backend.models.field import Field from web3.backend.models.field import Field
from web3.backend.models.rule import Rule from web3.backend.models.rule import Rule
...@@ -51,44 +61,70 @@ logger = get_logger("backend.routers.queries") ...@@ -51,44 +61,70 @@ logger = get_logger("backend.routers.queries")
# ── Pydantic ────────────────────────────────────────────── # ── Pydantic ──────────────────────────────────────────────
class RunQueryRequest(BaseModel): class StartQueryRequest(BaseModel):
task_id: int = PydField(..., description="任务 id(后端从 DB 读全部配置)") task_id: int = PydField(..., description="任务 id(后端从 DB 读全部配置)")
max_rows: int = PydField( page_size: int = PydField(
5000, ge=1, le=100000, 500, ge=1, le=10000,
description="最多**返回**多少不合规行(不限制扫描范围,扫描始终是全表)", description="每页行数(默认 500);前端并发拉页,受单页大小影响",
) )
class RunQueryRow(BaseModel): class StartQueryResponse(BaseModel):
"""单行结果:业务列 + 错误标记(与前端 ResultTable 行结构对齐)""" ok: bool
errorCells: list[str] = PydField(default_factory=list) message: Optional[str] = None
issues: list[dict] = PydField(default_factory=list) session_id: Optional[str] = None
__row_index: int = 0 # 内部行号(前端不展示,用于后续 P 走「行内编辑」时定位) task_id: int
total_rows: int = 0
total_pages: int = 0
page_size: int = 0
field_list: list[dict] = PydField(default_factory=list, description="顺带回当前任务的字段+规则数")
class PageQueryRequest(BaseModel):
session_id: str = PydField(..., description="/start 返回的 session token")
page_no: int = PydField(..., ge=1, description="1-based 页号")
class RunQueryResponse(BaseModel): class PageQueryResponse(BaseModel):
ok: bool ok: bool
message: Optional[str] = None message: Optional[str] = None
scanned: int = 0 # 全表扫过的行数 session_id: str
bad_rows: int = 0 # 全表里不合规的行数(即使 rows 被截断,这个数字也是完整的) page_no: int
total: int = 0 # = bad_rows page_size: int
truncated: bool = False # bad_rows > max_rows → rows 只是前 max_rows 条 row_start: int # (page_no-1)*page_size + 1
rows: list[dict] = PydField( bad_rows: list[dict] = PydField(default_factory=list)
default_factory=list, scanned_delta: int = 0 # 本页扫过的行数(应 = page_size,最后一页可能更少)
description="只含「不合规」的行(最多 max_rows 条);全部规则都通过的行不返回", bad_delta: int = 0 # 本页命中的不合规行数
) cancelled: bool = False # 当前 session 是否已取消(前端据此停拉剩余页)
field_list: list[dict] = PydField(default_factory=list, description="顺带回当前任务的字段+规则数,让前端能渲染表头") done: bool = False # 本页是最后一页
# ── 端点 ────────────────────────────────────────────────── class CancelQueryRequest(BaseModel):
@router.post("/run", response_model=RunQueryResponse, summary="按任务配置跑查询 + 校验") session_id: str
async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)):
# 1) 拉任务
task = db.query(Task).filter(Task.id == req.task_id).first() class CancelQueryResponse(BaseModel):
ok: bool
cancelled: bool
pages_scanned: int
bad_rows_so_far: int
# ── 公共工具 ──────────────────────────────────────────────
def _load_task_compiled(db: Session, task_id: int):
"""从 ORM 拉 Task + Fields + Rules,并按 plan §3 编译好。
Returns:
(task, fields_with_rules, field_list_out, compiled_fields, col_names)
fields_with_rules: [(Field, [Rule])]
field_list_out: 给 /start 返回的字段摘要
compiled_fields: [CompiledFieldRules](小写 field_key 归一)
col_names: 按 ord 升序的列名(小写,去重)
"""
task = db.query(Task).filter(Task.id == task_id).first()
if not task: if not task:
return RunQueryResponse(ok=False, message=f"任务不存在:{req.task_id}") return None, None, None, None, None
# 2) 拉字段+规则
fields = db.query(Field).filter(Field.task_id == task.id).order_by(Field.ord).all() fields = db.query(Field).filter(Field.task_id == task.id).order_by(Field.ord).all()
rules_by_field: dict[int, list[Rule]] = { rules_by_field: dict[int, list[Rule]] = {
f.id: db.query(Rule).filter(Rule.field_id == f.id).order_by(Rule.ord).all() f.id: db.query(Rule).filter(Rule.field_id == f.id).order_by(Rule.ord).all()
...@@ -98,8 +134,7 @@ async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)): ...@@ -98,8 +134,7 @@ async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)):
(f, rules_by_field[f.id]) for f in fields if rules_by_field[f.id] (f, rules_by_field[f.id]) for f in fields if rules_by_field[f.id]
] ]
# 顺带回一份字段列表(表头用);包含规则明细(rule_list),前端可渲染「规则明细」tooltip # 给 /start 返回的表头摘要(list 走窄通道,详情从 session.field_list_out 拿)
# 不含 field_comment —— 表头注释从数据源补,见下方「打开连接后从 info_schema 拉」
field_list_out = [ field_list_out = [
{ {
"id": f.id, "id": f.id,
...@@ -121,23 +156,41 @@ async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)): ...@@ -121,23 +156,41 @@ async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)):
for f in fields for f in fields
] ]
# 3) 没配数据表 → 早返回 # 编译:regex 提前 compile;number/date 不编译(每行 exec)
if not task.source_table: # 注意:r.rule_type / r.desc / r.regex / r.code 都是 SQLAlchemy Mapped 字段,
return RunQueryResponse( # 运行时是 str,但静态类型是 InstrumentedAttribute —— 用 str() 包一下让 Pylance 也满意
ok=False, compiled_fields: list[CompiledFieldRules] = []
message="任务未配置数据表(source_table 为空)", for f, rules in fields_with_rules:
field_list=field_list_out, snapshots: list[tuple[str, Optional[str], Optional[str], str]] = []
) for r in rules:
if not fields_with_rules: rt = str(r.rule_type or "regex")
return RunQueryResponse( desc = str(r.desc or "")
ok=True, if rt == "regex":
message="任务未配置任何校验规则", regex_src = str(r.regex or "").strip() or r"^.+$"
field_list=field_list_out, try:
) compiled_pat = re.compile(regex_src)
snapshots.append((rt, compiled_pat.pattern, None, desc))
except re.error as e:
logger.warning(f"[queries] 规则 id={r.id} regex 编译失败:{e}")
snapshots.append((rt, None, None, desc)) # None 标记"坏规则"
else:
snapshots.append((rt, None, str(r.code or ""), desc))
compiled_fields.append(CompiledFieldRules(field_key=f.field_key.lower(), rules=snapshots))
# 4) 拼 DBConfig # SELECT 列名:按 ord 升序的去重小写列表
col_names: list[str] = []
seen: set[str] = set()
for f in fields:
if f.field_key and f.field_key not in seen:
col_names.append(f.field_key)
seen.add(f.field_key)
return task, fields_with_rules, field_list_out, compiled_fields, col_names
def _build_db_config(task: Task) -> DBConfig:
conn = _parse_conn_json(task.conn_json) conn = _parse_conn_json(task.conn_json)
cfg = DBConfig( return DBConfig(
db_type=_normalize_db_type(task.db_type or "MySQL"), db_type=_normalize_db_type(task.db_type or "MySQL"),
host=conn.get("host") or "", host=conn.get("host") or "",
port=int(conn.get("port") or 3306), port=int(conn.get("port") or 3306),
...@@ -147,160 +200,216 @@ async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)): ...@@ -147,160 +200,216 @@ async def run_query(req: RunQueryRequest, db: Session = Depends(get_session)):
charset="utf8mb4", charset="utf8mb4",
connect_timeout=10, connect_timeout=10,
# Oracle 专用:编辑时前端回填的 Instant Client 路径(保存进 conn_json) # Oracle 专用:编辑时前端回填的 Instant Client 路径(保存进 conn_json)
# 不传 → _ensure_oracle_client 走默认搜索 PATH/注册表/ORACLE_HOME,
# 找不到就退回 thin 模式 → DPY-3010(参考 db_adapter.py:_ensure_oracle_client 注释)
oracle_client_dir=conn.get("oracleClientDir"), oracle_client_dir=conn.get("oracleClientDir"),
) )
# 5) 校验引擎:按 rule_type 准备执行器
# - regex:提前 re.compile,失败标 None("坏规则") def _build_select_sql(cfg: DBConfig, col_names: list[str], table: str) -> str:
# - number/date:用户函数源码直接传,run_rule 时走沙箱 + validate_user_code 兜底 """拼 base SQL(不含 LIMIT/OFFSET),按首列 ORDER BY 保证 OFFSET 确定性"""
compiled: list[tuple[Field, list[tuple[Rule, Optional[re.Pattern]]]]] = []
for f, rules in fields_with_rules:
rule_patterns: list[tuple[Rule, Optional[re.Pattern]]] = []
for r in rules:
rt = r.rule_type or "regex"
if rt == "regex":
regex_src = (r.regex or "").strip() or r"^.+$" # 空 regex = 非空校验
try:
rule_patterns.append((r, re.compile(regex_src)))
except re.error as e:
logger.warning(f"[queries] 规则 id={r.id} regex 编译失败:{e}")
rule_patterns.append((r, None)) # None 标记"坏规则"
else:
# number/date:编译阶段不校验代码(每行要 exec 一次太慢),
# 真正跑到这行时再让沙箱兜底;编译阶段先标个非 None 占位,
# 让「坏规则」分支只在执行时触发
rule_patterns.append((r, re.compile(r"^.*$"))) # 占位正则,永远 match
compiled.append((f, rule_patterns))
# 6) 连库 + 拉数据
# SELECT 的列 = 任务已配置的**全部**字段(plan §3 决策:结果表展示全部已配置列作为上下文,
# 规则只决定哪些列要标红,不决定哪些列要展示;早期版本限制为「rules ∪ show_default」,
# 但结果表表头用的是 field_list(全部已配置),列与数据错位 → 非规则列一片 NULL。
# 仍不做 SELECT * —— 只 SELECT 任务显式配置的列,防止暴露未配置的敏感列)
started = time.monotonic()
table = task.source_table
schema = (conn.get("schema") or conn.get("db") or "").strip()
col_names: list[str] = []
seen: set[str] = set()
for f in fields:
if f.field_key and f.field_key not in seen:
col_names.append(f.field_key)
seen.add(f.field_key)
col_sql = ", ".join(quote_ident(c, cfg.db_type) for c in col_names) col_sql = ", ".join(quote_ident(c, cfg.db_type) for c in col_names)
# 全表扫描:不带 LIMIT。靠 iter_rows 流式读,整表不进内存 first_col = quote_ident(col_names[0], cfg.db_type)
sql = f"SELECT {col_sql} FROM {quote_ident(table, cfg.db_type)}" return f"SELECT {col_sql} FROM {quote_ident(table, cfg.db_type)} ORDER BY {first_col}"
def _evaluate_row(compiled_fields: list[CompiledFieldRules], raw: dict[str, Any]) -> tuple[list[str], list[dict]]:
"""对一行业务数据跑所有规则,返 (error_cells, issues)。
error_cells: 不合规字段 key 列表(去重保序)
issues: 每条不通过的规则的 {field, desc}
"""
issues: list[dict] = []
error_cells: list[str] = []
for cf in compiled_fields:
val = raw.get(cf.field_key)
for rule_type, compiled_regex, code, desc in cf.rules:
passed: Optional[bool] = None
fail_desc: Optional[str] = None
if rule_type == "regex":
if compiled_regex is None:
passed = False
fail_desc = f"规则正则编译失败:{desc}"
else:
pat = re.compile(compiled_regex)
val_str = "" if val is None else str(val)
passed = pat.search(val_str) is not None
if not passed:
fail_desc = desc
else:
try:
passed = run_rule(rule_type, None, code, val)
if not passed:
fail_desc = desc
except RuleRunError as e:
passed = False
fail_desc = f"规则执行失败:{e}"
if passed is False:
issues.append({"field": cf.field_key, "desc": fail_desc})
if cf.field_key not in error_cells:
error_cells.append(cf.field_key)
return error_cells, issues
# ── 端点 ──────────────────────────────────────────────────
@router.post("/start", response_model=StartQueryResponse, summary="创建查询 session(跑 COUNT + 编译规则)")
async def start_query(req: StartQueryRequest, db: Session = Depends(get_session)):
task, fields_with_rules, field_list_out, compiled_fields, col_names = _load_task_compiled(db, req.task_id)
if task is None:
return StartQueryResponse(ok=False, message=f"任务不存在:{req.task_id}", task_id=req.task_id)
if not task.source_table:
return StartQueryResponse(
ok=False, message="任务未配置数据表(source_table 为空)",
task_id=task.id, field_list=field_list_out,
)
if not fields_with_rules:
return StartQueryResponse(
ok=True, message="任务未配置任何校验规则",
task_id=task.id, field_list=field_list_out,
)
cfg = _build_db_config(task)
started = time.monotonic()
logger.info("─" * 60) logger.info("─" * 60)
logger.info( logger.info(
f"POST /api/queries/run task_id={task.id} table={table!r} " f"POST /api/queries/start task_id={task.id} table={task.source_table!r} "
f"fields={len(col_names)} rules_total={sum(len(rp) for _, rp in compiled)} " f"fields={len(col_names)} rules_total={sum(len(cf.rules) for cf in compiled_fields)} "
f"max_rows={req.max_rows}(全表扫描)" f"page_size={req.page_size}"
) )
logger.info(f"SQL: {sql}")
# 1) COUNT(短连接,干净)
# 7) 流式扫全表 + 逐行校验,只留「不合规」的行
# 字段内多条规则 = AND(任一条不满足 → 该字段不合规,且所有未通过的规则都记进 issues)
# 行内多个字段 = AND(任一字段不合规 → 整行不合规 → 返回)
scanned = 0
bad_count = 0
out_rows: list[dict[str, Any]] = []
try: try:
with DBConnection(cfg) as dbc: with DBConnection(cfg) as dbc:
# 7.0) 从数据源补全字段注释(仅展示用,失败也不阻塞主查询) count_sql = f"SELECT COUNT(*) AS n FROM {quote_ident(task.source_table, cfg.db_type)}"
# 字段的 column_comment 没有落库(plan §4 决策),这里直接从 info_schema 拉一次 count_row = dbc.fetchone(count_sql)
# 失败时 field_comment 留空字符串,前端 tooltip 退化到「暂无注释」 total_rows = int((count_row or {}).get("n") or 0)
if schema and table:
try:
col_meta = dbc.list_columns(schema, table) or []
comments_by_key: dict[str, str] = {
(m.get("column_name") or "").lower(): (m.get("column_comment") or "")
for m in col_meta
if m.get("column_name")
}
for f_out in field_list_out:
c = comments_by_key.get(f_out["field_key"].lower(), "")
if c:
f_out["field_comment"] = c
except Exception as e:
logger.warning(f"[queries] 获取字段注释失败,跳过:{type(e).__name__}: {e}")
# 7.1) 流式扫全表 + 逐行校验
for raw in dbc.iter_rows(sql):
scanned += 1
# 列名归一为小写(达梦默认大写,MySQL 视 collation)
raw_lc = {(k.lower() if isinstance(k, str) else k): v for k, v in raw.items()}
issues: list[dict] = []
error_cells: list[str] = []
for f, rule_patterns in compiled:
val = raw_lc.get(f.field_key.lower())
for rule, pat in rule_patterns:
rt = rule.rule_type or "regex"
rule_passed: Optional[bool] = None
fail_desc: Optional[str] = None
if rt == "regex":
if pat is None:
# 坏规则:正则本身编译不过,按不合规处理并把原因暴露出来
rule_passed = False
fail_desc = f"规则正则编译失败:{rule.desc}"
else:
val_str = "" if val is None else str(val)
rule_passed = pat.search(val_str) is not None
if not rule_passed:
fail_desc = rule.desc
else:
# number / date:走沙箱执行用户函数
try:
rule_passed = run_rule(rt, None, rule.code, val)
if not rule_passed:
fail_desc = rule.desc
except RuleRunError as e:
rule_passed = False
fail_desc = f"规则执行失败:{e}"
if rule_passed is False:
issues.append({"field": f.field_key, "desc": fail_desc})
if f.field_key not in error_cells:
error_cells.append(f.field_key)
if not issues:
continue # 全部规则都通过 → 合规行,不返回
bad_count += 1
# 超出 max_rows 后继续扫(保证 scanned / bad_rows 是全表真值),但不再攒行
if len(out_rows) >= req.max_rows:
continue
out_row = dict(raw_lc)
out_row["errorCells"] = error_cells
out_row["issues"] = issues
out_row["__row_index"] = scanned - 1
out_rows.append(out_row)
except Exception as e: except Exception as e:
msg = f"查询失败: {type(e).__name__}: {e}" msg = f"COUNT 失败:{type(e).__name__}: {e}"
logger.exception(msg) logger.exception(msg)
return RunQueryResponse(ok=False, message=msg, field_list=field_list_out) return StartQueryResponse(ok=False, message=msg, task_id=task.id, field_list=field_list_out)
# 2) 补字段注释(info_schema;不影响主流程)
if (conn := _parse_conn_json(task.conn_json)) and (schema := (conn.get("schema") or conn.get("db") or "").strip()):
try:
with DBConnection(cfg) as dbc:
col_meta = dbc.list_columns(schema, task.source_table) or []
comments_by_key: dict[str, str] = {
(m.get("column_name") or "").lower(): (m.get("column_comment") or "")
for m in col_meta
if m.get("column_name")
}
for f_out in field_list_out:
c = comments_by_key.get(f_out["field_key"].lower(), "")
if c:
f_out["field_comment"] = c
except Exception as e:
logger.warning(f"[queries] 获取字段注释失败,跳过:{type(e).__name__}: {e}")
# 3) 拼 base SQL(ORDER BY 必加,否则 OFFSET 语义不一致)
base_sql = _build_select_sql(cfg, col_names, task.source_table)
total_pages = max(1, (total_rows + req.page_size - 1) // req.page_size)
# 4) 建 session
sess = QuerySession(
session_id=uuid.uuid4().hex,
task_id=task.id,
db_config=cfg,
db_type=cfg.db_type,
page_size=req.page_size,
total_rows=total_rows,
total_pages=total_pages,
select_sql_base=base_sql,
compiled=compiled_fields,
col_names=col_names,
field_list_out=field_list_out,
)
sm_put(sess)
elapsed_ms = (time.monotonic() - started) * 1000 elapsed_ms = (time.monotonic() - started) * 1000
truncated = bad_count > len(out_rows)
logger.info( logger.info(
f"✅ 全表扫描完成:{bad_count}/{scanned} 行不合规" f"✅ session={sess.session_id} total_rows={total_rows} total_pages={total_pages} "
f"{f'(返回前 {len(out_rows)} 条,已截断)' if truncated else ''}"
f"(耗时 {elapsed_ms:.0f}ms)" f"(耗时 {elapsed_ms:.0f}ms)"
) )
logger.info("─" * 60) logger.info("─" * 60)
return StartQueryResponse(
ok=True,
session_id=sess.session_id,
task_id=task.id,
total_rows=total_rows,
total_pages=total_pages,
page_size=req.page_size,
field_list=field_list_out,
)
message = None
if truncated:
message = (
f"全表扫描 {scanned} 行,命中 {bad_count} 行不合规;"
f"当前只返回前 {len(out_rows)} 条,其余请用导出获取"
)
return RunQueryResponse( @router.post("/page", response_model=PageQueryResponse, summary="拉一页")
async def fetch_page(req: PageQueryRequest, db: Session = Depends(get_session)):
sess = sm_get(req.session_id)
if sess is None:
raise HTTPException(status_code=404, detail=f"session 不存在或已过期:{req.session_id}")
if req.page_no > sess.total_pages:
raise HTTPException(status_code=400, detail=f"page_no 越界:{req.page_no} > total_pages={sess.total_pages}")
cfg = sess.db_config
offset = (req.page_no - 1) * sess.page_size
sql = paginate_sql(sess.select_sql_base, cfg.db_type, offset, sess.page_size)
logger.info(
f"[queries] session={sess.session_id[:8]} page={req.page_no}/{sess.total_pages} "
f"offset={offset} size={sess.page_size}"
)
# 每页新连接(避开 cursor 生命周期复杂度,~50ms 重连可接受)
bad_rows: list[dict[str, Any]] = []
scanned_delta = 0
try:
with DBConnection(cfg) as dbc:
rows = dbc.fetchall(sql)
scanned_delta = len(rows)
for idx, raw in enumerate(rows):
err_cells, issues = _evaluate_row(sess.compiled, raw)
if not issues:
continue
row_out = dict(raw)
row_out["errorCells"] = err_cells
row_out["issues"] = issues
# 全局行号:(page_no-1)*page_size + idx → 前端做 row-key(append 模式必用)
row_out["__row_index"] = offset + idx
bad_rows.append(row_out)
except Exception as e:
# 单页拉失败不摘 session —— 让前端重试或走 /cancel
msg = f"拉取 page={req.page_no} 失败:{type(e).__name__}: {e}"
logger.exception(msg)
raise HTTPException(status_code=500, detail=msg)
sess.mark_page_done(bad_delta=len(bad_rows))
return PageQueryResponse(
ok=True, ok=True,
message=message, session_id=sess.session_id,
scanned=scanned, page_no=req.page_no,
total=bad_count, page_size=sess.page_size,
bad_rows=bad_count, row_start=offset + 1,
truncated=truncated, bad_rows=bad_rows,
rows=out_rows, scanned_delta=scanned_delta,
field_list=field_list_out, bad_delta=len(bad_rows),
cancelled=False, # session 已被 cancel 时上面已 404
done=req.page_no == sess.total_pages,
) )
@router.post("/cancel", response_model=CancelQueryResponse, summary="取消查询 session")
async def cancel_query(req: CancelQueryRequest):
sess = sm_pop(req.session_id)
if sess is None:
# 已过期 / 已取消 / 不存在 —— 算 idempotent success
return CancelQueryResponse(ok=True, cancelled=False, pages_scanned=0, bad_rows_so_far=0)
sess.cancelled = True
logger.info(
f"[queries] cancel session={sess.session_id[:8]} "
f"pages_scanned={sess.pages_scanned} bad_rows={sess.bad_rows_total}"
)
return CancelQueryResponse(
ok=True,
cancelled=True,
pages_scanned=sess.pages_scanned,
bad_rows_so_far=sess.bad_rows_total,
)
\ No newline at end of file
"""验证:POST /api/queries/{start,page,cancel} 三端点(2026-08-21 新增)。
端点设计:见 routers/queries.py 顶部 docstring。
覆盖:
- /start COUNT + 建 session + 返 session_id + total_rows + total_pages
- /page 分页 + 行内规则 + __row_index 稳定 + 边界(最后一页/越界)
- /cancel 摘 session + 后续 /page → 404 + 幂等
- paginate_sql: mysql / dameng / oracle / sqlite(unsupported)
- session_manager: TTL 懒清理 + 并发 page 计数
跑:python -m pytest web3/tests/test_queries.py -v
"""
import json
import re
import sys
import threading
import time
from pathlib import Path
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
from web3.backend.core import session_manager as sm_mod # noqa: E402
from web3.backend.core.db_adapter import paginate_sql # noqa: E402
from web3.backend.db.database import Base, get_session # noqa: E402
from web3.backend.models.field import Field # noqa: E402
from web3.backend.models.rule import Rule # noqa: E402
from web3.backend.models.task import Task # noqa: E402
from web3.backend.models.task_group import TaskGroup # noqa: E402
from web3.backend.routers import queries as queries_mod # noqa: E402
# ── 目标表的假数据 ────────────────────────────────────────────
FAKE_TABLE_ROWS = [
{"id_card": "110101199003074512", "phone": "13800138000", "person_name": "张三"}, # 全合规
{"id_card": "1101", "phone": "13800138001", "person_name": "李四"}, # id_card 挂
{"id_card": "110101199003074512", "phone": "021-8888", "person_name": "王五"}, # phone 挂
{"id_card": "abc", "phone": "nope", "person_name": "赵六"}, # 两个都挂
{"id_card": "310101198801011234", "phone": "13900139000", "person_name": "钱七"}, # 全合规
]
class _FakeDBConnection:
"""替身:同时支持 fetchone (COUNT) + fetchall (PAGE) + list_columns。"""
count_n: int = 5 # COUNT(*) 默认返 5
page_rows: list[dict] = FAKE_TABLE_ROWS
column_meta: list[dict] = [
{"column_name": "id_card", "column_comment": "身份证号"},
{"column_name": "phone", "column_comment": "手机号"},
{"column_name": "person_name", "column_comment": "姓名"},
]
seen_sqls: list[str] = []
fail_sql: str | None = None # 匹配则抛 RuntimeError(按需注入)
def __init__(self, cfg):
self.cfg = cfg
def __enter__(self):
return self
def __exit__(self, *exc):
return False
def fetchone(self, sql, params=None):
_FakeDBConnection.seen_sqls.append(sql)
if _FakeDBConnection.fail_sql and _FakeDBConnection.fail_sql in sql:
raise RuntimeError("simulated fetchone failure")
return {"n": _FakeDBConnection.count_n}
def fetchall(self, sql, params=None, batch_size=1000):
_FakeDBConnection.seen_sqls.append(sql)
if _FakeDBConnection.fail_sql and _FakeDBConnection.fail_sql in sql:
raise RuntimeError("simulated fetchall failure")
offset, limit = _parse_limit_offset(sql)
rows = list(_FakeDBConnection.page_rows)
if offset is None or limit is None:
return rows
return rows[offset:offset + limit]
def list_columns(self, schema, table_name):
return list(_FakeDBConnection.column_meta)
def _parse_limit_offset(sql: str) -> tuple[int | None, int | None]:
"""从 SQL 里抽 LIMIT/OFFSET 或 OFFSET/FETCH NEXT。
返回 (offset, limit):
- 找到 → 两个 int
- 没匹配(无分页的 SQL)→ (None, None) → fetchall 返全量
"""
s = sql.upper()
m = re.search(r"LIMIT\s+(\d+)\s+OFFSET\s+(\d+)", s)
if m:
return int(m.group(2)), int(m.group(1))
m = re.search(r"OFFSET\s+(\d+)\s+ROWS\s+FETCH\s+NEXT\s+(\d+)\s+ROWS\s+ONLY", s)
if m:
return int(m.group(1)), int(m.group(2))
return None, None
@pytest.fixture()
def client(tmp_path, monkeypatch):
_FakeDBConnection.count_n = 5
_FakeDBConnection.page_rows = list(FAKE_TABLE_ROWS)
_FakeDBConnection.seen_sqls = []
_FakeDBConnection.fail_sql = None
engine = create_engine(
f"sqlite:///{tmp_path / 'test.sqlite'}",
connect_args={"check_same_thread": False},
)
Base.metadata.create_all(engine)
TestSession = sessionmaker(bind=engine, autocommit=False, autoflush=False)
with TestSession() as db:
db.add(TaskGroup(id=1, name="身份证"))
db.add(Task(
id=1, name="paginate-test", group_id=1,
source_table="t_user_info", db_type="MySQL",
conn_json=json.dumps({"host": "127.0.0.1", "port": 3306,
"user": "u", "password": "p", "db": "d"}),
))
db.flush()
db.add(Field(id=10, task_id=1, field_key="id_card", show_default=1, ord=0))
db.add(Rule(field_id=10, desc="必须 18 位", regex=r"^.{18}$", ord=0))
db.add(Field(id=11, task_id=1, field_key="phone", show_default=1, ord=1))
db.add(Rule(field_id=11, desc="手机号 11 位", regex=r"^1\d{10}$", ord=0))
db.add(Field(id=12, task_id=1, field_key="person_name", show_default=1, ord=2))
db.commit()
monkeypatch.setattr(queries_mod, "DBConnection", _FakeDBConnection)
sm_mod._REGISTRY.clear() # 每个 test 之间隔离
app = FastAPI()
app.include_router(queries_mod.router, prefix="/api")
app.dependency_overrides[get_session] = lambda: TestSession()
with TestClient(app) as c:
yield c
def _start(client, page_size=2) -> dict:
r = client.post("/api/queries/start", json={"task_id": 1, "page_size": page_size})
assert r.status_code == 200, r.text
body = r.json()
assert body["ok"] is True, body
return body
# ── /start ──────────────────────────────────────────────────
def test_start_returns_session_and_total(client):
"""COUNT=5 + page_size=2 → total_pages=3。"""
_FakeDBConnection.count_n = 5
body = _start(client, page_size=2)
assert body["session_id"]
assert body["task_id"] == 1
assert body["total_rows"] == 5
assert body["total_pages"] == 3
assert body["page_size"] == 2
assert len(body["field_list"]) == 3
# COUNT 的 SQL 应是纯 SELECT COUNT(*) + ORDER BY 没有 LIMIT
count_sql = next(s for s in _FakeDBConnection.seen_sqls if "COUNT" in s.upper())
assert "COUNT(*)" in count_sql.upper()
assert "LIMIT" not in count_sql.upper()
def test_start_unknown_task_returns_error(client):
r = client.post("/api/queries/start", json={"task_id": 999})
assert r.status_code == 200
body = r.json()
assert body["ok"] is False
assert "任务不存在" in body["message"]
# ── /page ───────────────────────────────────────────────────
def test_page_basic(client):
"""5 行里 3 行不合规;page=1 (size=2) → 扫 2 行、命中 1 行(李四)。"""
body = _start(client, page_size=2)
r = client.post("/api/queries/page", json={"session_id": body["session_id"], "page_no": 1})
assert r.status_code == 200, r.text
p = r.json()
assert p["ok"] is True
assert p["page_no"] == 1
assert p["page_size"] == 2
assert p["row_start"] == 1
assert p["scanned_delta"] == 2
assert p["bad_delta"] == 1
assert p["done"] is False
assert len(p["bad_rows"]) == 1
assert p["bad_rows"][0]["person_name"] == "李四"
# 全局行号:李四是全表第 2 行(0-based=1)
assert p["bad_rows"][0]["__row_index"] == 1
def test_page_last_partial_marks_done(client):
"""total=5, page_size=2 → 第 3 页只 1 行(钱七合规),done=True。"""
body = _start(client, page_size=2)
r = client.post("/api/queries/page", json={"session_id": body["session_id"], "page_no": 3})
p = r.json()
assert p["ok"] is True
assert p["page_no"] == 3
assert p["done"] is True
assert p["scanned_delta"] == 1
assert p["bad_delta"] == 0
assert p["bad_rows"] == []
def test_page_out_of_range_400(client):
body = _start(client, page_size=2)
r = client.post("/api/queries/page", json={"session_id": body["session_id"], "page_no": 99})
assert r.status_code == 400
assert "越界" in r.json()["detail"]
def test_page_unknown_session_404(client):
r = client.post("/api/queries/page", json={"session_id": "no-such-id", "page_no": 1})
assert r.status_code == 404
assert "不存在" in r.json()["detail"]
def test_row_index_stable_across_pages(client):
"""page=1 和 page=2 的 __row_index 不重叠、全局连续。"""
body = _start(client, page_size=2)
sid = body["session_id"]
r1 = client.post("/api/queries/page", json={"session_id": sid, "page_no": 1}).json()
r2 = client.post("/api/queries/page", json={"session_id": sid, "page_no": 2}).json()
idx1 = [row["__row_index"] for row in r1["bad_rows"]]
idx2 = [row["__row_index"] for row in r2["bad_rows"]]
assert all(0 <= i < 2 for i in idx1), f"page=1 的 row_index 应在 [0,2):{idx1}"
assert all(2 <= i < 4 for i in idx2), f"page=2 的 row_index 应在 [2,4):{idx2}"
assert set(idx1).isdisjoint(set(idx2)), "row_index 不该跨页重复"
# ── /cancel ─────────────────────────────────────────────────
def test_cancel_then_page_404(client):
body = _start(client, page_size=2)
sid = body["session_id"]
cr = client.post("/api/queries/cancel", json={"session_id": sid}).json()
assert cr["ok"] is True and cr["cancelled"] is True
# 取消后 /page → 404(session 已被 pop)
r = client.post("/api/queries/page", json={"session_id": sid, "page_no": 1})
assert r.status_code == 404
assert "不存在" in r.json()["detail"]
def test_cancel_idempotent(client):
"""连续 cancel 两次:第二次 cancelled=False 但 ok=True(不报错)。"""
body = _start(client, page_size=2)
sid = body["session_id"]
r1 = client.post("/api/queries/cancel", json={"session_id": sid}).json()
r2 = client.post("/api/queries/cancel", json={"session_id": sid}).json()
assert r1["cancelled"] is True
assert r2["cancelled"] is False
assert r2["ok"] is True
# ── paginate_sql 助手 ───────────────────────────────────────
def test_paginate_sql_mysql():
sql = paginate_sql("SELECT * FROM t ORDER BY id", "mysql", offset=1000, limit=500)
assert "LIMIT 500 OFFSET 1000" in sql.upper()
def test_paginate_sql_dameng():
sql = paginate_sql("SELECT * FROM t ORDER BY id", "dameng", offset=1000, limit=500)
assert "OFFSET 1000 ROWS FETCH NEXT 500 ROWS ONLY" in sql.upper()
def test_paginate_sql_oracle():
"""Oracle 走 ROWNUM 嵌套子查询(兼容 11g/10g;12c+ 也照样能跑)。
历史原因(2026-08-21):T4 任务(Oracle 11g 服务端)拉页全失败 ORA-00933,
原 OFFSET/FETCH NEXT 是 12c+ 新语法;改用 ROWNUM 三层嵌套。
"""
sql = paginate_sql("SELECT * FROM t ORDER BY id", "oracle", offset=1000, limit=500)
upper = sql.upper()
assert "ROWNUM" in upper, f"Oracle 应走 ROWNUM 嵌套:{sql}"
# end_row = 1000 + 500 = 1500;start_row = 1000
assert "ROWNUM <= 1500" in upper, f"内层 ROWNUM 阈值应是 1500:{sql}"
assert "RN > 1000" in upper, f"外层 rn 过滤应是 > 1000:{sql}"
# 不应该再含 OFFSET/FETCH NEXT(12c+ 语法,老 Oracle 不认)
assert "OFFSET" not in upper and "FETCH NEXT" not in upper, \
f"Oracle 不该用 OFFSET/FETCH NEXT(老服务端不认):{sql}"
def test_paginate_sql_unsupported_db_raises():
with pytest.raises(ValueError, match="Unsupported"):
paginate_sql("SELECT * FROM t", "sqlite", offset=0, limit=10)
# ── 并发 + TTL ──────────────────────────────────────────────
def test_concurrent_pages_count_once(client):
"""两个并发 /page → session.pages_scanned == 2(锁正确)。"""
body = _start(client, page_size=2)
sid = body["session_id"]
results: list[dict] = []
errors: list[Exception] = []
def worker(page_no: int):
try:
r = client.post("/api/queries/page", json={"session_id": sid, "page_no": page_no})
results.append(r.json())
except Exception as e:
errors.append(e)
t1 = threading.Thread(target=worker, args=(1,))
t2 = threading.Thread(target=worker, args=(2,))
t1.start(); t2.start()
t1.join(); t2.join()
assert not errors, f"并发 page 不该报错:{errors}"
assert len(results) == 2
sess = sm_mod.get(sid)
assert sess is not None
assert sess.pages_scanned == 2, f"并发两个 /page,pages_scanned 应是 2:{sess.pages_scanned}"
assert sess.bad_rows_total == sum(len(r["bad_rows"]) for r in results)
def test_ttl_cleanup_on_next_start(client):
"""改 created_at 为 31 分钟前,再 start → 旧 session 应被懒清理。"""
body = _start(client, page_size=2)
old_sid = body["session_id"]
sess = sm_mod.get(old_sid)
assert sess is not None
sess.created_at = time.time() - 31 * 60 # 模拟 31 分钟前
# 再 start → 触发 lazy cleanup → 旧 session 应消失
_start(client, page_size=2)
assert sm_mod.get(old_sid) is None, "TTL 过期 session 应在下次 start 时被清掉"
\ No newline at end of file
...@@ -48,7 +48,19 @@ class _FakeDBConnection: ...@@ -48,7 +48,19 @@ class _FakeDBConnection:
def __exit__(self, *exc): def __exit__(self, *exc):
return False return False
def fetchone(self, sql, params=None):
# 2026-08-21 分页重构:/start 需要 COUNT(*)
_FakeDBConnection.last_sql = sql
return {"n": len(FAKE_TABLE_ROWS)}
def fetchall(self, sql, params=None, batch_size=1000):
# 2026-08-21 分页重构:/page 走 fetchall(LIMIT/OFFSET)
# 这里 mock 简化处理:直接返全量(断言不依赖 LIMIT 切片,单测聚焦规则 dispatch)
_FakeDBConnection.last_sql = sql
return [dict(r) for r in FAKE_TABLE_ROWS]
def iter_rows(self, sql, params=None, batch_size=1000): def iter_rows(self, sql, params=None, batch_size=1000):
# 保留 iter_rows 以兼容老调用方(老测试可能还在用)
_FakeDBConnection.last_sql = sql _FakeDBConnection.last_sql = sql
for r in FAKE_TABLE_ROWS: for r in FAKE_TABLE_ROWS:
yield dict(r) yield dict(r)
...@@ -175,7 +187,7 @@ def test_default_rule_type_is_regex(client): ...@@ -175,7 +187,7 @@ def test_default_rule_type_is_regex(client):
# ── 2) queries/run dispatch ───────────────────────────── # ── 2) queries/run dispatch ─────────────────────────────
def test_queries_run_dispatches_by_rule_type(client): def test_queries_run_dispatches_by_rule_type(client):
"""queries/run 对 number/date 规则走 rule_runner,结果里该字段不合规的行被查出。""" """queries/start+page 对 number/date 规则走 rule_runner,结果里该字段不合规的行被查出。"""
c, _ = client c, _ = client
# 1) 先建一个含 number + date 规则的任务 # 1) 先建一个含 number + date 规则的任务
create = c.post("/api/tasks", json={ create = c.post("/api/tasks", json={
...@@ -217,10 +229,20 @@ def test_queries_run_dispatches_by_rule_type(client): ...@@ -217,10 +229,20 @@ def test_queries_run_dispatches_by_rule_type(client):
assert create.status_code == 200, create.text assert create.status_code == 200, create.text
task_id = create.json()["id"] task_id = create.json()["id"]
# 2) 跑查询 # 2) 跑查询:/start + 聚合 /page(2026-08-21 分页重构,/run 端点已删)
r = c.post("/api/queries/run", json={"task_id": task_id}) start = c.post("/api/queries/start", json={"task_id": task_id, "page_size": 1000}).json()
assert r.status_code == 200, r.text assert start["ok"], start
body = r.json() rows: list[dict] = []
for n in range(1, start["total_pages"] + 1):
page = c.post("/api/queries/page", json={"session_id": start["session_id"], "page_no": n}).json()
assert page["ok"], page
rows.extend(page["bad_rows"])
body = {
"ok": True,
"scanned": start["total_rows"],
"rows": rows,
"field_list": start["field_list"],
}
assert body["ok"] is True assert body["ok"] is True
assert body["scanned"] == 4 assert body["scanned"] == 4
...@@ -266,13 +288,18 @@ def test_queries_run_rule_failure_marks_row_bad_with_reason(client): ...@@ -266,13 +288,18 @@ def test_queries_run_rule_failure_marks_row_bad_with_reason(client):
assert create.status_code == 200, create.text assert create.status_code == 200, create.text
task_id = create.json()["id"] task_id = create.json()["id"]
r = c.post("/api/queries/run", json={"task_id": task_id}) # 2026-08-21 分页重构:/run 已删,改用 /start + /page
body = r.json() start = c.post("/api/queries/start", json={"task_id": task_id, "page_size": 1000}).json()
assert body["scanned"] == 4 assert start["ok"], start
rows: list[dict] = []
for n in range(1, start["total_pages"] + 1):
page = c.post("/api/queries/page", json={"session_id": start["session_id"], "page_no": n}).json()
rows.extend(page["bad_rows"])
assert start["total_rows"] == 4
# 所有行 age 字段都算不合规 → 4 行全返回 # 所有行 age 字段都算不合规 → 4 行全返回
assert len(body["rows"]) == 4 assert len(rows) == 4
# 每行的 issues 都应包含「规则执行失败」字样 # 每行的 issues 都应包含「规则执行失败」字样
for row in body["rows"]: for row in rows:
assert any("规则执行失败" in i["desc"] for i in row["issues"]), row assert any("规则执行失败" in i["desc"] for i in row["issues"]), row
......
""" """验证:POST /api/queries/start + /api/queries/page 只返回「不合规」的行(2026-08-21 重写)。
验证:POST /api/queries/run 只返回「不合规」的行。
本文件早期版本基于 /api/queries/run(一次性同步返回)。分页重构后:
/run 端点已删 → 改为 /start + 轮询 /page 聚合结果
业务行为不变(只返不合规行 + issues + errorCells),断言适配新接口
判定口径(本次需求): 判定口径(与原版一致):
- 字段内多条规则是 AND —— 任一条不满足,该字段就不合规,且每条未通过的规则都进 issues - 字段内多条规则是 AND —— 任一条不满足,该字段就不合规,且每条未通过的规则都进 issues
- 行内多个字段也是 AND —— 任一字段不合规,整行就要被查出来 - 行内多个字段也是 AND —— 任一字段不合规,整行就要被查出来
- 全部规则都通过的行 **不返回** - 全部规则都通过的行 **不返回**
做法:用临时 sqlite 当配置库(task/field/rule),monkeypatch 掉 DBConnection, 做法:临时 sqlite 当配置库,monkeypatch 掉 DBConnection,喂一批构造好的目标表数据。
喂一批构造好的目标表数据,不依赖真实 MySQL / 达梦。 FakeDBConnection 同时支持 fetchone (COUNT) + fetchall (PAGE) + 记录 SQL。
跑:python -m pytest web3/tests/test_run_query_only_bad_rows.py -v 跑:python -m pytest web3/tests/test_run_query_only_bad_rows.py -v
""" """
import json import json
import re
import sys import sys
from pathlib import Path from pathlib import Path
...@@ -23,6 +27,7 @@ from sqlalchemy.orm import sessionmaker ...@@ -23,6 +27,7 @@ from sqlalchemy.orm import sessionmaker
sys.path.insert(0, str(Path(__file__).resolve().parents[2])) sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
from web3.backend.core import session_manager as sm_mod # noqa: E402
from web3.backend.db.database import Base, get_session # noqa: E402 from web3.backend.db.database import Base, get_session # noqa: E402
from web3.backend.models.field import Field # noqa: E402 from web3.backend.models.field import Field # noqa: E402
from web3.backend.models.rule import Rule # noqa: E402 from web3.backend.models.rule import Rule # noqa: E402
...@@ -31,8 +36,9 @@ from web3.backend.models.task_group import TaskGroup # noqa: E402 ...@@ -31,8 +36,9 @@ from web3.backend.models.task_group import TaskGroup # noqa: E402
from web3.backend.routers import queries as queries_mod # noqa: E402 from web3.backend.routers import queries as queries_mod # noqa: E402
# ── 目标表的假数据 ─────────────────────────────────────────── # ── 目标表的假数据 ────────────────────────────────────────────
# 字段:id_card(有 2 条规则)/ phone(有 1 条规则)/ person_name(无规则但 show_default) # 字段:id_card(有 2 条规则)/ phone(有 1 条规则)/ person_name(无规则但 show_default)
# nickname(无规则 + show_default=false)
FAKE_TABLE_ROWS = [ FAKE_TABLE_ROWS = [
# 0 全合规 → 不该被返回 # 0 全合规 → 不该被返回
{"id_card": "110101199003074512", "phone": "13800138000", "person_name": "张三"}, {"id_card": "110101199003074512", "phone": "13800138000", "person_name": "张三"},
...@@ -48,17 +54,15 @@ FAKE_TABLE_ROWS = [ ...@@ -48,17 +54,15 @@ FAKE_TABLE_ROWS = [
class _FakeDBConnection: class _FakeDBConnection:
"""替身:吞掉 DBConfig,iter_rows 流式吐 rows(默认 FAKE_TABLE_ROWS)。""" """替身:吞 DBConfig;支持 fetchone (COUNT) + fetchall (PAGE) + list_columns + 记 SQL。
last_sql = None 注意:不要在类外给 fetchall/fetchone 赋 lambda,因为 closure 引用会跨测试残留。
rows = FAKE_TABLE_ROWS 需要自定义行集时改 _FakeDBConnection.override_page_rows = [...] 后由 fetchall 读取。
column_meta = [ # 假装数据源 info_schema 返回的元数据 """
{"column_name": "id_card", "column_comment": "身份证号"},
{"column_name": "phone", "column_comment": "手机号"}, seen_sqls: list[str] = []
{"column_name": "person_name", "column_comment": "姓名"}, count_n: int = 5 # COUNT(*) 默认 5
{"column_name": "nickname", "column_comment": "昵称"}, override_page_rows: list[dict] | None = None # 非 None 时覆盖 fetchall 返回
]
list_columns_called_with = None # 用来断言「打开连接后去拉过字段注释」
def __init__(self, cfg): def __init__(self, cfg):
self.cfg = cfg self.cfg = cfg
...@@ -69,23 +73,33 @@ class _FakeDBConnection: ...@@ -69,23 +73,33 @@ class _FakeDBConnection:
def __exit__(self, *exc): def __exit__(self, *exc):
return False return False
def iter_rows(self, sql, params=None, batch_size=1000): def fetchone(self, sql, params=None):
_FakeDBConnection.last_sql = sql _FakeDBConnection.seen_sqls.append(sql)
for r in _FakeDBConnection.rows: return {"n": _FakeDBConnection.count_n}
yield dict(r)
def fetchall(self, sql, params=None, batch_size=1000):
_FakeDBConnection.seen_sqls.append(sql)
if _FakeDBConnection.override_page_rows is not None:
return list(_FakeDBConnection.override_page_rows)
# 不模拟 LIMIT/OFFSET 切片,单测聚焦规则 dispatch(断言不依赖切片)
return [dict(r) for r in FAKE_TABLE_ROWS]
def list_columns(self, schema, table_name): def list_columns(self, schema, table_name):
"""新接口:从 info_schema 拉字段元数据(含 column_comment) return [
queries.py 在打开连接后会调一次,把 column_comment 写进 field_list {"column_name": "id_card", "column_comment": "身份证号"},
""" {"column_name": "phone", "column_comment": "手机号"},
_FakeDBConnection.list_columns_called_with = (schema, table_name) {"column_name": "person_name", "column_comment": "姓名"},
return list(_FakeDBConnection.column_meta) {"column_name": "nickname", "column_comment": "昵称"},
]
@pytest.fixture() @pytest.fixture()
def client(tmp_path, monkeypatch): def client(tmp_path, monkeypatch):
_FakeDBConnection.rows = FAKE_TABLE_ROWS # 每个用例复位,避免相互污染 # 重置类级别 mock 状态,避免前面 test 残留
_FakeDBConnection.last_sql = None _FakeDBConnection.seen_sqls = []
_FakeDBConnection.count_n = 5
_FakeDBConnection.override_page_rows = None
engine = create_engine( engine = create_engine(
f"sqlite:///{tmp_path / 'test.sqlite'}", f"sqlite:///{tmp_path / 'test.sqlite'}",
connect_args={"check_same_thread": False}, connect_args={"check_same_thread": False},
...@@ -117,6 +131,7 @@ def client(tmp_path, monkeypatch): ...@@ -117,6 +131,7 @@ def client(tmp_path, monkeypatch):
db.commit() db.commit()
monkeypatch.setattr(queries_mod, "DBConnection", _FakeDBConnection) monkeypatch.setattr(queries_mod, "DBConnection", _FakeDBConnection)
sm_mod._REGISTRY.clear() # 每个用例隔离
app = FastAPI() app = FastAPI()
app.include_router(queries_mod.router, prefix="/api") app.include_router(queries_mod.router, prefix="/api")
...@@ -125,13 +140,26 @@ def client(tmp_path, monkeypatch): ...@@ -125,13 +140,26 @@ def client(tmp_path, monkeypatch):
yield c yield c
def _run(client, **payload): def _run(client, **payload) -> dict:
body_in = {"task_id": 1, **payload} """等价旧 /run:/start + 聚合 /page bad_rows,构造旧响应形状以兼容旧断言。"""
r = client.post("/api/queries/run", json=body_in) start = client.post("/api/queries/start", json={"task_id": 1, "page_size": 1000}).json()
assert r.status_code == 200, r.text assert start["ok"], start.get("message")
body = r.json() bad_rows: list[dict] = []
assert body["ok"] is True, body.get("message") for n in range(1, start["total_pages"] + 1):
return body page = client.post("/api/queries/page", json={
"session_id": start["session_id"], "page_no": n,
}).json()
assert page["ok"], page
bad_rows.extend(page["bad_rows"])
return {
"ok": True,
"scanned": start["total_rows"],
"bad_rows": len(bad_rows),
"total": len(bad_rows),
"truncated": False,
"rows": bad_rows,
"field_list": start["field_list"],
}
def test_only_bad_rows_returned(client): def test_only_bad_rows_returned(client):
...@@ -175,84 +203,57 @@ def test_any_field_failing_flags_the_row(client): ...@@ -175,84 +203,57 @@ def test_any_field_failing_flags_the_row(client):
def test_show_default_field_without_rules_is_selected(client): def test_show_default_field_without_rules_is_selected(client):
"""没规则但勾了「默认展示」的字段也要 SELECT 出来,否则结果表这列全空。""" """没规则但勾了「默认展示」的字段也要 SELECT 出来,否则结果表这列全空。"""
_run(client) _run(client)
sql = _FakeDBConnection.last_sql sql = next((s for s in _FakeDBConnection.seen_sqls if s.upper().startswith("SELECT") and "LIMIT" in s.upper()), "")
assert "`person_name`" in sql, f"show_default 字段应进 SELECT:{sql}" assert "`person_name`" in sql, f"show_default 字段应进 SELECT:{sql}"
assert "`id_card`" in sql and "`phone`" in sql assert "`id_card`" in sql and "`phone`" in sql
def test_full_table_scan_no_limit_clause(client): def test_page_sql_contains_order_by_for_stable_offset(client):
"""全表扫描:SQL 不该带 LIMIT。""" """分页:base SQL 必须 ORDER BY,不然 OFFSET 语义在不同方言间不一致。
_run(client)
sql = _FakeDBConnection.last_sql
assert "LIMIT" not in sql.upper(), f"应全表扫描,SQL 不该有 LIMIT:{sql}"
def test_max_rows_truncates_returned_rows_but_not_counts(client):
"""max_rows 只截断返回的行,不截断扫描 —— scanned / bad_rows 仍是全表真值。"""
body = _run(client, max_rows=2)
assert len(body["rows"]) == 2, "返回行数应被 max_rows 截断"
assert body["truncated"] is True
assert body["scanned"] == 5, "扫描不受 max_rows 影响,仍是全表 5 行"
assert body["bad_rows"] == 3, "计数不受 max_rows 影响,仍是全表 3 行不合规"
assert body["total"] == 3
assert body["message"], "被截断时应给出提示文案"
def test_not_truncated_when_under_max_rows(client):
"""没超上限时 truncated=False,且不带截断提示。"""
body = _run(client, max_rows=5000)
assert body["truncated"] is False
assert body["message"] is None
assert len(body["rows"]) == body["bad_rows"] == 3
注:旧版本断言『SQL 不该带 LIMIT』(全表扫描),现在分页必须带 LIMIT/OFFSET,
改成断言 OFFSET 稳定性的前置条件 ORDER BY。
"""
_run(client)
page_sql = next(s for s in _FakeDBConnection.seen_sqls if "LIMIT" in s.upper() or "OFFSET" in s.upper())
assert "ORDER BY" in page_sql.upper(), f"分页 SQL 必须 ORDER BY:{page_sql}"
def test_scans_beyond_old_1000_row_limit(client):
"""回归:旧实现固定 LIMIT 1000,只校验前 1000 行。现在要扫完整张表。"""
good = {"id_card": "110101199003074512", "phone": "13800138000", "person_name": "合规"}
bad = {"id_card": "x", "phone": "y", "person_name": "第2500行的脏数据"}
_FakeDBConnection.rows = [dict(good) for _ in range(2499)] + [bad]
body = _run(client) def test_page_sql_uses_dialect_limit(client):
assert body["scanned"] == 2500, "应扫完 2500 行,不是停在 1000" """MySQL 用 LIMIT <n> OFFSET <m> 语法。"""
assert body["bad_rows"] == 1 _run(client)
assert body["rows"][0]["person_name"] == "第2500行的脏数据", "第 1000 行之后的脏数据也要被查出来" page_sql = next(s for s in _FakeDBConnection.seen_sqls if "LIMIT" in s.upper() or "OFFSET" in s.upper())
assert re.search(r"LIMIT\s+\d+\s+OFFSET\s+\d+", page_sql, re.IGNORECASE) or \
re.search(r"OFFSET\s+\d+\s+ROWS\s+FETCH\s+NEXT\s+\d+\s+ROWS\s+ONLY", page_sql, re.IGNORECASE), \
f"分页 SQL 缺少方言 LIMIT/OFFSET:{page_sql}"
def test_row_index_points_at_position_in_full_scan(client): def test_row_index_points_at_position_in_full_scan(client):
"""__row_index 记的是「全表第几行」,截断后也不能错位。""" """__row_index 记的是「全表第几行」,跨页后也不能错位。"""
body = _run(client, max_rows=1) body = _run(client)
# 李四是全表第 2 行(0-based = 1) # 李四是全表第 2 行(0-based = 1)
assert body["rows"][0]["__row_index"] == 1 li_si = next(r for r in body["rows"] if r["person_name"] == "李四")
assert li_si["__row_index"] == 1, f"李四应是全表第 2 行,实际 __row_index={li_si['__row_index']}"
# ── 2026-08-21 追加:结果表非规则列显示实际值 / 表头字段注释 ───────── # ── 2026-08-21 追加:结果表非规则列显示实际值 / 表头字段注释 ─────────
def test_show_default_false_no_rules_field_is_selected(client): def test_show_default_false_no_rules_field_is_selected(client):
"""没规则 + show_default=false 的字段也要进 SELECT(2026-08-21 行为变更)。 """没规则 + show_default=false 的字段也要进 SELECT。
早期版本限制「rules ∪ show_default」,结果是结果表表头有这列、值全 NULL。 SELECT 全部已配置字段;规则只决定哪些列要标红。
现在改成:SELECT 全部已配置字段;规则只决定哪些列要标红。
""" """
_run(client) _run(client)
sql = _FakeDBConnection.last_sql sql = next((s for s in _FakeDBConnection.seen_sqls if s.upper().startswith("SELECT") and "LIMIT" in s.upper()), "")
assert "`nickname`" in sql, ( assert "`nickname`" in sql, (
f"show_default=false + 无规则的字段也要 SELECT,否则结果表这列全 NULL:{sql}" f"show_default=false + 无规则的字段也要 SELECT,否则结果表这列全 NULL:{sql}"
) )
def test_field_comment_returned_for_header_tooltip(client): def test_field_comment_returned_for_header_tooltip(client):
"""field_list 每项应带回 field_comment(来自数据源 COLUMN_COMMENT), """field_list 每项应带回 field_comment(来自数据源 COLUMN_COMMENT)。"""
供前端表头 tooltip 展示字段注释。
"""
_FakeDBConnection.list_columns_called_with = None # 复位
body = _run(client) body = _run(client)
# 1) 后端确实打开连接去拉过字段注释
assert _FakeDBConnection.list_columns_called_with == ("d", "t_user_info"), (
"应打开连接后从 info_schema 拉一次字段注释"
)
# 2) field_list 里每个有注释的字段都填上了 field_comment
fl = {f["field_key"]: f for f in body["field_list"]} fl = {f["field_key"]: f for f in body["field_list"]}
assert fl["id_card"]["field_comment"] == "身份证号" assert fl["id_card"]["field_comment"] == "身份证号"
assert fl["phone"]["field_comment"] == "手机号" assert fl["phone"]["field_comment"] == "手机号"
...@@ -261,9 +262,7 @@ def test_field_comment_returned_for_header_tooltip(client): ...@@ -261,9 +262,7 @@ def test_field_comment_returned_for_header_tooltip(client):
def test_rule_list_included_for_tooltip(client): def test_rule_list_included_for_tooltip(client):
"""field_list 每项应带回 rule_list(rule 的 desc + regex), """field_list 每项应带回 rule_list(rule 的 desc + regex)。"""
供前端「⚠ N规则」badge hover 出规则明细。
"""
body = _run(client) body = _run(client)
fl = {f["field_key"]: f for f in body["field_list"]} fl = {f["field_key"]: f for f in body["field_list"]}
...@@ -280,4 +279,4 @@ def test_rule_list_included_for_tooltip(client): ...@@ -280,4 +279,4 @@ def test_rule_list_included_for_tooltip(client):
# 无规则的字段:rule_list 是空 list(不是 None,前端可以放心 .map) # 无规则的字段:rule_list 是空 list(不是 None,前端可以放心 .map)
assert fl["person_name"]["rule_list"] == [] assert fl["person_name"]["rule_list"] == []
assert fl["nickname"]["rule_list"] == [] assert fl["nickname"]["rule_list"] == []
\ No newline at end of file
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