deerflow-code/offline-backend-20260512/backend/tests/test_ensure_orm_columns_sync.py
2026-09-07 18:24:55 +08:00

329 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""回归:``_ensure_orm_columns_sync`` 必须把**既有**表升级到最新 ORM schema。
模拟生产「已上线 MySQL → 切到带 Phase 3 改动的代码」的真实场景:``roundtable_jobs``
表已存在但只有旧列(无 attempt/version/active_dedupe_key/…,无唯一约束)。
``create_all`` 不会动既有表,alembic 自动升级在部署里是注释掉的,唯一兜底就是
``_ensure_orm_columns_sync``。若它漏补一列或漏建唯一索引:
* ``attempt``/``version`` 缺失 → 每次 ``claim_job``/条件 UPDATE 引用该列 → SQL 报错 → 作业永远不被领取(并发报错);
* ``active_dedupe_key`` 唯一性缺失 → 同 scope 两条活跃作业共存 → 去重失效(用户报的并发双跑)。
本测试用 SQLite(CREATE UNIQUE INDEX 全平台支持)锁死这条升级路径。
"""
from __future__ import annotations
import pytest
from sqlalchemy import create_engine, inspect, text
def _legacy_roundtable_jobs_ddl() -> str:
"""线上旧版 roundtable_jobs(Phase 3 之前)的最小列集。"""
return (
"CREATE TABLE roundtable_jobs ("
"id VARCHAR(64) PRIMARY KEY,"
"draft_id VARCHAR(64),"
"user_id VARCHAR(64),"
"task_id VARCHAR(128),"
"status VARCHAR(32) NOT NULL DEFAULT 'queued',"
"created_at DATETIME NOT NULL,"
"updated_at DATETIME NOT NULL"
")"
)
def test_legacy_roundtable_jobs_upgraded_with_phase3_columns_and_unique_index():
"""既有表缺 Phase 3 列 → sync 补齐列 + 建唯一索引 + 旧行默认值 0。"""
from deerflow.persistence.engine import _ensure_orm_columns_sync
engine = create_engine("sqlite:///:memory:")
with engine.begin() as conn:
conn.execute(text(_legacy_roundtable_jobs_ddl()))
# 一条旧行(Phase 3 列不存在,验证补列后默认值生效)。
conn.execute(
text(
"INSERT INTO roundtable_jobs (id, draft_id, status, created_at, updated_at) "
"VALUES ('legacy-1', 'd1', 'done', '2026-01-01 00:00:00', '2026-01-01 00:00:00')"
)
)
# 触发升级。
_ensure_orm_columns_sync(conn)
inspector = inspect(engine)
cols = {c["name"] for c in inspector.get_columns("roundtable_jobs")}
# 1) Phase 3 必备列全部补齐(attempt/version 曾缺 server_default 会被漏掉)。
for required in (
"attempt",
"version",
"active_dedupe_key",
"request_id",
"request_hash",
"input_snapshot",
"lease_owner",
"lease_until",
"draft_persist_pending",
):
assert required in cols, f"missing Phase 3 column after sync: {required}"
# 2) 唯一索引建好(去重正确性的根)。
idx_names = {ix["name"] for ix in inspector.get_indexes("roundtable_jobs")}
assert "uq_roundtable_jobs_active_dedupe_key" in idx_names
with engine.begin() as conn:
# 3) 旧行的 NOT NULL 计数/CAS 列拿到 server_default 0(不是 NULL,否则查询报错)。
row = conn.execute(
text("SELECT attempt, version, draft_persist_pending FROM roundtable_jobs WHERE id='legacy-1'")
).fetchone()
assert row is not None
assert int(row[0]) == 0 and int(row[1]) == 0 and int(row[2]) in (0, False)
# 4) 唯一索引真的强制唯一:两条同 active_dedupe_key 的活跃行第二条必冲突。
conn.execute(
text(
"INSERT INTO roundtable_jobs (id, status, active_dedupe_key, created_at, updated_at) "
"VALUES ('a', 'running', 'K', '2026-01-01 00:00:00', '2026-01-01 00:00:00')"
)
)
with pytest.raises(Exception):
conn.execute(
text(
"INSERT INTO roundtable_jobs (id, status, active_dedupe_key, created_at, updated_at) "
"VALUES ('b', 'running', 'K', '2026-01-01 00:00:00', '2026-01-01 00:00:00')"
)
)
# 5) 多个 NULL 不冲突(终态历史作业共存)。
conn.execute(
text(
"INSERT INTO roundtable_jobs (id, status, active_dedupe_key, created_at, updated_at) "
"VALUES ('c', 'done', NULL, '2026-01-01 00:00:00', '2026-01-01 00:00:00')"
)
)
conn.execute(
text(
"INSERT INTO roundtable_jobs (id, status, active_dedupe_key, created_at, updated_at) "
"VALUES ('d', 'done', NULL, '2026-01-01 00:00:00', '2026-01-01 00:00:00')"
)
)
def test_sync_idempotent_repeated_runs_do_not_duplicate_index():
"""重复启动(多次 sync)不应重复建索引(同名已存在即跳过)。"""
from deerflow.persistence.engine import _ensure_orm_columns_sync
engine = create_engine("sqlite:///:memory:")
with engine.begin() as conn:
conn.execute(text(_legacy_roundtable_jobs_ddl()))
_ensure_orm_columns_sync(conn)
# 第二次:不应抛「索引已存在」。
_ensure_orm_columns_sync(conn)
inspector = inspect(engine)
idx_names = [ix["name"] for ix in inspector.get_indexes("roundtable_jobs")]
assert idx_names.count("uq_roundtable_jobs_active_dedupe_key") == 1
def _legacy_task_buttons_ddl() -> str:
"""线上旧版 task_buttons(login_params_json / append_auth 之前)的最小列集。"""
return (
"CREATE TABLE task_buttons ("
"id VARCHAR(128) PRIMARY KEY,"
"business VARCHAR(32) NOT NULL,"
"label VARCHAR(255),"
"link_type VARCHAR(16),"
"target VARCHAR(1024),"
"append_task_id BOOLEAN,"
"enabled BOOLEAN,"
"sort_order INTEGER,"
"updated_by VARCHAR(64),"
"updated_at DATETIME"
")"
)
def test_legacy_task_buttons_gains_login_params_json():
"""既有 task_buttons 缺 login_params_json → sync 必须补上,不能因 NOT NULL 跳过。"""
from deerflow.persistence.engine import _ensure_orm_columns_sync
engine = create_engine("sqlite:///:memory:")
with engine.begin() as conn:
conn.execute(text(_legacy_task_buttons_ddl()))
conn.execute(
text(
"INSERT INTO task_buttons (id, business, label, link_type, target, append_task_id, enabled, sort_order) "
"VALUES ('rwfx-next', 'rwfx', '下一步', 'url', 'https://example.test', 1, 1, 0)"
)
)
_ensure_orm_columns_sync(conn)
cols = {c["name"] for c in inspect(engine).get_columns("task_buttons")}
assert "login_params_json" in cols, "login_params_json was skipped by auto-add"
assert "append_auth" in cols
assert "auth_token_param" in cols
assert "auth_name_param" in cols
with engine.begin() as conn:
row = conn.execute(text("SELECT login_params_json, append_auth FROM task_buttons WHERE id='rwfx-next'")).fetchone()
assert row is not None
# nullable JSON 补列后旧行为 NULL;append_auth 有 server_default 0。
assert row[0] in (None, "[]", "null", b"null")
assert int(row[1]) in (0, False)
def _legacy_workflow_data_sources_ddl() -> str:
"""工作流数据源表在 HTTP 接口元数据列出现前的最小完整结构。"""
return (
"CREATE TABLE workflow_data_sources ("
"id VARCHAR(64) PRIMARY KEY,"
"name VARCHAR(191) NOT NULL,"
"description VARCHAR(512) NOT NULL DEFAULT '',"
"kind VARCHAR(32) NOT NULL DEFAULT 'sql',"
"driver VARCHAR(64) NOT NULL DEFAULT '',"
"secret_ciphertext TEXT NOT NULL DEFAULT '',"
"secret_encrypted BOOLEAN NOT NULL DEFAULT 1,"
"masked_target VARCHAR(512) NOT NULL DEFAULT '',"
"max_rows INTEGER NOT NULL DEFAULT 1000,"
"allowed_tables TEXT,"
"enabled BOOLEAN NOT NULL DEFAULT 1,"
"owner_id VARCHAR(64),"
"created_by VARCHAR(64),"
"created_at DATETIME NOT NULL,"
"updated_at DATETIME NOT NULL"
")"
)
def test_legacy_workflow_data_source_gains_http_metadata_columns():
"""既有数据源表在启动同步后可保存 HTTP 地址和方法白名单。"""
from deerflow.persistence.engine import _ensure_orm_columns_sync
engine = create_engine("sqlite:///:memory:")
with engine.begin() as conn:
conn.execute(text(_legacy_workflow_data_sources_ddl()))
conn.execute(
text(
"INSERT INTO workflow_data_sources (id, name, created_at, updated_at) "
"VALUES ('legacy-http', 'legacy', '2026-01-01 00:00:00', '2026-01-01 00:00:00')"
)
)
_ensure_orm_columns_sync(conn)
columns = {column["name"] for column in inspect(engine).get_columns("workflow_data_sources")}
assert "base_url" in columns
assert "allowed_methods" in columns
with engine.begin() as conn:
row = conn.execute(
text("SELECT base_url, allowed_methods FROM workflow_data_sources WHERE id='legacy-http'")
).fetchone()
assert row is not None
assert row[0] == ""
assert row[1] is None
def _legacy_report_collaboration_003_ddl() -> list[str]:
"""线上若已按 RC-BE-003 建表、尚未跑 015/016 补丁时的最小列集。"""
return [
(
"CREATE TABLE report_collaboration_runs ("
"id VARCHAR(64) PRIMARY KEY,"
"session_id VARCHAR(64) NOT NULL,"
"plan_id VARCHAR(64) NOT NULL,"
"status VARCHAR(32) NOT NULL DEFAULT 'queued',"
"cancel_requested BOOLEAN NOT NULL DEFAULT 0,"
"next_event_seq INTEGER NOT NULL DEFAULT 0,"
"current_phase VARCHAR(64),"
"active_dedupe_key VARCHAR(64),"
"idempotency_key VARCHAR(128),"
"lease_owner VARCHAR(64),"
"lease_until DATETIME,"
"attempt INTEGER NOT NULL DEFAULT 0,"
"version INTEGER NOT NULL DEFAULT 0,"
"created_at DATETIME NOT NULL,"
"updated_at DATETIME NOT NULL"
")"
),
(
"CREATE TABLE report_collaboration_report_versions ("
"id VARCHAR(64) PRIMARY KEY,"
"session_id VARCHAR(64) NOT NULL,"
"version INTEGER NOT NULL,"
"parent_version_id VARCHAR(64),"
"change_type VARCHAR(32) NOT NULL DEFAULT 'initial',"
"change_summary TEXT,"
"markdown TEXT NOT NULL,"
"source_index_json TEXT NOT NULL,"
"is_head BOOLEAN NOT NULL DEFAULT 0,"
"created_at DATETIME NOT NULL"
")"
),
]
def test_legacy_report_collaboration_gains_015_016_columns_and_tables():
"""既有 003 表缺模板/预算/幂等列 → 启动 sync 补列;create_all 建改写/审计新表。"""
from deerflow.persistence.base import Base
from deerflow.persistence.engine import _ensure_orm_columns_sync
import deerflow.persistence.models # noqa: F401
engine = create_engine("sqlite:///:memory:")
with engine.begin() as conn:
for ddl in _legacy_report_collaboration_003_ddl():
conn.execute(text(ddl))
conn.execute(
text(
"INSERT INTO report_collaboration_runs "
"(id, session_id, plan_id, status, created_at, updated_at) "
"VALUES ('run-old', 's1', 'p1', 'queued', '2026-09-05 00:00:00', '2026-09-05 00:00:00')"
)
)
conn.execute(
text(
"INSERT INTO report_collaboration_report_versions "
"(id, session_id, version, markdown, source_index_json, created_at) "
"VALUES ('ver-old', 's1', 1, '# 旧稿', '[]', '2026-09-05 00:00:00')"
)
)
Base.metadata.create_all(conn)
_ensure_orm_columns_sync(conn)
inspector = inspect(engine)
run_cols = {item["name"] for item in inspector.get_columns("report_collaboration_runs")}
version_cols = {item["name"] for item in inspector.get_columns("report_collaboration_report_versions")}
tables = set(inspector.get_table_names())
for required in ("template_snapshot_json", "budget_snapshot_json", "usage_json", "deadline_at", "budget_exhausted_reason"):
assert required in run_cols, f"missing run column after sync: {required}"
for required in ("idempotency_key", "origin_run_id", "quality_result_json", "template_snapshot_json"):
assert required in version_cols, f"missing version column after sync: {required}"
assert "report_collaboration_rewrite_proposals" in tables
assert "report_collaboration_audits" in tables
idx_names = {item["name"] for item in inspector.get_indexes("report_collaboration_report_versions")}
assert "uq_rc_report_versions_idempotency" in idx_names
assert "uq_rc_report_versions_session_version" in idx_names
with engine.begin() as conn:
row = conn.execute(
text("SELECT template_snapshot_json, usage_json, budget_exhausted_reason FROM report_collaboration_runs WHERE id='run-old'")
).fetchone()
assert row is not None
assert row[0] is None and row[2] is None
conn.execute(
text(
"UPDATE report_collaboration_runs SET template_snapshot_json='{}', budget_snapshot_json='{}', usage_json='{}' WHERE id='run-old'"
)
)
conn.execute(
text(
"INSERT INTO report_collaboration_rewrite_proposals "
"(id, session_id, parent_version_id, instruction, markdown, source_index_json, idempotency_key, created_at, updated_at) "
"VALUES ('rw-1', 's1', 'ver-old', 'more-cautious', '# rewrite', '[]', 'ik-rw', '2026-09-06 00:00:00', '2026-09-06 00:00:00')"
)
)
conn.execute(
text(
"INSERT INTO report_collaboration_audits "
"(id, session_id, category, action, created_at) "
"VALUES ('aud-1', 's1', 'policy', 'source.added', '2026-09-06 00:00:00')"
)
)