#!/usr/bin/env python3 """embed_drafts 并发与向量 owner 规则的纯离线测试。""" import hashlib import pathlib import sys import unittest from unittest.mock import patch SCRIPT_DIR = pathlib.Path(__file__).resolve().parent sys.path.insert(0, str(SCRIPT_DIR)) import embed_drafts as embed # noqa: E402 TEXT = "文本" CONTENT_HASH = hashlib.sha256(f"{TEXT}|{embed.MODEL}".encode()).hexdigest() NEW_TEXT = "新文本" NEW_HASH = hashlib.sha256(f"{NEW_TEXT}|{embed.MODEL}".encode()).hexdigest() class _Result: """提供 psycopg 查询结果所需的最小读取接口。""" 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 _EmbeddingConnection: """模拟 draft 与同 hash 唯一行,记录写事务内的锁和 UPSERT 顺序。""" def __init__(self, *, candidate_deleted=False, candidate_status="pending", candidate_payload=None, draft_live_embeddings=None, owner_draft_id=None, owner_entity_id=None, owner_embedding_deleted=False, owner_draft_deleted=None, owner_tenant=embed.TENANT, inject_active_owner_before_upsert=False): self.candidate = { "tenant": embed.TENANT, "deleted": candidate_deleted, "status": candidate_status, "payload": candidate_payload or {"embed_text": TEXT}, } self.draft_live_embeddings = [ {"id": row_id, "content_hash": content_hash, "entity_id": entity_id, "deleted": False} for row_id, content_hash, entity_id in (draft_live_embeddings or []) ] self.embedding = None if owner_draft_id is not None or owner_entity_id is not None: self.embedding = { "draft_id": owner_draft_id, "entity_id": owner_entity_id, "deleted": owner_embedding_deleted, "owner_deleted": owner_draft_deleted, "owner_tenant": owner_tenant, } self.events = [] self.sql = [] self.inject_active_owner_before_upsert = inject_active_owner_before_upsert def execute(self, sql, params=None): normalized = " ".join(sql.split()).lower() self.sql.append((normalized, params)) if normalized.startswith("lock table"): self.events.append("table_lock") self.assert_table_lock_sql(normalized, params) return _Result() if normalized.startswith( "select tenant_id, deleted, status, draft_payload from muse_knowledge_draft"): self.events.append("candidate_lock") self.assert_candidate_lock_sql(normalized, params) return _Result(( self.candidate["tenant"], self.candidate["deleted"], self.candidate["status"], self.candidate["payload"], )) if normalized.startswith("select id, content_hash, entity_id from example_knowledge_embedding"): self.events.append("draft_vectors_lock") self.assert_draft_vectors_lock_sql(normalized, params) rows = [ (row["id"], row["content_hash"], row["entity_id"]) for row in self.draft_live_embeddings if not row["deleted"] ] return _Result(rows=rows) if normalized.startswith("select e.draft_id, e.entity_id, e.deleted"): self.events.append("owner_lock" if "for update of e" in normalized else "owner_read") self.assert_owner_sql(normalized, params) if self.embedding is None: return _Result(None) row = self.embedding owner_deleted = row["owner_deleted"] if owner_deleted is None: owner_deleted = True return _Result(( row["draft_id"], row["entity_id"], row["deleted"], owner_deleted, row["owner_tenant"], )) if normalized.startswith("insert into example_knowledge_embedding"): self.events.append("upsert") self.assert_upsert_sql(normalized, params) candidate_draft_id = params[0] # 模拟 owner 预查时尚无唯一行,随后并发事务先插入另一活跃 owner。 if self.inject_active_owner_before_upsert and self.embedding is None: self.embedding = { "draft_id": 90, "entity_id": None, "deleted": False, "owner_deleted": False, "owner_tenant": embed.TENANT, } if self.embedding is not None: row = self.embedding owner_is_active = ( row["draft_id"] is not None and row["owner_deleted"] is False ) can_take = ( row["entity_id"] is None and (row["draft_id"] == candidate_draft_id or not owner_is_active) ) if not can_take: return _Result(None) self.embedding = { "draft_id": candidate_draft_id, "entity_id": None, "deleted": False, "owner_deleted": False, "owner_tenant": embed.TENANT, } self.draft_live_embeddings.append({ "id": 999, "content_hash": params[1], "entity_id": None, "deleted": False, }) return _Result((candidate_draft_id,)) if normalized.startswith("update example_knowledge_embedding set deleted=true"): self.events.append("old_vectors_soft_delete") self.assert_old_vectors_update_sql(normalized, params) for row in self.draft_live_embeddings: if (not row["deleted"] and row["entity_id"] is None and row["content_hash"] != params[2]): row["deleted"] = True return _Result() raise AssertionError(f"未覆盖的离线 SQL:{normalized}") @staticmethod def assert_table_lock_sql(sql, params): """写事务第一条 SQL 必须按统一顺序取得两表 ROW EXCLUSIVE 锁。""" assert sql == ( "lock table muse_knowledge_draft, example_knowledge_embedding " "in row exclusive mode" ) assert params is None @staticmethod def assert_candidate_lock_sql(sql, params): """candidate 必须按主键锁行,并在同一快照读取当前 payload。""" assert "where id=%s" in sql assert "for update" in sql assert params == (101,) @staticmethod def assert_draft_vectors_lock_sql(sql, params): """必须锁定当前 draft 的全部活向量,而非只看目标 hash。""" assert "where tenant_id=%s and draft_id=%s and deleted=false" in sql assert "for update" in sql assert params == (embed.TENANT, 101) @staticmethod def assert_old_vectors_update_sql(sql, params): """只软删当前 draft 的无 entity 旧 hash 活向量。""" assert "entity_id is null" in sql assert "content_hash!=%s" in sql assert params == (embed.ACTOR, embed.TENANT, 101, CONTENT_HASH) @staticmethod def assert_owner_sql(sql, params): """hash 查询必须读取 owner 全状态;写段还必须锁住唯一行。""" for fragment in ( "e.draft_id", "e.entity_id", "e.deleted", "coalesce(d.deleted, true)", "d.tenant_id", "left join muse_knowledge_draft d on d.id=e.draft_id", "e.tenant_id=%s", "e.content_hash=%s", "e.model=%s"): assert fragment in sql assert params == (embed.TENANT, CONTENT_HASH, embed.MODEL) @staticmethod def assert_upsert_sql(sql, params): """唯一键冲突只能迁移无 entity 且 owner 已失活的行,并校验返回 owner。""" for fragment in ( "on conflict (tenant_id, content_hash, model)", "do update set draft_id=excluded.draft_id", "example_knowledge_embedding.entity_id is null", "example_knowledge_embedding.draft_id=excluded.draft_id", "not exists ( select 1 from muse_knowledge_draft owner", "owner.id=example_knowledge_embedding.draft_id", "owner.deleted=false", "returning draft_id"): assert fragment in sql assert params[0] == 101 assert params[1] == CONTENT_HASH class EmbedDraftsOfflineTest(unittest.TestCase): """覆盖 reset/embed 两种先后顺序与 hash owner 反例。""" VECTOR = [0.1, 0.2] def test_reset先完成时软删draft零写入且有可追踪输出(self): conn = _EmbeddingConnection( candidate_deleted=True, owner_draft_id=90, owner_embedding_deleted=True, owner_draft_deleted=True, ) with patch.object(embed.click, "echo") as echo: written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertFalse(written) self.assertEqual(conn.events, ["table_lock", "candidate_lock"]) self.assertEqual(conn.embedding["draft_id"], 90) echo.assert_called_once() self.assertIn("draft=101", echo.call_args.args[0]) self.assertIn("软删", echo.call_args.args[0]) def test_embed先锁draft时先锁后写且写入会造成reset摘要漂移(self): conn = _EmbeddingConnection() before = conn.embedding with patch.object(embed.click, "echo"): written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertTrue(written) self.assertEqual( conn.events, ["table_lock", "candidate_lock", "draft_vectors_lock", "owner_lock", "upsert"], ) self.assertIsNone(before) self.assertEqual(conn.embedding["draft_id"], 101) self.assertFalse(conn.embedding["deleted"]) def test_active同did才允许幂等跳过(self): conn = _EmbeddingConnection( owner_draft_id=101, owner_embedding_deleted=False, owner_draft_deleted=False, ) action = embed._embedding_owner_action(conn, 101, CONTENT_HASH) self.assertEqual(action, "skip") self.assertEqual(conn.events, ["owner_read"]) def test软删旧draft_owner允许迁移到新did(self): conn = _EmbeddingConnection( owner_draft_id=90, owner_embedding_deleted=True, owner_draft_deleted=True, ) self.assertEqual(embed._embedding_owner_action(conn, 101, CONTENT_HASH), "write") self.assertTrue(embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR)) self.assertEqual(conn.embedding["draft_id"], 101) self.assertFalse(conn.embedding["deleted"]) def test_entity_owner明确冲突(self): conn = _EmbeddingConnection( owner_draft_id=90, owner_entity_id=9001, owner_embedding_deleted=True, owner_draft_deleted=True, ) with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "entity"): embed._embedding_owner_action(conn, 101, CONTENT_HASH) def test其它active_draft_owner明确冲突且绝不迁移(self): conn = _EmbeddingConnection( owner_draft_id=90, owner_embedding_deleted=True, owner_draft_deleted=False, ) with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "owner=90"): embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertNotIn("upsert", conn.events) self.assertEqual(conn.embedding["draft_id"], 90) def test_owner检查后出现active_owner时条件upsert失败关闭(self): conn = _EmbeddingConnection(inject_active_owner_before_upsert=True) with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "未绑定当前 draft"): embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertEqual( conn.events, ["table_lock", "candidate_lock", "draft_vectors_lock", "owner_lock", "upsert"], ) self.assertEqual(conn.embedding["draft_id"], 90) def test_candidate租户不匹配时失败且零写入(self): conn = _EmbeddingConnection() conn.candidate["tenant"] = 2 with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "租户"): embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertEqual(conn.events, ["table_lock", "candidate_lock"]) def test_confirm先完成时非pending_draft可追踪跳过且零写入(self): conn = _EmbeddingConnection(candidate_status="confirmed") with patch.object(embed.click, "echo") as echo: written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertFalse(written) self.assertEqual(conn.events, ["table_lock", "candidate_lock"]) self.assertIsNone(conn.embedding) echo.assert_called_once() self.assertIn("status=confirmed", echo.call_args.args[0]) def test_HTTP期间payload变化时重算hash后跳过且零写入(self): conn = _EmbeddingConnection(candidate_payload={"embed_text": NEW_TEXT}) with patch.object(embed.click, "echo") as echo: written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertFalse(written) self.assertEqual(conn.events, ["table_lock", "candidate_lock"]) self.assertIsNone(conn.embedding) self.assertIn("payload", echo.call_args.args[0]) self.assertIn("漂移", echo.call_args.args[0]) def test_parse并发写入新hash时陈旧结果不软删新向量(self): conn = _EmbeddingConnection( candidate_payload={"embed_text": NEW_TEXT}, draft_live_embeddings=[(701, NEW_HASH, None)], ) with patch.object(embed.click, "echo"): written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertFalse(written) self.assertEqual(conn.events, ["table_lock", "candidate_lock"]) self.assertFalse(conn.draft_live_embeddings[0]["deleted"]) def test_entity旧活向量阻断写入且绝不软删(self): conn = _EmbeddingConnection( draft_live_embeddings=[(701, "old-hash", 9001)], ) with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "entity"): embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertEqual(conn.events, ["table_lock", "candidate_lock", "draft_vectors_lock"]) self.assertFalse(conn.draft_live_embeddings[0]["deleted"]) def test其它hash的draft_owner旧活向量先软删再写入(self): conn = _EmbeddingConnection( draft_live_embeddings=[(701, "old-hash", None)], ) written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertTrue(written) self.assertEqual( conn.events, ["table_lock", "candidate_lock", "draft_vectors_lock", "owner_lock", "old_vectors_soft_delete", "upsert"], ) active = [row for row in conn.draft_live_embeddings if not row["deleted"]] self.assertEqual([(row["content_hash"], row["entity_id"]) for row in active], [(CONTENT_HASH, None)]) def test写前同hash当前活向量幂等跳过(self): conn = _EmbeddingConnection( draft_live_embeddings=[(701, CONTENT_HASH, None)], ) with patch.object(embed.click, "echo"): written = embed._write_embedding(conn, 101, CONTENT_HASH, TEXT, self.VECTOR) self.assertFalse(written) self.assertEqual(conn.events, ["table_lock", "candidate_lock", "draft_vectors_lock"]) if __name__ == "__main__": unittest.main()