862 lines
37 KiB
Python
862 lines
37 KiB
Python
#!/usr/bin/env python3
|
||
"""升格命令锁接线的纯离线测试:锁失败时禁止任何业务调用。"""
|
||
|
||
import contextlib
|
||
import copy
|
||
import pathlib
|
||
import sys
|
||
import unittest
|
||
from unittest.mock import Mock, patch
|
||
|
||
from click.testing import CliRunner
|
||
|
||
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
|
||
SKILLS_DIR = PROJECT_ROOT / ".agent" / "skills"
|
||
SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "entity" / "skills" / "ingest" / "重置作品抽取结果" / "scripts"
|
||
BACKUP_SCRIPTS = PROJECT_ROOT / "muse" / "content" / "entity" / "skills" / "ingest" / "备份作品抽取结果" / "scripts"
|
||
EXTRACTION_SCRIPTS = PROJECT_ROOT / "muse" / "content" / "entity" / "skills" / "ingest" / "抽取作品知识" / "scripts"
|
||
sys.path.insert(0, str(SCRIPT_DIR))
|
||
sys.path.insert(0, str(BACKUP_SCRIPTS))
|
||
sys.path.insert(0, str(EXTRACTION_SCRIPTS))
|
||
|
||
import backup_upgrade_work as backup # noqa: E402
|
||
import reset_upgrade_work as reset # noqa: E402
|
||
import upgrade as parse # noqa: E402
|
||
from upgrade_work_lock import UpgradeWorkLockUnavailable # noqa: E402
|
||
|
||
BACKUP_ID = "11111111-1111-4111-8111-111111111111"
|
||
CONFIRMATION_SHA = "c" * 64
|
||
BACKUP_ARGS = [
|
||
"--backup-dir", "/private/tmp",
|
||
"--backup-id", BACKUP_ID,
|
||
"--confirmation-sha", CONFIRMATION_SHA,
|
||
]
|
||
EXPECTED_RESET_LOCK_TABLES = (
|
||
"muse_content_chapter",
|
||
"muse_content_block",
|
||
"muse_meta_schema",
|
||
"muse_meta_schema_version",
|
||
"muse_knowledge_draft",
|
||
"example_upgrade_window",
|
||
"example_upgrade_alias",
|
||
"example_upgrade_presence",
|
||
"example_upgrade_card_state",
|
||
"example_upgrade_audit",
|
||
"example_knowledge_embedding",
|
||
)
|
||
|
||
|
||
def _locked_tables(sql):
|
||
"""从规范化 LOCK TABLE SQL 中提取有序表名。"""
|
||
|
||
prefix = "lock table "
|
||
suffix = " in share row exclusive mode"
|
||
assert sql.startswith(prefix) and sql.endswith(suffix)
|
||
return tuple(table.strip() for table in sql[len(prefix):-len(suffix)].split(","))
|
||
|
||
|
||
def _lock_failure(*_args, **_kwargs):
|
||
"""模拟同书已有命令持锁。"""
|
||
|
||
raise UpgradeWorkLockUnavailable(tenant_id=1, work_id=8)
|
||
|
||
|
||
class _Result:
|
||
"""模拟 psycopg 查询结果,提供 reset 与备份快照使用的读取接口。"""
|
||
|
||
def __init__(self, row=None, rows=None):
|
||
self.row = row
|
||
self.rows = list(rows or [])
|
||
|
||
def fetchone(self):
|
||
return self.row
|
||
|
||
def fetchall(self):
|
||
return list(self.rows)
|
||
|
||
|
||
class _CursorContext:
|
||
"""模拟 dict_row cursor;同一 SQL 在默认 connection 上仍返回 tuple。"""
|
||
|
||
def __init__(self, conn, row_factory):
|
||
self.conn = conn
|
||
self.row_factory = row_factory
|
||
|
||
def __enter__(self):
|
||
return self
|
||
|
||
def __exit__(self, exc_type, exc, traceback):
|
||
return False
|
||
|
||
def execute(self, sql, params=None):
|
||
normalized = " ".join(sql.split()).lower()
|
||
self.conn.executions.append((normalized, params))
|
||
return self.conn._input_query_result(normalized, mapping=True)
|
||
|
||
|
||
class _ResetConnection:
|
||
"""有状态 fake DB:验证向量 SQL 边界,并保留不应被 reset 误伤的反例。"""
|
||
|
||
def __init__(self, *, postcheck_failure=None, confirmed_draft_id=None,
|
||
input_drift=None):
|
||
self.postcheck_failure = postcheck_failure
|
||
self.input_drift = input_drift
|
||
self.commits = 0
|
||
self.executions = []
|
||
self.cursor_row_factories = []
|
||
self.drafts = {
|
||
# 目标作品本轮活卡,以及上次 reset 已软删但向量仍活跃的旧卡。
|
||
101: {"tenant": 1, "work": 8, "source": "upgrade_book", "deleted": False},
|
||
102: {"tenant": 1, "work": 8, "source": "upgrade_book", "deleted": True},
|
||
# 三个反例:其它作品、其它来源、其它租户均不得被本次 reset 误伤。
|
||
201: {"tenant": 1, "work": 9, "source": "upgrade_book", "deleted": False},
|
||
202: {"tenant": 1, "work": 8, "source": "parse_book", "deleted": False},
|
||
203: {"tenant": 2, "work": 8, "source": "upgrade_book", "deleted": False},
|
||
}
|
||
self.embeddings = {
|
||
draft_id: {
|
||
"deleted": False,
|
||
"updater": "",
|
||
"entity_id": 9001 if draft_id == confirmed_draft_id else None,
|
||
}
|
||
for draft_id in self.drafts
|
||
}
|
||
self.window_statuses = ["done", "pending"]
|
||
self.alias_count = 3
|
||
self.presence_count = 4
|
||
self.card_state_count = 2
|
||
self.audit_count = 5
|
||
|
||
def __enter__(self):
|
||
# 进入连接事务时保留快照,异常退出须像 psycopg 一样回滚全部已执行写入。
|
||
self._transaction_snapshot = (
|
||
copy.deepcopy(self.drafts),
|
||
copy.deepcopy(self.embeddings),
|
||
list(self.window_statuses),
|
||
self.alias_count,
|
||
self.presence_count,
|
||
self.card_state_count,
|
||
self.audit_count,
|
||
)
|
||
return self
|
||
|
||
def __exit__(self, exc_type, exc, traceback):
|
||
if exc_type is not None:
|
||
(self.drafts, self.embeddings, self.window_statuses,
|
||
self.alias_count, self.presence_count,
|
||
self.card_state_count, self.audit_count) = self._transaction_snapshot
|
||
return False
|
||
|
||
def commit(self):
|
||
self.commits += 1
|
||
|
||
def cursor(self, *, row_factory=None):
|
||
"""只允许实现代码显式请求 dict_row 快照游标。"""
|
||
|
||
self.cursor_row_factories.append(row_factory)
|
||
return _CursorContext(self, row_factory)
|
||
|
||
def _input_query_result(self, normalized, *, mapping):
|
||
"""为完整输入查询提供真实列形;默认 connection 故意返回 tuple。"""
|
||
|
||
if normalized.startswith("select current_database() as database"):
|
||
row = {
|
||
"database": "other-db" if self.input_drift == "database" else "muse-example",
|
||
"server_addr": "100.64.0.8",
|
||
"server_port": 5433,
|
||
"server_version": "160000",
|
||
}
|
||
return _Result(row=row if mapping else tuple(row.values()))
|
||
if normalized.startswith("select c.id as chapter_id"):
|
||
rows = [
|
||
{
|
||
"chapter_id": 11,
|
||
"chapter_order": 1,
|
||
"chapter_title": "第一章",
|
||
"chapter_revision": 1,
|
||
"block_id": 101,
|
||
"block_order": 1,
|
||
"block_type": "scene",
|
||
"block_title": None,
|
||
"content_doc": None,
|
||
"content_text": "漂移正文" if self.input_drift == "canonical" else "第一章正文",
|
||
"block_revision": 1,
|
||
},
|
||
{
|
||
"chapter_id": 12,
|
||
"chapter_order": 2,
|
||
"chapter_title": "第二章",
|
||
"chapter_revision": 1,
|
||
"block_id": 102,
|
||
"block_order": 1,
|
||
"block_type": "scene",
|
||
"block_title": None,
|
||
"content_doc": None,
|
||
"content_text": "第二章正文",
|
||
"block_revision": 1,
|
||
},
|
||
]
|
||
return _Result(rows=rows if mapping else [tuple(row.values()) for row in rows])
|
||
if normalized.startswith("select s.schema_key, s.active_version_id"):
|
||
rows = [
|
||
{
|
||
"schema_key": schema_key,
|
||
"active_version_id": index,
|
||
"field_contract_snapshot": {"字段": []},
|
||
}
|
||
for index, schema_key in enumerate(sorted(backup.UPGRADE_SCHEMA_KEYS), start=1)
|
||
]
|
||
return _Result(rows=rows if mapping else [tuple(row.values()) for row in rows])
|
||
raise AssertionError(f"未覆盖的完整输入 SQL:{normalized}")
|
||
|
||
@staticmethod
|
||
def _is_target_draft(draft):
|
||
return (draft["tenant"], draft["work"], draft["source"]) == (1, 8, "upgrade_book")
|
||
|
||
def _active_target_vectors(self, *, unconfirmed_only=False):
|
||
matched = []
|
||
for draft_id, embedding in self.embeddings.items():
|
||
draft = self.drafts[draft_id]
|
||
if embedding["deleted"] or not self._is_target_draft(draft):
|
||
continue
|
||
if unconfirmed_only and embedding["entity_id"] is not None:
|
||
continue
|
||
matched.append(draft_id)
|
||
return matched
|
||
|
||
@staticmethod
|
||
def _assert_vector_boundary(sql, params, *, require_active=True, require_unconfirmed=False):
|
||
"""SQL 必须同时约束向量租户、卡租户、作品和来源。"""
|
||
|
||
assert "e.tenant_id=%s" in sql
|
||
if require_active:
|
||
assert "e.deleted=false" in sql
|
||
assert "d.id=e.draft_id" in sql
|
||
assert "d.tenant_id=%s" in sql
|
||
assert "d.work_id=%s" in sql
|
||
assert "d.source_type=%s" in sql
|
||
assert params == (1, 1, 8, "upgrade_book")
|
||
if require_unconfirmed:
|
||
assert "e.entity_id is null" in sql
|
||
|
||
def execute(self, sql, params=None):
|
||
normalized = " ".join(sql.split()).lower()
|
||
self.executions.append((normalized, params))
|
||
if normalized.startswith("lock table"):
|
||
assert _locked_tables(normalized) == EXPECTED_RESET_LOCK_TABLES
|
||
return _Result()
|
||
if normalized.startswith((
|
||
"select current_database() as database",
|
||
"select c.id as chapter_id",
|
||
"select s.schema_key, s.active_version_id",
|
||
)):
|
||
return self._input_query_result(normalized, mapping=False)
|
||
if normalized.startswith("select title from muse_content_work"):
|
||
return _Result(("离线测试书",))
|
||
if normalized.startswith("select count(*) from muse_knowledge_draft"):
|
||
count = sum(
|
||
not draft["deleted"] and self._is_target_draft(draft)
|
||
for draft in self.drafts.values()
|
||
)
|
||
return _Result((count,))
|
||
if (normalized.startswith("select count(*), count(*) filter")
|
||
or normalized.startswith("select count(*) filter")):
|
||
done = sum(status == "done" for status in self.window_statuses)
|
||
pending = sum(status == "pending" for status in self.window_statuses)
|
||
if "where status='done'" in normalized:
|
||
return _Result((len(self.window_statuses), done))
|
||
return _Result((pending, len(self.window_statuses)))
|
||
if normalized.startswith("select count(*) from example_upgrade_alias"):
|
||
return _Result((self.alias_count,))
|
||
if normalized.startswith("select count(*) from example_upgrade_presence"):
|
||
return _Result((self.presence_count,))
|
||
if normalized.startswith("select count(*) from example_upgrade_card_state"):
|
||
return _Result((self.card_state_count,))
|
||
if normalized.startswith("select count(*) from example_upgrade_audit"):
|
||
return _Result((self.audit_count,))
|
||
if normalized.startswith("select count(*) from example_knowledge_embedding e"):
|
||
if "e.entity_id is not null" in normalized:
|
||
self._assert_vector_boundary(normalized, params, require_active=False)
|
||
count = sum(
|
||
embedding["entity_id"] is not None
|
||
for draft_id, embedding in self.embeddings.items()
|
||
if self._is_target_draft(self.drafts[draft_id])
|
||
)
|
||
return _Result((count,))
|
||
self._assert_vector_boundary(normalized, params)
|
||
unconfirmed_only = "e.entity_id is null" in normalized
|
||
count = len(self._active_target_vectors(unconfirmed_only=unconfirmed_only))
|
||
return _Result((count,))
|
||
if normalized.startswith("update muse_knowledge_draft"):
|
||
for draft_id, draft in self.drafts.items():
|
||
if self._is_target_draft(draft) and not draft["deleted"]:
|
||
if self.postcheck_failure == "cards" and draft_id == 101:
|
||
continue
|
||
draft["deleted"] = True
|
||
return _Result(None)
|
||
if normalized.startswith("update example_knowledge_embedding e"):
|
||
self._assert_vector_boundary(normalized, params, require_unconfirmed=True)
|
||
assert "set deleted=true, updater='upgrade-reset'" in normalized
|
||
targets = self._active_target_vectors(unconfirmed_only=True)
|
||
if self.postcheck_failure == "vectors":
|
||
targets = targets[:-1]
|
||
for draft_id in targets:
|
||
self.embeddings[draft_id].update(deleted=True, updater="upgrade-reset")
|
||
return _Result(None)
|
||
if normalized.startswith("update example_upgrade_window"):
|
||
self.window_statuses = ["pending"] * len(self.window_statuses)
|
||
if self.postcheck_failure == "windows":
|
||
self.window_statuses[-1] = "done"
|
||
return _Result(None)
|
||
if normalized.startswith("delete from example_upgrade_audit"):
|
||
self.audit_count = 1 if self.postcheck_failure == "audit" else 0
|
||
return _Result(None)
|
||
if normalized.startswith("delete from example_upgrade_alias"):
|
||
self.alias_count = 1 if self.postcheck_failure == "alias" else 0
|
||
return _Result(None)
|
||
if normalized.startswith("delete from example_upgrade_presence"):
|
||
self.presence_count = 1 if self.postcheck_failure == "presence" else 0
|
||
return _Result(None)
|
||
if normalized.startswith("delete from example_upgrade_card_state"):
|
||
self.card_state_count = 1 if self.postcheck_failure == "state" else 0
|
||
return _Result(None)
|
||
raise AssertionError(f"未覆盖的离线 SQL:{normalized}")
|
||
|
||
|
||
def _snapshot_domains(*, entity_id=None):
|
||
"""构造七域最小一致快照;摘要全部调用备份模块既有规则。"""
|
||
|
||
return {
|
||
"drafts": [
|
||
{"id": 101, "tenant_id": 1, "work_id": 8,
|
||
"source_type": "upgrade_book", "deleted": False},
|
||
{"id": 102, "tenant_id": 1, "work_id": 8,
|
||
"source_type": "upgrade_book", "deleted": True},
|
||
],
|
||
"windows": [
|
||
{"id": 301, "tenant_id": 1, "work_id": 8,
|
||
"window_no": 1, "from_chapter": 1, "to_chapter": 2,
|
||
"status": "done", "deleted": False},
|
||
],
|
||
"aliases": [{"id": 401, "tenant_id": 1, "work_id": 8}],
|
||
"presence": [{"id": 501, "tenant_id": 1, "work_id": 8}],
|
||
"card_state": [{"draft_id": 101, "tenant_id": 1, "work_id": 8}],
|
||
"audits": [{"id": 601, "tenant_id": 1, "draft_id": 101}],
|
||
"embeddings": [
|
||
{"id": 701, "tenant_id": 1, "draft_id": 101,
|
||
"entity_id": entity_id, "deleted": False},
|
||
],
|
||
}
|
||
|
||
|
||
def _code_files():
|
||
"""固定当前代码身份;reset execute 测试只改变待验证的单个 fileSha。"""
|
||
|
||
return {
|
||
"backup_upgrade_work.py": "b" * 64,
|
||
"reset_upgrade_work.py": "r" * 64,
|
||
"parse_upgrade.py": "u" * 64,
|
||
"parse_llm.py": "p" * 64,
|
||
"embed_drafts.py": "e" * 64,
|
||
"upgrade_work_lock.py": "k" * 64,
|
||
"llm.py": "l" * 64,
|
||
"parse-book/SKILL.md": "s" * 64,
|
||
}
|
||
|
||
|
||
def _input_snapshot(*, code_files=None):
|
||
"""构造与真实备份 manifest 同形的完整 reset 输入身份。"""
|
||
|
||
return {
|
||
"databaseIdentity": {
|
||
"database": "muse-example",
|
||
"serverAddr": "100.64.0.8",
|
||
"serverPort": 5433,
|
||
"serverVersion": "160000",
|
||
},
|
||
"canonicalContent": {
|
||
"chapterCount": 2,
|
||
"blockCount": 2,
|
||
"sha": "c" * 64,
|
||
},
|
||
"windowBoundary": {"count": 1, "sha": "w" * 64},
|
||
"activeFieldContracts": {
|
||
"count": 7,
|
||
"schemaKeys": sorted(backup.UPGRADE_SCHEMA_KEYS),
|
||
"sha": "f" * 64,
|
||
},
|
||
"codeFiles": dict(code_files or _code_files()),
|
||
"expectedChapters": 2,
|
||
"expectedWindows": 1,
|
||
"plan": {
|
||
"model": backup.PLANNED_MODEL,
|
||
"semanticDedup": backup.PLANNED_SEMANTIC_DEDUP,
|
||
},
|
||
"sourceType": backup.SOURCE_TYPE,
|
||
"tenant": 1,
|
||
"work": 8,
|
||
}
|
||
|
||
|
||
def _manifest_for(domains, *, code_files=None):
|
||
"""按备份模块的排序、主键摘要和内容摘要构造已离线 verify 的清单替身。"""
|
||
|
||
domain_manifest = {}
|
||
normalized = {}
|
||
for name, spec in backup.DOMAIN_SPECS.items():
|
||
rows = backup.sort_rows(name, domains[name])
|
||
normalized[name] = rows
|
||
domain_manifest[name] = {
|
||
"rowCount": len(rows),
|
||
"primaryKeySha": backup._primary_key_sha(spec, rows),
|
||
"contentSha": backup._content_sha(rows),
|
||
}
|
||
input_snapshot = _input_snapshot(code_files=code_files)
|
||
return {
|
||
"backup_id": BACKUP_ID,
|
||
"confirmationSha": CONFIRMATION_SHA,
|
||
"work": 8,
|
||
"tenant": 1,
|
||
"input": input_snapshot,
|
||
"inputSha": backup._sha256(
|
||
backup.canonical_json(input_snapshot).encode("utf-8")
|
||
),
|
||
"domains": domain_manifest,
|
||
}, normalized
|
||
|
||
|
||
@contextlib.contextmanager
|
||
def _lock_success(*_args, **_kwargs):
|
||
"""离线成功锁,不连接 PostgreSQL。"""
|
||
|
||
yield
|
||
|
||
|
||
class UpgradeCommandLockOfflineTest(unittest.TestCase):
|
||
"""验证锁位于全部业务 SQL、写入和模型调用之前。"""
|
||
|
||
def setUp(self):
|
||
self.runner = CliRunner()
|
||
|
||
def test_parse_run_lock_failure_skips_business(self):
|
||
with patch.object(parse, "upgrade_work_lock", side_effect=_lock_failure) as lock, \
|
||
patch.object(parse, "_run") as business:
|
||
result = self.runner.invoke(parse.cli, ["run", "--work-id", "8"])
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("已有升格命令正在处理", result.output)
|
||
lock.assert_called_once_with(parse.DSN, parse.TENANT, 8)
|
||
business.assert_not_called()
|
||
|
||
def test_parse_windows_lock_failure_skips_database(self):
|
||
with patch.object(parse, "upgrade_work_lock", side_effect=_lock_failure) as lock, \
|
||
patch.object(parse, "connect") as connect:
|
||
result = self.runner.invoke(parse.cli, ["windows", "--work-id", "8"])
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("已有升格命令正在处理", result.output)
|
||
lock.assert_called_once_with(parse.DSN, parse.TENANT, 8)
|
||
connect.assert_not_called()
|
||
|
||
def test_parse_status_does_not_request_lock(self):
|
||
lock = Mock(side_effect=AssertionError("status 不应请求锁"))
|
||
with patch.object(parse, "upgrade_work_lock", lock), \
|
||
patch.object(parse, "connect", side_effect=RuntimeError("离线中止")):
|
||
result = self.runner.invoke(parse.cli, ["status", "--work-id", "8"])
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
lock.assert_not_called()
|
||
|
||
def test_reset_preview_lock_failure_skips_database(self):
|
||
self._assert_reset_lock_failure_skips_database(execute=False)
|
||
|
||
def test_reset_execute_lock_failure_skips_database(self):
|
||
self._assert_reset_lock_failure_skips_database(execute=True)
|
||
|
||
def _assert_reset_lock_failure_skips_database(self, *, execute):
|
||
args = ["--work-id", "8"] + (["--execute", *BACKUP_ARGS] if execute else [])
|
||
with patch.object(reset, "upgrade_work_lock", side_effect=_lock_failure) as lock, \
|
||
patch.object(reset, "connect") as connect:
|
||
result = self.runner.invoke(reset.main, args)
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("已有升格命令正在处理", result.output)
|
||
lock.assert_called_once_with(reset.DSN, reset.TENANT, 8)
|
||
connect.assert_not_called()
|
||
|
||
def test_lock_success_wraps_parse_run_business(self):
|
||
events = []
|
||
|
||
@contextlib.contextmanager
|
||
def fake_lock(*_args, **_kwargs):
|
||
events.append("lock-enter")
|
||
try:
|
||
yield
|
||
finally:
|
||
events.append("lock-exit")
|
||
|
||
def fake_business(*_args, **_kwargs):
|
||
events.append("business")
|
||
|
||
with patch.object(parse, "upgrade_work_lock", fake_lock), \
|
||
patch.object(parse, "_run", side_effect=fake_business):
|
||
result = self.runner.invoke(parse.cli, ["run", "--work-id", "8"])
|
||
|
||
self.assertEqual(result.exit_code, 0, result.output)
|
||
self.assertEqual(events, ["lock-enter", "business", "lock-exit"])
|
||
|
||
def _invoke_reset_execute(self, conn, *, manifest=None, current_domains=None,
|
||
backup_id=BACKUP_ID, confirmation_sha=CONFIRMATION_SHA,
|
||
lock_context=_lock_success, current_code_files=None,
|
||
current_input=None, capture_input=None):
|
||
"""在锁、磁盘 verify 与同事务快照均为离线替身时调用 execute。"""
|
||
|
||
default_manifest, default_domains = _manifest_for(_snapshot_domains())
|
||
manifest = manifest or default_manifest
|
||
current_domains = current_domains or default_domains
|
||
current_input = manifest["input"] if current_input is None else current_input
|
||
capture_input = capture_input or Mock(return_value=current_input)
|
||
args = [
|
||
"--work-id", "8", "--execute",
|
||
"--backup-dir", "/private/tmp",
|
||
"--backup-id", backup_id,
|
||
"--confirmation-sha", confirmation_sha,
|
||
]
|
||
with patch.object(reset, "upgrade_work_lock", lock_context), \
|
||
patch.object(reset, "connect", return_value=conn), \
|
||
patch.object(backup, "verify_backup", return_value=manifest), \
|
||
patch.object(backup, "capture_code_identity", return_value={
|
||
"codeFiles": dict(current_code_files or manifest["input"]["codeFiles"]),
|
||
}), \
|
||
patch.object(backup, "capture_input_snapshot", capture_input), \
|
||
patch.object(backup, "_read_snapshot",
|
||
return_value=("离线测试书", current_domains)):
|
||
return self.runner.invoke(reset.main, args)
|
||
|
||
def test_reset_execute_requires_complete_backup_confirmation(self):
|
||
"""execute 缺少任一备份确认参数时必须在取锁和连接数据库前失败。"""
|
||
|
||
required = {
|
||
"--backup-dir": "/private/tmp",
|
||
"--backup-id": BACKUP_ID,
|
||
"--confirmation-sha": CONFIRMATION_SHA,
|
||
}
|
||
for missing in required:
|
||
args = ["--work-id", "8", "--execute"]
|
||
for option, value in required.items():
|
||
if option != missing:
|
||
args.extend((option, value))
|
||
with self.subTest(missing=missing), \
|
||
patch.object(reset, "upgrade_work_lock") as lock, \
|
||
patch.object(reset, "connect") as connect:
|
||
result = self.runner.invoke(reset.main, args)
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn(missing, result.output)
|
||
lock.assert_not_called()
|
||
connect.assert_not_called()
|
||
|
||
def test_reset_preview_reports_all_target_active_vectors_without_writes(self):
|
||
"""预览包含旧 reset 遗留向量,但不得执行任何 UPDATE/DELETE。"""
|
||
|
||
conn = _ResetConnection()
|
||
with patch.object(reset, "upgrade_work_lock", _lock_success), \
|
||
patch.object(reset, "connect", return_value=conn):
|
||
result = self.runner.invoke(reset.main, ["--work-id", "8"])
|
||
|
||
self.assertEqual(result.exit_code, 0, result.output)
|
||
self.assertIn("《离线测试书》", result.output)
|
||
self.assertIn("软删活向量 2", result.output)
|
||
self.assertFalse(any(sql.startswith(("update ", "delete "))
|
||
for sql, _params in conn.executions))
|
||
self.assertEqual(conn.commits, 0)
|
||
self.assertEqual(conn.cursor_row_factories, [])
|
||
|
||
def test_reset_execute_soft_deletes_only_target_upgrade_vectors(self):
|
||
"""执行只软删当前租户、作品、升格来源向量,并用明确 actor 留痕。"""
|
||
|
||
conn = _ResetConnection()
|
||
result = self._invoke_reset_execute(conn)
|
||
|
||
self.assertEqual(result.exit_code, 0, result.output)
|
||
self.assertIn("七域清理断言全部通过", result.output)
|
||
self.assertEqual(conn.commits, 1)
|
||
self.assertTrue(conn.executions[0][0].startswith("lock table"))
|
||
self.assertEqual(len(conn.cursor_row_factories), 1)
|
||
for draft_id in (101, 102):
|
||
self.assertTrue(conn.embeddings[draft_id]["deleted"])
|
||
self.assertEqual(conn.embeddings[draft_id]["updater"], "upgrade-reset")
|
||
for draft_id in (201, 202, 203):
|
||
self.assertFalse(conn.embeddings[draft_id]["deleted"])
|
||
self.assertEqual(conn.embeddings[draft_id]["updater"], "")
|
||
vector_writes = [sql for sql, _params in conn.executions
|
||
if "example_knowledge_embedding" in sql
|
||
and sql.startswith(("update ", "delete "))]
|
||
self.assertEqual(len(vector_writes), 1)
|
||
self.assertTrue(vector_writes[0].startswith("update "))
|
||
self.assertIn("e.entity_id is null", vector_writes[0])
|
||
|
||
def test_reset_execute_verifies_backup_and_input_inside_advisory_lock(self):
|
||
"""磁盘 verify 与完整输入读取均须位于同书 advisory lock 内。"""
|
||
|
||
events = []
|
||
|
||
@contextlib.contextmanager
|
||
def tracked_lock(*_args, **_kwargs):
|
||
events.append("lock-enter")
|
||
try:
|
||
yield
|
||
finally:
|
||
events.append("lock-exit")
|
||
|
||
manifest, domains = _manifest_for(_snapshot_domains())
|
||
conn = _ResetConnection()
|
||
|
||
def verified_manifest(_path):
|
||
self.assertEqual(events, ["lock-enter"])
|
||
events.append("backup-verified")
|
||
return manifest
|
||
|
||
def captured_input(*_args, **_kwargs):
|
||
self.assertEqual(events, ["lock-enter", "backup-verified"])
|
||
events.append("input-captured")
|
||
return manifest["input"]
|
||
|
||
with patch.object(reset, "upgrade_work_lock", tracked_lock), \
|
||
patch.object(reset, "connect", return_value=conn), \
|
||
patch.object(backup, "verify_backup", side_effect=verified_manifest), \
|
||
patch.object(backup, "capture_code_identity", return_value={
|
||
"codeFiles": _code_files(),
|
||
}), \
|
||
patch.object(backup, "capture_input_snapshot",
|
||
side_effect=captured_input), \
|
||
patch.object(backup, "_read_snapshot",
|
||
return_value=("离线测试书", domains)):
|
||
result = self.runner.invoke(
|
||
reset.main, ["--work-id", "8", "--execute", *BACKUP_ARGS]
|
||
)
|
||
|
||
self.assertEqual(result.exit_code, 0, result.output)
|
||
self.assertEqual(
|
||
events, ["lock-enter", "backup-verified", "input-captured", "lock-exit"]
|
||
)
|
||
|
||
def test_reset_execute_rejects_snapshot_drift_before_writes(self):
|
||
"""七域任一摘要漂移必须在所有 reset 写语句前失败。"""
|
||
|
||
manifest, domains = _manifest_for(_snapshot_domains())
|
||
drifted = copy.deepcopy(domains)
|
||
drifted["aliases"].append({"id": 999, "tenant_id": 1, "work_id": 8})
|
||
conn = _ResetConnection()
|
||
result = self._invoke_reset_execute(
|
||
conn, manifest=manifest, current_domains=drifted,
|
||
)
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("aliases", result.output)
|
||
self.assertEqual(conn.commits, 0)
|
||
self.assertEqual(len(conn.executions), 1)
|
||
self.assertTrue(conn.executions[0][0].startswith("lock table"))
|
||
|
||
def test_reset_execute_captures_complete_input_inside_table_lock(self):
|
||
"""输入源与七域须按固定顺序先锁定,再读完整输入和执行 destructive SQL。"""
|
||
|
||
manifest, domains = _manifest_for(_snapshot_domains())
|
||
current_identity = {"codeFiles": dict(manifest["input"]["codeFiles"])}
|
||
conn = _ResetConnection()
|
||
|
||
def capture_input(*args, **kwargs):
|
||
self.assertTrue(conn.executions)
|
||
self.assertTrue(conn.executions[0][0].startswith("lock table"))
|
||
self.assertEqual(
|
||
_locked_tables(conn.executions[0][0]), EXPECTED_RESET_LOCK_TABLES
|
||
)
|
||
self.assertFalse(any(sql.startswith(("update ", "delete "))
|
||
for sql, _params in conn.executions))
|
||
capture_cursor = args[0]
|
||
self.assertIsInstance(capture_cursor, _CursorContext)
|
||
self.assertIs(capture_cursor.conn, conn)
|
||
self.assertIs(capture_cursor.row_factory, reset.dict_row)
|
||
self.assertEqual(args[1:], (8, 1, current_identity, 2, 1))
|
||
self.assertEqual(kwargs, {"windows": domains["windows"]})
|
||
return manifest["input"]
|
||
|
||
capture = Mock(side_effect=capture_input)
|
||
result = self._invoke_reset_execute(
|
||
conn,
|
||
manifest=manifest,
|
||
current_domains=domains,
|
||
current_code_files=current_identity["codeFiles"],
|
||
capture_input=capture,
|
||
)
|
||
|
||
self.assertEqual(result.exit_code, 0, result.output)
|
||
capture.assert_called_once()
|
||
self.assertEqual(conn.commits, 1)
|
||
|
||
def test_reset_execute_real_input_capture_uses_dict_cursor_and_guards_drift(self):
|
||
"""默认 connection 为 tuple 时,真实 input capture 仍须成功并保持写前漂移闸门。"""
|
||
|
||
domains = backup.sort_rows("windows", _snapshot_domains()["windows"])
|
||
all_domains = _snapshot_domains()
|
||
all_domains["windows"] = domains
|
||
manifest, normalized_domains = _manifest_for(all_domains)
|
||
identity = {"codeFiles": dict(manifest["input"]["codeFiles"])}
|
||
baseline_conn = _ResetConnection()
|
||
with baseline_conn.cursor(row_factory=reset.dict_row) as cursor:
|
||
expected_input = backup.capture_input_snapshot(
|
||
cursor, 8, 1, identity, 2, 1, windows=domains
|
||
)
|
||
manifest["input"] = expected_input
|
||
manifest["inputSha"] = backup._sha256(
|
||
backup.canonical_json(expected_input).encode("utf-8")
|
||
)
|
||
|
||
for case, input_drift, expected_exit in (
|
||
("consistent", None, 0),
|
||
("canonical-drift", "canonical", 1)):
|
||
with self.subTest(case=case):
|
||
conn = _ResetConnection(input_drift=input_drift)
|
||
capture = Mock(wraps=backup.capture_input_snapshot)
|
||
result = self._invoke_reset_execute(
|
||
conn,
|
||
manifest=manifest,
|
||
current_domains=normalized_domains,
|
||
current_code_files=identity["codeFiles"],
|
||
capture_input=capture,
|
||
)
|
||
|
||
self.assertEqual(result.exit_code, expected_exit, result.output)
|
||
capture.assert_called_once()
|
||
capture_cursor = capture.call_args.args[0]
|
||
self.assertIsInstance(capture_cursor, _CursorContext)
|
||
self.assertIs(capture_cursor.row_factory, reset.dict_row)
|
||
writes = [sql for sql, _params in conn.executions
|
||
if sql.startswith(("update ", "delete "))]
|
||
if input_drift:
|
||
self.assertIn("input 漂移", result.output)
|
||
self.assertEqual(writes, [])
|
||
self.assertEqual(conn.commits, 0)
|
||
else:
|
||
self.assertTrue(writes)
|
||
self.assertEqual(conn.commits, 1)
|
||
|
||
def test_reset_execute_rejects_complete_input_drift_before_writes(self):
|
||
"""正文、字段合同或数据库身份漂移均必须在 destructive SQL 前失败。"""
|
||
|
||
cases = (
|
||
("canonical", ("canonicalContent", "sha"), "d" * 64, "input 漂移"),
|
||
("contracts", ("activeFieldContracts", "sha"), "e" * 64, "input 漂移"),
|
||
("database", ("databaseIdentity", "database"), "other-db", "数据库 identity"),
|
||
)
|
||
for case, path, value, expected_message in cases:
|
||
with self.subTest(case=case):
|
||
manifest, domains = _manifest_for(_snapshot_domains())
|
||
current_input = copy.deepcopy(manifest["input"])
|
||
current_input[path[0]][path[1]] = value
|
||
conn = _ResetConnection()
|
||
|
||
result = self._invoke_reset_execute(
|
||
conn,
|
||
manifest=manifest,
|
||
current_domains=domains,
|
||
current_input=current_input,
|
||
)
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn(expected_message, result.output)
|
||
self.assertEqual(conn.commits, 0)
|
||
self.assertEqual(len(conn.executions), 1)
|
||
self.assertTrue(conn.executions[0][0].startswith("lock table"))
|
||
self.assertFalse(any(sql.startswith(("update ", "delete "))
|
||
for sql, _params in conn.executions))
|
||
|
||
def test_reset_execute_rejects_missing_or_drifted_code_identity_before_database(self):
|
||
"""七域不变时,旧 manifest 缺 reset 或 reset/parse 任一 fileSha 漂移仍须拒绝。"""
|
||
|
||
for case in ("missing-reset", "reset-drift", "parse-drift"):
|
||
with self.subTest(case=case):
|
||
manifest_files = _code_files()
|
||
current_files = _code_files()
|
||
expected_name = "reset_upgrade_work.py"
|
||
if case == "missing-reset":
|
||
manifest_files.pop("reset_upgrade_work.py")
|
||
elif case == "reset-drift":
|
||
current_files["reset_upgrade_work.py"] = "x" * 64
|
||
else:
|
||
current_files["parse_upgrade.py"] = "x" * 64
|
||
expected_name = "parse_upgrade.py"
|
||
manifest, domains = _manifest_for(
|
||
_snapshot_domains(), code_files=manifest_files,
|
||
)
|
||
conn = _ResetConnection()
|
||
result = self._invoke_reset_execute(
|
||
conn,
|
||
manifest=manifest,
|
||
current_domains=domains,
|
||
current_code_files=current_files,
|
||
)
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn(expected_name, result.output)
|
||
self.assertEqual(conn.executions, [])
|
||
self.assertEqual(conn.commits, 0)
|
||
|
||
def test_reset_execute_rejects_confirmed_entity_vectors_before_writes(self):
|
||
"""目标 draft 任一 entity_id 非空时失败关闭,不能软删已确认实体向量。"""
|
||
|
||
manifest, domains = _manifest_for(_snapshot_domains(entity_id=9001))
|
||
conn = _ResetConnection(confirmed_draft_id=101)
|
||
result = self._invoke_reset_execute(
|
||
conn, manifest=manifest, current_domains=domains,
|
||
)
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("entity_id", result.output)
|
||
self.assertEqual(conn.commits, 0)
|
||
self.assertFalse(any(sql.startswith(("update ", "delete "))
|
||
for sql, _params in conn.executions))
|
||
self.assertFalse(conn.embeddings[101]["deleted"])
|
||
|
||
def test_reset_execute_rejects_confirmation_mismatch_before_database(self):
|
||
"""backup_id 或 confirmationSha 不匹配时不得创建业务连接。"""
|
||
|
||
manifest, _domains = _manifest_for(_snapshot_domains())
|
||
with patch.object(reset, "upgrade_work_lock", _lock_success), \
|
||
patch.object(reset, "connect") as connect, \
|
||
patch.object(backup, "verify_backup", return_value=manifest):
|
||
result = self.runner.invoke(
|
||
reset.main,
|
||
["--work-id", "8", "--execute", "--backup-dir", "/private/tmp",
|
||
"--backup-id", BACKUP_ID, "--confirmation-sha", "wrong"],
|
||
)
|
||
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("confirmation-sha", result.output)
|
||
connect.assert_not_called()
|
||
|
||
def test_reset_execute_rolls_back_when_any_postcheck_fails(self):
|
||
"""七类提交前断言任一失败都必须回滚全部已执行 reset 写入。"""
|
||
|
||
for failure in ("cards", "windows", "alias", "presence", "state", "audit", "vectors"):
|
||
with self.subTest(failure=failure):
|
||
conn = _ResetConnection(postcheck_failure=failure)
|
||
result = self._invoke_reset_execute(conn)
|
||
self.assertNotEqual(result.exit_code, 0)
|
||
self.assertIn("清后复核失败", result.output)
|
||
self.assertEqual(conn.commits, 0)
|
||
self.assertFalse(conn.drafts[101]["deleted"])
|
||
self.assertTrue(conn.drafts[102]["deleted"])
|
||
self.assertEqual(conn.window_statuses, ["done", "pending"])
|
||
self.assertEqual(
|
||
(conn.alias_count, conn.presence_count,
|
||
conn.card_state_count, conn.audit_count),
|
||
(3, 4, 2, 5),
|
||
)
|
||
for embedding in conn.embeddings.values():
|
||
self.assertFalse(embedding["deleted"])
|
||
self.assertEqual(embedding["updater"], "")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|