muse-agent-example/tests/skills/重置作品抽取结果/test_reset_upgrade_work_offline.py

864 lines
37 KiB
Python
Raw 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.

#!/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 upgrade as parse # noqa: E402
import backup_upgrade_work as backup # noqa: E402
import reset_upgrade_work as reset # 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()