411 lines
16 KiB
Python
411 lines
16 KiB
Python
#!/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()
|