muse-agent-example/tests/skills/embed-knowledge/test_embed_drafts_offline.py
zizi 061010ba1b 框架: 共享运行时装成可安装包,切断 Skill 之间的 sys.path 互指
被多个 Skill 或看板消费的连接、模型、嵌入、Claude 运行时与声音账入口
各只保留一份实现;Skill 只留 CLI/落库,看板只读 muse-db,门禁锁死跨域注入。

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-20 00:45:20 +08:00

1001 lines
40 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
"""embed_drafts 并发与向量 owner 规则的纯离线测试。"""
import copy
import hashlib
import importlib.util
import pathlib
import unittest
from unittest.mock import MagicMock, Mock, patch
from click.testing import CliRunner
import muse_embed as embed
# CLI 壳留在 Skill 里、不随包安装,只能按路径加载
CLI_PATH = (pathlib.Path(__file__).resolve().parents[3]
/ ".claude/skills/embed-knowledge/scripts/embed_drafts.py")
_spec = importlib.util.spec_from_file_location("embed_drafts_cli", CLI_PATH)
cli = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(cli)
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(cli, "_session") as session, \
patch.object(cli, "connect", return_value=connection_context) as connect, \
patch.object(cli, "_run_bulk") as run_bulk, \
patch.object(cli, "embed_texts") as http:
result = CliRunner().invoke(cli.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()