#!/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()