muse-agent-example/.claude/skills/embed/scripts/test_embed_drafts_offline.py

411 lines
16 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 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()