#!/usr/bin/env python3 """embed_drafts 并发与向量 owner 规则的纯离线测试。""" import copy import hashlib import pathlib import sys import unittest from unittest.mock import MagicMock, Mock, patch from click.testing import CliRunner PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3] SCRIPT_DIR = PROJECT_ROOT / ".claude" / "skills" / "embed-knowledge" / "scripts" 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, "model": embed.MODEL, "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, model, 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["model"], 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], "model": params[3], "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[3] or row["model"] != params[4])): 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 或旧 model 活向量。""" assert "entity_id is null" in sql assert "content_hash!=%s or model!=%s" in sql assert params == (embed.ACTOR, embed.TENANT, 101, CONTENT_HASH, embed.MODEL) @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 _BulkConnection: """模拟 bulk 的候选读取、owner 预查和逐 draft 写事务。""" def __init__(self, drafts, vectors=None): self.drafts = { draft_id: { "tenant": embed.TENANT, "deleted": False, "status": "pending", "payload": payload, "source_id": source_id, "work_id": work_id, "source_type": source_type, } for draft_id, payload, source_id, source_type, work_id in ( (item[0], item[1], item[2], item[3], item[4]) if len(item) == 5 else (item[0], item[1], item[2], item[3], item[2]) if len(item) == 4 else (*item, None, item[2]) for item in drafts ) } self.vectors = [dict(vector) for vector in (vectors or [])] self.events = [] self.sql = [] self._next_vector_id = 1000 def execute(self, sql, params=None): normalized = " ".join(sql.split()).lower() self.sql.append((normalized, params)) if normalized.startswith("select d.id, d.draft_payload"): self._assert_bulk_select_sql(normalized) work_id = params[2] if len(params) >= 3 else None rows = [] for draft_id, draft in sorted(self.drafts.items()): if (draft["tenant"] != embed.TENANT or draft["deleted"] or draft["status"] != "pending"): continue if "d.work_id=%s" in normalized and "d.source_type=%s" in normalized: if (draft["work_id"] != work_id or draft["source_type"] != params[3]): continue elif work_id is not None and draft["source_id"] != work_id: continue active = sorted( (row for row in self.vectors if row["tenant"] == embed.TENANT and row["draft_id"] == draft_id and not row["deleted"]), key=lambda row: row["id"], ) if not active: rows.append((draft_id, draft["payload"], None, None, None, None)) continue rows.extend(( draft_id, draft["payload"], row["id"], row["content_hash"], row["model"], row["entity_id"], ) for row in active) self.events.append("bulk_select") return _Result(rows=rows) if normalized.startswith("lock table"): self.events.append("table_lock") return _Result() if normalized.startswith( "select tenant_id, deleted, status, draft_payload from muse_knowledge_draft"): draft_id = params[0] draft = self.drafts.get(draft_id) self.events.append(f"candidate_lock:{draft_id}") if draft is None: return _Result(None) return _Result(( draft["tenant"], draft["deleted"], draft["status"], draft["payload"], )) if normalized.startswith( "select id, content_hash, model, entity_id from example_knowledge_embedding"): self.assertInSql(normalized, "for update") draft_id = params[1] rows = [ (row["id"], row["content_hash"], row["model"], row["entity_id"]) for row in self.vectors if row["tenant"] == params[0] and row["draft_id"] == draft_id and not row["deleted"] ] self.events.append(f"draft_vectors_lock:{draft_id}") return _Result(rows=rows) if normalized.startswith("select e.draft_id, e.entity_id, e.deleted"): tenant, content_hash, model = params owner = next(( row for row in self.vectors if row["tenant"] == tenant and row["content_hash"] == content_hash and row["model"] == model ), None) self.events.append("owner_lock" if "for update of e" in normalized else "owner_read") if owner is None: return _Result(None) owner_draft = self.drafts.get(owner["draft_id"]) return _Result(( owner["draft_id"], owner["entity_id"], owner["deleted"], True if owner_draft is None else owner_draft["deleted"], None if owner_draft is None else owner_draft["tenant"], )) if normalized.startswith("update example_knowledge_embedding set deleted=true"): self.assertInSql(normalized, "content_hash!=%s or model!=%s") _, tenant, draft_id, content_hash, model = params for row in self.vectors: if (row["tenant"] == tenant and row["draft_id"] == draft_id and not row["deleted"] and row["entity_id"] is None and (row["content_hash"] != content_hash or row["model"] != model)): row["deleted"] = True self.events.append(f"old_vectors_soft_delete:{draft_id}") return _Result() if normalized.startswith("insert into example_knowledge_embedding"): draft_id, content_hash, _, model = params[:4] owner = next(( row for row in self.vectors if row["tenant"] == embed.TENANT and row["content_hash"] == content_hash and row["model"] == model ), None) if owner is not None: owner_draft = self.drafts.get(owner["draft_id"]) owner_active = owner_draft is not None and not owner_draft["deleted"] if (owner["entity_id"] is not None or (owner["draft_id"] != draft_id and owner_active)): return _Result(None) owner.update({"draft_id": draft_id, "entity_id": None, "deleted": False}) else: self.vectors.append({ "id": self._next_vector_id, "tenant": embed.TENANT, "draft_id": draft_id, "content_hash": content_hash, "model": model, "entity_id": None, "deleted": False, }) self._next_vector_id += 1 self.events.append(f"upsert:{draft_id}") return _Result((draft_id,)) raise AssertionError(f"未覆盖的 bulk 离线 SQL:{normalized}") def commit(self): self.events.append("commit") def transaction(self): return self._Transaction(self) @staticmethod def assertInSql(sql, fragment): assert fragment in sql, f"SQL 缺少约束:{fragment}\n{sql}" @classmethod def _assert_bulk_select_sql(cls, sql): """候选查询必须读取全部活向量,筛选和 limit 留给 Python。""" for fragment in ( "left join example_knowledge_embedding e", "e.deleted=false", "e.content_hash", "e.model", "e.entity_id", "d.deleted=false", "d.status='pending'", "order by d.id", "e.id"): cls.assertInSql(sql, fragment) assert " limit " not in f" {sql} " class _Transaction: """异常时恢复向量状态,模拟 PostgreSQL 单 draft 事务回滚。""" def __init__(self, conn): self.conn = conn self.snapshot = None def __enter__(self): self.snapshot = copy.deepcopy(self.conn.vectors) return self def __exit__(self, exc_type, exc_value, traceback): if exc_type is not None: self.conn.vectors = self.snapshot return False def _draft(draft_id, text, source_id=1, source_type=None, work_id=None): if work_id is None: return draft_id, {"embed_text": text}, source_id, source_type return draft_id, {"embed_text": text}, source_id, source_type, work_id def _vector(vector_id, draft_id, content_hash, *, model=embed.MODEL, entity_id=None, deleted=False): return { "id": vector_id, "tenant": embed.TENANT, "draft_id": draft_id, "content_hash": content_hash, "model": model, "entity_id": entity_id, "deleted": deleted, } class _EmbeddingResponse: """模拟 embeddings HTTP 响应,仅提供生产代码实际读取的方法。""" def __init__(self, data): self.data = data def raise_for_status(self): return None def json(self): return {"data": self.data} class _EmbeddingSession: """按请求文本动态返回带原始 index 的离线 embeddings 响应。""" def __init__(self, responder): self.responder = responder self.inputs = [] def post(self, _url, *, json, timeout): assert timeout == 120 inputs = list(json["input"]) self.inputs.append(inputs) return _EmbeddingResponse(self.responder(inputs)) class EmbedDraftsOfflineTest(unittest.TestCase): """覆盖 reset/embed 两种先后顺序与 hash owner 反例。""" VECTOR = [0.1, 0.2] def test英文作品卡的名称摘要字段进入嵌入文本(self): first = embed.build_embed_text({ "type": "character", "name": "林深", "brief": "机师", "fields": {"身份": "驾驶员"}, }) second = embed.build_embed_text({ "type": "character", "name": "何岚", "brief": "指挥员", "fields": {"身份": "舰长"}, }) self.assertIn("林深", first) self.assertIn("机师", first) self.assertIn("身份:驾驶员", first) self.assertNotEqual(embed._content_hash(first), embed._content_hash(second)) 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"]) class BulkSelfHealingOfflineTest(unittest.TestCase): """直接覆盖 bulk 候选筛选、HTTP 边界和补嵌自愈结果。""" VECTOR = [0.1, 0.2] @staticmethod def _candidate_ids(conn, limit=0): candidates, failures = embed._load_bulk_candidates(conn, work_id=None, limit=limit) return [candidate[0] for candidate in candidates], failures def test无活向量入选且limit作用于Python筛选结果(self): conn = _BulkConnection( drafts=[_draft(100, TEXT), _draft(101, NEW_TEXT)], vectors=[_vector(700, 100, CONTENT_HASH)], ) candidate_ids, failures = self._candidate_ids(conn, limit=1) self.assertEqual(candidate_ids, [101]) self.assertEqual(failures, []) def test旧hash与旧model均入选并替换为唯一当前活向量(self): conn = _BulkConnection( drafts=[_draft(101, TEXT), _draft(102, NEW_TEXT)], vectors=[ _vector(701, 101, "old-hash"), _vector(702, 102, NEW_HASH, model="Old/Embedding-Model"), ], ) with patch.object( embed, "embed_texts", return_value=([self.VECTOR, self.VECTOR], set())): stats = embed._run_bulk(conn, sess=object(), work_id=None, limit=0) self.assertEqual(stats, {"done": 2, "skip": 0, "fail": 0}) for draft_id, expected_hash in ((101, CONTENT_HASH), (102, NEW_HASH)): active = [ row for row in conn.vectors if row["draft_id"] == draft_id and not row["deleted"] ] self.assertEqual(len(active), 1) self.assertEqual((active[0]["content_hash"], active[0]["model"]), (expected_hash, embed.MODEL)) def test批响应缺少index零时不压缩错写并逐条降级(self): first_vector = [0.1] second_vector = [0.2] def respond(inputs): if len(inputs) == 2: # 服务端只返回第二条;若按排序后压缩,会被错误写给第一个 draft。 return [{"index": 1, "embedding": second_vector}] if inputs == [TEXT]: return [] if inputs == [NEW_TEXT]: return [{"index": 0, "embedding": second_vector}] raise AssertionError(f"未预期的输入:{inputs}") conn = _BulkConnection( drafts=[_draft(101, TEXT), _draft(102, NEW_TEXT)], vectors=[ _vector(701, 101, "old-hash-1"), _vector(702, 102, "old-hash-2"), ], ) sess = _EmbeddingSession(respond) with patch.object(embed.time, "sleep"): stats = embed._run_bulk(conn, sess=sess, work_id=None, limit=0) self.assertEqual(stats, {"done": 1, "skip": 0, "fail": 1}) self.assertEqual(sess.inputs, [ [TEXT, NEW_TEXT], [TEXT, NEW_TEXT], [TEXT, NEW_TEXT], [TEXT], [NEW_TEXT], ]) active_101 = [row for row in conn.vectors if row["draft_id"] == 101 and not row["deleted"]] active_102 = [row for row in conn.vectors if row["draft_id"] == 102 and not row["deleted"]] self.assertEqual([(row["content_hash"], row["model"]) for row in active_101], [("old-hash-1", embed.MODEL)]) self.assertEqual([(row["content_hash"], row["model"]) for row in active_102], [(NEW_HASH, embed.MODEL)]) candidate_ids, failures = self._candidate_ids(conn) self.assertEqual(candidate_ids, [101]) self.assertEqual(failures, []) def test乱序但完整的响应按原始index还原(self): first_vector = [0.1] second_vector = [0.2] sess = _EmbeddingSession(lambda _inputs: [ {"index": 1, "embedding": second_vector}, {"index": 0, "embedding": first_vector}, ]) vectors, bad = embed.embed_texts(sess, [TEXT, NEW_TEXT]) self.assertEqual(vectors, [first_vector, second_vector]) self.assertEqual(bad, set()) self.assertEqual(sess.inputs, [[TEXT, NEW_TEXT]]) def test当前hash与model已有唯一活向量时跳过HTTP(self): conn = _BulkConnection( drafts=[_draft(101, TEXT)], vectors=[_vector(701, 101, CONTENT_HASH)], ) http = Mock() with patch.object(embed, "embed_texts", http): stats = embed._run_bulk(conn, sess=object(), work_id=None, limit=0) self.assertEqual(stats, {"done": 0, "skip": 0, "fail": 0}) http.assert_not_called() def test目标hash已有entity_owner时HTTP前失败且绝不迁移(self): conn = _BulkConnection( drafts=[_draft(101, TEXT)], vectors=[_vector(700, 90, CONTENT_HASH, entity_id=9001)], ) http = Mock() with patch.object(embed, "embed_texts", http), \ patch.object(embed.click, "echo"): with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "entity"): embed._run_bulk(conn, sess=object(), work_id=None, limit=0) http.assert_not_called() self.assertEqual(conn.vectors[0]["draft_id"], 90) self.assertEqual(conn.vectors[0]["entity_id"], 9001) def test目标hash已有其它active_draft_owner时向上抛且绝不迁移(self): conn = _BulkConnection( drafts=[_draft(90, TEXT), _draft(101, TEXT)], vectors=[_vector(700, 90, CONTENT_HASH)], ) conn.drafts[90]["status"] = "confirmed" http = Mock() with patch.object(embed, "embed_texts", http): with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "owner=90"): embed._run_bulk(conn, sess=object(), work_id=None, limit=0) http.assert_not_called() self.assertEqual(conn.vectors[0]["draft_id"], 90) self.assertFalse(conn.vectors[0]["deleted"]) def test_limit范围外候选被entity_owner占用时HTTP前向上抛(self): conn = _BulkConnection( drafts=[_draft(90, NEW_TEXT), _draft(101, TEXT), _draft(102, NEW_TEXT)], vectors=[_vector(700, 90, NEW_HASH, entity_id=9001)], ) conn.drafts[90]["status"] = "confirmed" http = Mock() with patch.object(embed, "embed_texts", http): with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "entity"): embed._run_bulk(conn, sess=object(), work_id=None, limit=1) http.assert_not_called() self.assertEqual(len(conn.vectors), 1) self.assertEqual((conn.vectors[0]["draft_id"], conn.vectors[0]["entity_id"]), (90, 9001)) def test_limit范围外候选被其它active_owner占用时HTTP前向上抛(self): conn = _BulkConnection( drafts=[_draft(90, NEW_TEXT), _draft(101, TEXT), _draft(102, NEW_TEXT)], vectors=[_vector(700, 90, NEW_HASH)], ) conn.drafts[90]["status"] = "confirmed" http = Mock() with patch.object(embed, "embed_texts", http): with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "owner=90"): embed._run_bulk(conn, sess=object(), work_id=None, limit=1) http.assert_not_called() self.assertEqual(len(conn.vectors), 1) self.assertEqual(conn.vectors[0]["draft_id"], 90) self.assertFalse(conn.vectors[0]["deleted"]) def test_limit前预检全部owner但仅对限量候选发HTTP并写入(self): conn = _BulkConnection( drafts=[_draft(101, TEXT), _draft(102, NEW_TEXT)], ) sess = object() def respond(_sess, texts): self.assertEqual(texts, [TEXT]) self.assertEqual(conn.events.count("owner_read"), 3) self.assertEqual(conn.events[-1], "commit") return [self.VECTOR], set() with patch.object(embed, "embed_texts", side_effect=respond) as http: stats = embed._run_bulk(conn, sess=sess, work_id=None, limit=1) self.assertEqual(stats, {"done": 1, "skip": 0, "fail": 0}) http.assert_called_once() first_commit = conn.events.index("commit") self.assertEqual(conn.events[:first_commit].count("owner_read"), 2) active = [row for row in conn.vectors if not row["deleted"]] self.assertEqual([(row["draft_id"], row["content_hash"]) for row in active], [(101, CONTENT_HASH)]) def test读段发现同draft多条活向量时失败关闭并明确报告(self): conn = _BulkConnection( drafts=[_draft(100, TEXT), _draft(101, NEW_TEXT)], vectors=[ _vector(701, 101, "old-hash-1"), _vector(702, 101, "old-hash-2"), ], ) http = Mock() with patch.object(embed, "embed_texts", http), \ patch.object(embed.click, "echo") as echo: with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "多条活向量"): embed._run_bulk(conn, sess=object(), work_id=None, limit=1) http.assert_not_called() self.assertTrue(all(not row["deleted"] for row in conn.vectors)) self.assertTrue(any("多条活向量" in call.args[0] for call in echo.call_args_list)) def test_HTTP后写前出现多条活向量时事务失败关闭并明确报告(self): conn = _BulkConnection(drafts=[_draft(101, TEXT)]) def race_after_http(_sess, _texts): conn.vectors.extend([ _vector(701, 101, "race-old-1"), _vector(702, 101, "race-old-2"), ]) return [self.VECTOR], set() with patch.object(embed, "embed_texts", side_effect=race_after_http) as http, \ patch.object(embed.click, "echo"): with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "多条活向量"): embed._run_bulk(conn, sess=object(), work_id=None, limit=0) http.assert_called_once() self.assertEqual(len(conn.vectors), 2) self.assertTrue(all(not row["deleted"] for row in conn.vectors)) def test_HTTP坏结果零写入且下一轮仍可入选(self): bad_results = ( ([], set()), ([None], set()), ([self.VECTOR], {0}), ) for result in bad_results: with self.subTest(result=result): conn = _BulkConnection( drafts=[_draft(101, TEXT)], vectors=[_vector(701, 101, "old-hash")], ) with patch.object(embed, "embed_texts", return_value=result): stats = embed._run_bulk(conn, sess=object(), work_id=None, limit=0) self.assertEqual(stats, {"done": 0, "skip": 0, "fail": 1}) self.assertFalse(conn.vectors[0]["deleted"]) candidate_ids, failures = self._candidate_ids(conn) self.assertEqual(candidate_ids, [101]) self.assertEqual(failures, []) def test_HTTP网络异常零写入且下一轮仍可入选(self): conn = _BulkConnection( drafts=[_draft(101, TEXT)], vectors=[_vector(701, 101, "old-hash")], ) with patch.object(embed, "embed_texts", side_effect=OSError("network down")): stats = embed._run_bulk(conn, sess=object(), work_id=None, limit=0) self.assertEqual(stats, {"done": 0, "skip": 0, "fail": 1}) self.assertFalse(conn.vectors[0]["deleted"]) candidate_ids, failures = self._candidate_ids(conn) self.assertEqual(candidate_ids, [101]) self.assertEqual(failures, []) def test同批目标hash冲突在limit前失败关闭且无人抢owner(self): conn = _BulkConnection( drafts=[_draft(101, TEXT), _draft(102, TEXT)], ) http = Mock() with patch.object(embed, "embed_texts", http), \ patch.object(embed.click, "echo") as echo: with self.assertRaisesRegex(embed.EmbeddingOwnershipConflict, "目标 hash 冲突"): embed._run_bulk(conn, sess=object(), work_id=None, limit=1) http.assert_not_called() self.assertEqual(conn.vectors, []) messages = "\n".join(call.args[0] for call in echo.call_args_list) self.assertIn("目标 hash 冲突", messages) self.assertIn("101", messages) self.assertIn("102", messages) def test章后抽卡按draft_work_id和source_type筛选(self): conn = _BulkConnection( drafts=[ _draft(201, TEXT, source_id=17473, source_type="chapter_extract", work_id=12), _draft(202, NEW_TEXT, source_id=17474, source_type="chapter_extract", work_id=12), _draft(203, TEXT, source_id=12, source_type="parse_book"), ], ) candidates, failures = embed._load_bulk_candidates( conn, work_id=12, limit=0, source_type="chapter_extract" ) self.assertEqual([row[0] for row in candidates], [201, 202]) self.assertEqual(failures, []) select_sql, params = next(sql for sql in conn.sql if sql[0].startswith("select d.id")) self.assertIn("d.work_id=%s", select_sql) self.assertIn("d.source_type=%s", select_sql) self.assertEqual(params, [embed.TENANT, embed.TENANT, 12, "chapter_extract"]) class EmbedDraftsCliTest(unittest.TestCase): """验证 CLI 参数在创建 session 和访问外部资源前完成校验。""" def test_limit负数在任何session数据库或HTTP前被拒绝(self): connection_context = MagicMock() connection_context.__enter__.return_value = MagicMock() connection_context.__exit__.return_value = False with patch.object(embed, "_session") as session, \ patch.object(embed.psycopg, "connect", return_value=connection_context) as connect, \ patch.object(embed, "_run_bulk") as run_bulk, \ patch.object(embed, "embed_texts") as http: result = CliRunner().invoke(embed.main, ["--limit", "-1"]) self.assertNotEqual(result.exit_code, 0) self.assertIn("--limit", result.output) self.assertIn("x>=0", result.output) session.assert_not_called() connect.assert_not_called() run_bulk.assert_not_called() http.assert_not_called() if __name__ == "__main__": unittest.main()