一、技能重组(动作-对象命名) - 旧目录 clean/confirm/continuation/db/detect/embed/… 重组为 clean-book-text/decide-candidate/write-next-chapter/access-database/ check-content-consistency/embed-knowledge/…(git 识别为 rename,内容保持) - agents/*.md、AGENTS.md/CLAUDE.md 收编、example_skill 登记表同步新名 二、先审后入创作闭环(本次核心) 正文接受从"机械门一过就写正典"改为"机械门+语义审查双通过+用户批准+单事务原子提交", DB 级兜底,编排层跳步即被硬拒。 - candidate_cas.py + example_candidate_cas(109):持久化 CAS 状态链 - fact_delta.py + example_fact_delta/example_fact_ledger(106):结构化事实增量, 模型只提六型闭集增量+正文证据引文,仅用户批准的增量随正文同事务入账本 - projection_registry.py + example_projection_run(107):投影登记与恢复 - acceptance_state.py:接受前置实时状态重读 - lesson_registry.py + example_lesson(108):经验升格链,禁止自动升格 - DDL 105:example_candidate 增 semantic_status/semantic_report_sha256 - write_canonical.accept:语义兜底+同事务合并增量+登记投影; run_writer_pipeline/persist_writer_run/run_writer_semantic_detector/step2 接入全链 - claude_runtime:兼容新 CLI modelUsage 信息字段 三、审查修复(独立子代理四维审查后) - 事实增量 propose→approve 翻态正道,不撞唯一键 - 冻结配置探针重刷(CLI 2.1.211→2.1.231 漂移),profileSha256/adapterVersion 再登记 - 可视化合同悬空路径/五六空间矛盾、 SoT 旧技能名漂移、行尾空白清理 测试:离线 65 套 + 真实库集成 5 套(CAS/接受故障注入/事实增量/投影/经验升格)+ 回放 79 项全绿。 创作内容(docs/design、生成正文 artifacts)按"框架与创作分开"未入本提交。
998 lines
40 KiB
Python
998 lines
40 KiB
Python
#!/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
|
||
|
||
|
||
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,
|
||
"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()
|