From bc2f521c2ea8783a0c2871ee76f4e1947ef9bd96 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BC=A0=E8=83=9C=E5=AE=87?= <17412268+zzzzz11122222@user.noreply.gitee.com> Date: Thu, 10 Sep 2026 22:08:15 +0800 Subject: [PATCH] feat: add guarded knowledge publication tool --- app/service/knowledge_publication_service.py | 156 ++++++++++ .../test_knowledge_publication_service.py | 119 ++++++++ ...test_publish_customer_service_knowledge.py | 40 +++ tools/publish_customer_service_knowledge.py | 283 ++++++++++++++++++ 4 files changed, 598 insertions(+) create mode 100644 app/service/knowledge_publication_service.py create mode 100644 tests/unit/service/test_knowledge_publication_service.py create mode 100644 tests/unit/tools/test_publish_customer_service_knowledge.py create mode 100644 tools/publish_customer_service_knowledge.py diff --git a/app/service/knowledge_publication_service.py b/app/service/knowledge_publication_service.py new file mode 100644 index 0000000..57d86ff --- /dev/null +++ b/app/service/knowledge_publication_service.py @@ -0,0 +1,156 @@ +"""管理员知识发布编排:客服运行期只读,本模块仅供显式发布工具调用。""" + +from collections import defaultdict +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Protocol + +from app.core.knowledge_contracts import ALLOWED_KNOWLEDGE_COLLECTIONS + +VECTOR_DIMENSION = 1024 + + +class KnowledgePublicationError(RuntimeError): + """发布前置条件、阶段写入或补偿失败时的明确错误。""" + + +@dataclass(frozen=True) +class KnowledgePublicationRecord: + """预检清单中一条已经审核、可供管理员发布的公开知识。""" + + qa_id: str + milvus_collection: str + retrieval_text: str + title: str + snippet: str + tags: tuple[str, ...] + version: str + metadata: Mapping[str, object] + + +@dataclass(frozen=True) +class KnowledgePublicationResult: + """仅返回可审计的业务编号和数据库主键映射,不返回正文或密钥。""" + + knowledge_ids: dict[str, int] + collections: tuple[str, ...] + + +class KnowledgeEmbedder(Protocol): + async def embed(self, text: str) -> list[float]: ... + + +class KnowledgePublicationStore(Protocol): + async def stage(self, records: tuple[KnowledgePublicationRecord, ...]) -> dict[str, int]: ... + + async def publish(self, knowledge_ids: tuple[int, ...], reviewer_id: int) -> None: ... + + async def disable(self, knowledge_ids: tuple[int, ...]) -> None: ... + + +class KnowledgeVectorPublisher(Protocol): + async def upsert(self, collection: str, records: tuple[dict[str, object], ...]) -> None: ... + + async def delete(self, collection: str, knowledge_ids: tuple[str, ...]) -> None: ... + + +class KnowledgePublicationService: + """把发布动作拆为可补偿阶段,任何中断都不能让未索引知识对客可见。""" + + def __init__( + self, + embedder: KnowledgeEmbedder, + store: KnowledgePublicationStore, + vectors: KnowledgeVectorPublisher, + ) -> None: + self._embedder = embedder + self._store = store + self._vectors = vectors + + async def publish( + self, records: Sequence[KnowledgePublicationRecord], *, reviewer_id: int + ) -> KnowledgePublicationResult: + immutable_records = tuple(records) + self._validate(immutable_records, reviewer_id) + embeddings = await self._embeddings(immutable_records) + knowledge_ids = await self._store.stage(immutable_records) + self._validate_staged_ids(immutable_records, knowledge_ids) + payloads = self._payloads(immutable_records, embeddings, knowledge_ids) + try: + for collection, collection_payloads in payloads.items(): + await self._vectors.upsert(collection, tuple(collection_payloads)) + except Exception as exc: + await self._compensate(payloads, tuple(knowledge_ids.values())) + raise KnowledgePublicationError("向量写入失败,知识保持未发布状态") from exc + await self._store.publish(tuple(knowledge_ids.values()), reviewer_id) + return KnowledgePublicationResult( + knowledge_ids=knowledge_ids, + collections=tuple(payloads), + ) + + @staticmethod + def _validate(records: tuple[KnowledgePublicationRecord, ...], reviewer_id: int) -> None: + if reviewer_id <= 0: + raise KnowledgePublicationError("reviewer_id 必须是正整数") + if not records: + raise KnowledgePublicationError("没有可发布的公开知识") + qa_ids = [record.qa_id for record in records] + if len(qa_ids) != len(set(qa_ids)): + raise KnowledgePublicationError("发布清单存在重复 qa_id") + for record in records: + if record.milvus_collection not in ALLOWED_KNOWLEDGE_COLLECTIONS: + raise KnowledgePublicationError("发布清单包含未授权集合") + if not record.retrieval_text.strip(): + raise KnowledgePublicationError(f"{record.qa_id}: 检索文本不能为空") + + async def _embeddings( + self, records: tuple[KnowledgePublicationRecord, ...] + ) -> dict[str, list[float]]: + embeddings: dict[str, list[float]] = {} + for record in records: + vector = await self._embedder.embed(record.retrieval_text) + if len(vector) != VECTOR_DIMENSION: + raise KnowledgePublicationError( + f"{record.qa_id}: 向量维度必须为 {VECTOR_DIMENSION}" + ) + embeddings[record.qa_id] = vector + return embeddings + + @staticmethod + def _validate_staged_ids( + records: tuple[KnowledgePublicationRecord, ...], knowledge_ids: dict[str, int] + ) -> None: + expected = {record.qa_id for record in records} + if set(knowledge_ids) != expected or any(value <= 0 for value in knowledge_ids.values()): + raise KnowledgePublicationError("MySQL 暂存结果与发布清单不一致") + + @staticmethod + def _payloads( + records: tuple[KnowledgePublicationRecord, ...], + embeddings: dict[str, list[float]], + knowledge_ids: dict[str, int], + ) -> dict[str, list[dict[str, object]]]: + payloads: dict[str, list[dict[str, object]]] = defaultdict(list) + for record in records: + payloads[record.milvus_collection].append({ + "knowledge_id": str(knowledge_ids[record.qa_id]), + "embedding": embeddings[record.qa_id], + "title": record.title, + "snippet": record.snippet, + "tags": list(record.tags), + "version": record.version, + }) + return dict(payloads) + + async def _compensate( + self, payloads: dict[str, list[dict[str, object]]], knowledge_ids: tuple[int, ...] + ) -> None: + for collection, items in payloads.items(): + try: + await self._vectors.delete( + collection, tuple(str(item["knowledge_id"]) for item in items) + ) + except Exception: + # MySQL 行仍会被停用,因此清理失败的残余向量无法对客返回。 + pass + await self._store.disable(knowledge_ids) diff --git a/tests/unit/service/test_knowledge_publication_service.py b/tests/unit/service/test_knowledge_publication_service.py new file mode 100644 index 0000000..e1efa8a --- /dev/null +++ b/tests/unit/service/test_knowledge_publication_service.py @@ -0,0 +1,119 @@ +from dataclasses import dataclass + +import pytest + +from app.service.knowledge_publication_service import ( + KnowledgePublicationError, + KnowledgePublicationService, +) + + +@dataclass(frozen=True) +class Record: + qa_id: str + milvus_collection: str + retrieval_text: str + title: str = "基金开户" + snippet: str = "基金开户" + tags: tuple[str, ...] = ("开户",) + version: str = "v5.8" + + +class FakeEmbedder: + def __init__(self, vector: list[float] | None = None) -> None: + self.vector = vector or [0.1] * 1024 + self.calls: list[str] = [] + + async def embed(self, text: str) -> list[float]: + self.calls.append(text) + return self.vector + + +class FakeStore: + def __init__(self) -> None: + self.staged: list[Record] = [] + self.published: list[int] = [] + self.disabled: list[int] = [] + + async def stage(self, records: tuple[Record, ...]) -> dict[str, int]: + self.staged.extend(records) + return {record.qa_id: index for index, record in enumerate(records, start=101)} + + async def publish(self, knowledge_ids: tuple[int, ...], reviewer_id: int) -> None: + assert reviewer_id == 9 + self.published.extend(knowledge_ids) + + async def disable(self, knowledge_ids: tuple[int, ...]) -> None: + self.disabled.extend(knowledge_ids) + + +class FakeVectors: + def __init__(self, *, fail_collection: str | None = None) -> None: + self.fail_collection = fail_collection + self.upserts: list[tuple[str, tuple[dict[str, object], ...]]] = [] + self.deleted: list[tuple[str, tuple[str, ...]]] = [] + + async def upsert(self, collection: str, records: tuple[dict[str, object], ...]) -> None: + self.upserts.append((collection, records)) + if collection == self.fail_collection: + raise RuntimeError("milvus unavailable") + + async def delete(self, collection: str, knowledge_ids: tuple[str, ...]) -> None: + self.deleted.append((collection, knowledge_ids)) + + +@pytest.mark.asyncio +async def test_publication_stages_vectors_then_publishes_after_all_collections_succeed() -> None: + store = FakeStore() + vectors = FakeVectors() + service = KnowledgePublicationService(FakeEmbedder(), store, vectors) + records = ( + Record("FAQ-001", "fin_faq_collection", "标准问题:基金开户"), + Record("POL-001", "fin_policy_collection", "标准问题:风险测评"), + ) + + result = await service.publish(records, reviewer_id=9) + + assert result.knowledge_ids == {"FAQ-001": 101, "POL-001": 102} + assert [collection for collection, _payload in vectors.upserts] == [ + "fin_faq_collection", "fin_policy_collection" + ] + assert store.published == [101, 102] + assert store.disabled == [] + payload = vectors.upserts[0][1][0] + assert payload["knowledge_id"] == "101" + assert len(payload["embedding"]) == 1024 + + +@pytest.mark.asyncio +async def test_vector_failure_keeps_staged_rows_unpublished_and_compensates_vectors() -> None: + store = FakeStore() + vectors = FakeVectors(fail_collection="fin_policy_collection") + service = KnowledgePublicationService(FakeEmbedder(), store, vectors) + records = ( + Record("FAQ-001", "fin_faq_collection", "标准问题:基金开户"), + Record("POL-001", "fin_policy_collection", "标准问题:风险测评"), + ) + + with pytest.raises(KnowledgePublicationError, match="向量写入失败"): + await service.publish(records, reviewer_id=9) + + assert store.published == [] + assert store.disabled == [101, 102] + assert vectors.deleted == [ + ("fin_faq_collection", ("101",)), + ("fin_policy_collection", ("102",)), + ] + + +@pytest.mark.asyncio +async def test_invalid_embedding_dimension_blocks_all_database_and_vector_writes() -> None: + store = FakeStore() + vectors = FakeVectors() + service = KnowledgePublicationService(FakeEmbedder([0.1] * 512), store, vectors) + + with pytest.raises(KnowledgePublicationError, match="1024"): + await service.publish((Record("FAQ-001", "fin_faq_collection", "基金开户"),), reviewer_id=9) + + assert store.staged == [] + assert vectors.upserts == [] diff --git a/tests/unit/tools/test_publish_customer_service_knowledge.py b/tests/unit/tools/test_publish_customer_service_knowledge.py new file mode 100644 index 0000000..a5409a8 --- /dev/null +++ b/tests/unit/tools/test_publish_customer_service_knowledge.py @@ -0,0 +1,40 @@ +import pytest + +from tools.publish_customer_service_knowledge import load_pending_manifest + + +def manifest() -> dict[str, object]: + return { + "summary": {"eligible_records": 1, "publication_state": "pending_review"}, + "records": [{ + "qa_id": "FAQ-001", + "milvus_collection": "fin_faq_collection", + "retrieval_text": "标准问题:基金开户", + "title": "基金开户", + "snippet": "基金开户", + "tags": ["开户"], + "version": "v5.8", + "content_text": "{\"answer\":\"请在官方页面开户。\"}", + "source_file": "qa-v5.8.txt", + "effective_date": None, + "expire_date": None, + "review_status": "pending_review", + "status": "active", + }], + } + + +def test_pending_manifest_is_converted_to_a_publishable_administrator_payload() -> None: + records = load_pending_manifest(manifest()) + + assert len(records) == 1 + assert records[0].qa_id == "FAQ-001" + assert records[0].metadata["content_text"] == "{\"answer\":\"请在官方页面开户。\"}" + + +def test_manifest_that_claims_to_be_published_is_rejected() -> None: + invalid = manifest() + invalid["summary"] = {"eligible_records": 1, "publication_state": "published"} + + with pytest.raises(ValueError, match="pending_review"): + load_pending_manifest(invalid) diff --git a/tools/publish_customer_service_knowledge.py b/tools/publish_customer_service_knowledge.py new file mode 100644 index 0000000..374e6ad --- /dev/null +++ b/tools/publish_customer_service_knowledge.py @@ -0,0 +1,283 @@ +"""管理员显式批准后发布一期客服公开知识;默认只验证清单,不写外部服务。""" + +import argparse +import asyncio +import json +import sys +from collections.abc import Mapping +from datetime import UTC, datetime +from pathlib import Path +from typing import Any, cast + +from sqlalchemy import select, text, update + +# 直接执行 tools 脚本时优先解析当前工作树,避免误导入相邻 worktree 的 app 包。 +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from app.service.knowledge_publication_service import ( # noqa: E402 + KnowledgePublicationRecord, + KnowledgePublicationService, +) + +# 外部服务依赖仅在显式 --apply 时载入,dry-run 不需要本地 .env 或服务可达。 +SessionFactory: Any +InteractionAudit: Any +FinKnowledgeMeta: Any +DatabaseModelGateway: Any +get_settings: Any + + +class DatabaseKnowledgeEmbedder: + """发布工具只走专用知识向量端点,禁止意外使用聊天端点。""" + + async def embed(self, text_value: str) -> list[float]: + settings = get_settings() + endpoint_code = settings.knowledge_embedding_endpoint_code + if not endpoint_code: + raise RuntimeError("KNOWLEDGE_EMBEDDING_ENDPOINT_CODE 未配置") + return cast( + list[float], + await DatabaseModelGateway().embed( + endpoint_code=endpoint_code, + text=text_value, + timeout_ms=settings.knowledge_embedding_timeout_ms, + ), + ) + + +class SqlAlchemyKnowledgePublicationStore: + """使用现有知识表暂存、发布和停用记录,不变更数据库表结构。""" + + def __init__(self, reviewer_id: int) -> None: + self._reviewer_id = reviewer_id + + async def stage(self, records: tuple[KnowledgePublicationRecord, ...]) -> dict[str, int]: + now = datetime.now(UTC).replace(tzinfo=None) + async with SessionFactory() as session, session.begin(): + await self._assert_reviewer(session) + await self._reject_existing_qa_ids(session, records) + rows: list[Any] = [] + for record in records: + row = FinKnowledgeMeta( + knowledge_type=str(record.metadata["knowledge_type"]), + title=record.title, + source_file=str(record.metadata["source_file"]), + minio_path=None, + milvus_collection=record.milvus_collection, + version=record.version, + effective_date=record.metadata.get("effective_date"), + expire_date=record.metadata.get("expire_date"), + content_text=str(record.metadata["content_text"]), + tags=list(record.tags), + reviewer_id=None, + review_status="pending", + status="disabled", + created_at=now, + updated_at=now, + ) + session.add(row) + rows.append(row) + await session.flush() + session.add(InteractionAudit( + actor_type="admin", + actor_id=self._reviewer_id, + portal="admin", + action_type="knowledge.publication_staged", + detail={"qa_ids": [record.qa_id for record in records]}, + created_at=now, + )) + return {record.qa_id: int(row.id) for record, row in zip(records, rows, strict=True)} + + async def publish(self, knowledge_ids: tuple[int, ...], reviewer_id: int) -> None: + now = datetime.now(UTC).replace(tzinfo=None) + async with SessionFactory() as session, session.begin(): + await session.execute( + update(FinKnowledgeMeta) + .where(FinKnowledgeMeta.id.in_(knowledge_ids)) + .values( + reviewer_id=reviewer_id, + review_status="published", + status="active", + updated_at=now, + ) + ) + session.add(InteractionAudit( + actor_type="admin", + actor_id=reviewer_id, + portal="admin", + action_type="knowledge.publication_completed", + detail={"knowledge_ids": list(knowledge_ids)}, + created_at=now, + )) + + async def disable(self, knowledge_ids: tuple[int, ...]) -> None: + now = datetime.now(UTC).replace(tzinfo=None) + async with SessionFactory() as session, session.begin(): + await session.execute( + update(FinKnowledgeMeta) + .where(FinKnowledgeMeta.id.in_(knowledge_ids)) + .values(status="disabled", updated_at=now) + ) + session.add(InteractionAudit( + actor_type="system", + actor_id=None, + portal="admin", + action_type="knowledge.publication_failed", + detail={"knowledge_ids": list(knowledge_ids)}, + created_at=now, + )) + + async def _assert_reviewer(self, session: Any) -> None: + row = await session.execute( + text( + "SELECT id FROM sys_user " + "WHERE id = :reviewer_id AND status = 'active' " + "AND user_type IN ('employee', 'admin')" + ), + {"reviewer_id": self._reviewer_id}, + ) + if row.scalar_one_or_none() is None: + raise RuntimeError("reviewer_id 不是有效的在职管理员或员工账号") + + @staticmethod + async def _reject_existing_qa_ids( + session: Any, records: tuple[KnowledgePublicationRecord, ...] + ) -> None: + collections = tuple({record.milvus_collection for record in records}) + rows = await session.scalars( + select(FinKnowledgeMeta) + .where(FinKnowledgeMeta.milvus_collection.in_(collections)) + .with_for_update() + ) + existing_ids: set[str] = set() + for row in rows: + try: + content = json.loads(row.content_text) + except json.JSONDecodeError: + continue + qa_id = content.get("qa_id") if isinstance(content, dict) else None + if isinstance(qa_id, str): + existing_ids.add(qa_id) + duplicates = sorted(existing_ids & {record.qa_id for record in records}) + if duplicates: + raise RuntimeError(f"qa_id 已存在,拒绝重复发布: {', '.join(duplicates)}") + + +class MilvusKnowledgePublicationStore: + """管理员发布期的最小 Milvus 写适配器;客服 Agent 运行期仍只能检索。""" + + def __init__(self) -> None: + settings = get_settings() + self._uri = settings.milvus_uri + self._token = settings.milvus_token or None + self._client: Any | None = None + + async def _client_instance(self) -> Any: + if self._client is None: + from pymilvus import AsyncMilvusClient # type: ignore[import-untyped] + + self._client = AsyncMilvusClient(uri=self._uri, token=self._token) + return self._client + + async def upsert(self, collection: str, records: tuple[dict[str, object], ...]) -> None: + client = await self._client_instance() + payload = [ + {**record, "tags": json.dumps(record["tags"], ensure_ascii=False)} + for record in records + ] + await client.upsert(collection_name=collection, data=payload) + + async def delete(self, collection: str, knowledge_ids: tuple[str, ...]) -> None: + if not all(knowledge_id.isdecimal() for knowledge_id in knowledge_ids): + raise ValueError("knowledge_id 必须为十进制主键") + client = await self._client_instance() + values = ", ".join(json.dumps(knowledge_id) for knowledge_id in knowledge_ids) + await client.delete(collection_name=collection, filter=f"knowledge_id in [{values}]") + + +def load_pending_manifest(value: Mapping[str, object]) -> tuple[KnowledgePublicationRecord, ...]: + """只接受预检工具输出的 pending_review 清单,拒绝手工伪造已发布状态。""" + summary = value.get("summary") + records = value.get("records") + if not isinstance(summary, dict) or summary.get("publication_state") != "pending_review": + raise ValueError("发布清单必须处于 pending_review 状态") + if not isinstance(records, list) or summary.get("eligible_records") != len(records): + raise ValueError("发布清单记录数与汇总不一致") + result: list[KnowledgePublicationRecord] = [] + for raw in records: + if not isinstance(raw, dict): + raise ValueError("发布清单记录必须是对象") + required = ("qa_id", "milvus_collection", "retrieval_text", "title", "snippet", "version") + if any(not isinstance(raw.get(field), str) or not raw[field].strip() for field in required): + raise ValueError("发布清单缺少字符串字段") + tags = raw.get("tags") + if not isinstance(tags, list) or not all(isinstance(tag, str) and tag for tag in tags): + raise ValueError("发布清单标签无效") + if raw.get("review_status") != "pending_review" or raw.get("status") != "active": + raise ValueError("只有预检待审核记录可以发布") + result.append(KnowledgePublicationRecord( + qa_id=str(raw["qa_id"]), + milvus_collection=str(raw["milvus_collection"]), + retrieval_text=str(raw["retrieval_text"]), + title=str(raw["title"]), + snippet=str(raw["snippet"]), + tags=tuple(tags), + version=str(raw["version"]), + metadata=dict(raw), + )) + return tuple(result) + + +def _read_manifest(path: Path) -> tuple[KnowledgePublicationRecord, ...]: + value = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(value, dict): + raise ValueError("发布清单根节点必须是对象") + return load_pending_manifest(value) + + +async def _apply(records: tuple[KnowledgePublicationRecord, ...], reviewer_id: int) -> None: + global DatabaseModelGateway, FinKnowledgeMeta, InteractionAudit, SessionFactory, get_settings + from app.core.config import get_settings + from app.infrastructure.db import SessionFactory + from app.model.audit import InteractionAudit + from app.model.knowledge import FinKnowledgeMeta + from app.service.model_gateway import DatabaseModelGateway + + service = KnowledgePublicationService( + DatabaseKnowledgeEmbedder(), + SqlAlchemyKnowledgePublicationStore(reviewer_id), + MilvusKnowledgePublicationStore(), + ) + result = await service.publish(records, reviewer_id=reviewer_id) + print( + json.dumps( + {"published_records": len(result.knowledge_ids), "collections": result.collections}, + ensure_ascii=False, + ) + ) + + +def main() -> int: + """默认 dry-run,且 apply 必须三重确认,防止候选资料被误发布。""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--input", type=Path, required=True) + parser.add_argument("--reviewer-id", type=int) + parser.add_argument("--apply", action="store_true") + parser.add_argument("--confirm-count", type=int) + arguments = parser.parse_args() + records = _read_manifest(arguments.input) + if not arguments.apply: + print(f"DRY RUN: {len(records)} records are pending administrator review") + return 0 + if arguments.reviewer_id is None or arguments.reviewer_id <= 0: + raise SystemExit("--apply requires a positive --reviewer-id") + if arguments.confirm_count != len(records): + raise SystemExit("--apply requires --confirm-count equal to the manifest record count") + asyncio.run(_apply(records, arguments.reviewer_id)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main())