|
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|
|
|
|
|
|
|
import asyncio
|
|
import asyncio
|
|
|
import hashlib
|
|
import hashlib
|
|
|
|
|
+import json
|
|
|
import re
|
|
import re
|
|
|
from datetime import UTC, datetime, timedelta
|
|
from datetime import UTC, datetime, timedelta
|
|
|
from typing import Any
|
|
from typing import Any
|
|
@@ -16,6 +17,9 @@ from wbb import BOT_PROFILE_ID, db
|
|
|
connectionsdb = db.business_assistant_connections
|
|
connectionsdb = db.business_assistant_connections
|
|
|
settingsdb = db.business_assistant_settings
|
|
settingsdb = db.business_assistant_settings
|
|
|
knowledgedb = db.business_assistant_knowledge
|
|
knowledgedb = db.business_assistant_knowledge
|
|
|
|
|
+knowledge_sourcesdb = db.business_assistant_knowledge_sources
|
|
|
|
|
+source_eventsdb = db.business_assistant_source_events
|
|
|
|
|
+knowledge_candidatesdb = db.business_assistant_knowledge_candidates
|
|
|
conversationsdb = db.business_assistant_conversations
|
|
conversationsdb = db.business_assistant_conversations
|
|
|
messagesdb = db.business_assistant_messages
|
|
messagesdb = db.business_assistant_messages
|
|
|
usagedb = db.business_assistant_usage
|
|
usagedb = db.business_assistant_usage
|
|
@@ -44,6 +48,10 @@ DEFAULT_ACCOUNT_SETTINGS: dict[str, Any] = {
|
|
|
|
|
|
|
|
VALID_NOTIFICATION_DESTINATIONS = {"owner", "ops", "both"}
|
|
VALID_NOTIFICATION_DESTINATIONS = {"owner", "ops", "both"}
|
|
|
VALID_TONES = {"professional", "friendly", "concise"}
|
|
VALID_TONES = {"professional", "friendly", "concise"}
|
|
|
|
|
+VALID_SOURCE_TYPES = {"channel", "group", "business"}
|
|
|
|
|
+VALID_PUBLICATION_MODES = {"review", "auto"}
|
|
|
|
|
+VALID_AUTHOR_POLICIES = {"admins_or_allowlist", "admins", "allowlist", "all"}
|
|
|
|
|
+VALID_CANDIDATE_STATUSES = {"pending", "published", "rejected", "stale"}
|
|
|
HANDOFF_TERMS = (
|
|
HANDOFF_TERMS = (
|
|
|
"人工",
|
|
"人工",
|
|
|
"真人",
|
|
"真人",
|
|
@@ -142,6 +150,23 @@ def _normalize_string_list(value: Any, *, max_items: int, max_length: int) -> li
|
|
|
return normalized
|
|
return normalized
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
+def _normalize_int_list(value: Any, *, max_items: int = 100) -> list[int]:
|
|
|
|
|
+ values = value if isinstance(value, (list, tuple, set)) else str(value or "").split(",")
|
|
|
|
|
+ normalized: list[int] = []
|
|
|
|
|
+ for item in values:
|
|
|
|
|
+ if item in {None, ""}:
|
|
|
|
|
+ continue
|
|
|
|
|
+ try:
|
|
|
|
|
+ parsed = int(item)
|
|
|
|
|
+ except (TypeError, ValueError) as exc:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "白名单用户 ID 必须是整数。") from exc
|
|
|
|
|
+ if parsed and parsed not in normalized:
|
|
|
|
|
+ normalized.append(parsed)
|
|
|
|
|
+ if len(normalized) > max_items:
|
|
|
|
|
+ raise AssistantDataError("too_many_items", f"最多允许 {max_items} 项。")
|
|
|
|
|
+ return normalized
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
async def ensure_assistant_indexes() -> None:
|
|
async def ensure_assistant_indexes() -> None:
|
|
|
global _indexes_ready
|
|
global _indexes_ready
|
|
|
if _indexes_ready:
|
|
if _indexes_ready:
|
|
@@ -161,6 +186,51 @@ async def ensure_assistant_indexes() -> None:
|
|
|
await knowledgedb.create_index(
|
|
await knowledgedb.create_index(
|
|
|
[("bot_id", ASCENDING), ("entry_id", ASCENDING)], unique=True
|
|
[("bot_id", ASCENDING), ("entry_id", ASCENDING)], unique=True
|
|
|
)
|
|
)
|
|
|
|
|
+ await knowledge_sourcesdb.create_index(
|
|
|
|
|
+ [("bot_id", ASCENDING), ("source_id", ASCENDING)], unique=True
|
|
|
|
|
+ )
|
|
|
|
|
+ await knowledge_sourcesdb.create_index(
|
|
|
|
|
+ [("bot_id", ASCENDING), ("source_key", ASCENDING)], unique=True
|
|
|
|
|
+ )
|
|
|
|
|
+ await knowledge_sourcesdb.create_index(
|
|
|
|
|
+ [
|
|
|
|
|
+ ("bot_id", ASCENDING),
|
|
|
|
|
+ ("connection_id", ASCENDING),
|
|
|
|
|
+ ("enabled", ASCENDING),
|
|
|
|
|
+ ("updated_at", DESCENDING),
|
|
|
|
|
+ ]
|
|
|
|
|
+ )
|
|
|
|
|
+ await source_eventsdb.create_index(
|
|
|
|
|
+ [("bot_id", ASCENDING), ("event_id", ASCENDING)], unique=True
|
|
|
|
|
+ )
|
|
|
|
|
+ await source_eventsdb.create_index(
|
|
|
|
|
+ [
|
|
|
|
|
+ ("bot_id", ASCENDING),
|
|
|
|
|
+ ("source_id", ASCENDING),
|
|
|
|
|
+ ("event_key", ASCENDING),
|
|
|
|
|
+ ],
|
|
|
|
|
+ unique=True,
|
|
|
|
|
+ )
|
|
|
|
|
+ await source_eventsdb.create_index("expires_at", expireAfterSeconds=0)
|
|
|
|
|
+ await knowledge_candidatesdb.create_index(
|
|
|
|
|
+ [("bot_id", ASCENDING), ("candidate_id", ASCENDING)], unique=True
|
|
|
|
|
+ )
|
|
|
|
|
+ await knowledge_candidatesdb.create_index(
|
|
|
|
|
+ [
|
|
|
|
|
+ ("bot_id", ASCENDING),
|
|
|
|
|
+ ("event_id", ASCENDING),
|
|
|
|
|
+ ("item_index", ASCENDING),
|
|
|
|
|
+ ],
|
|
|
|
|
+ unique=True,
|
|
|
|
|
+ )
|
|
|
|
|
+ await knowledge_candidatesdb.create_index(
|
|
|
|
|
+ [
|
|
|
|
|
+ ("bot_id", ASCENDING),
|
|
|
|
|
+ ("connection_id", ASCENDING),
|
|
|
|
|
+ ("status", ASCENDING),
|
|
|
|
|
+ ("updated_at", DESCENDING),
|
|
|
|
|
+ ]
|
|
|
|
|
+ )
|
|
|
await knowledgedb.create_index(
|
|
await knowledgedb.create_index(
|
|
|
[
|
|
[
|
|
|
("bot_id", ASCENDING),
|
|
("bot_id", ASCENDING),
|
|
@@ -432,6 +502,12 @@ async def create_knowledge_entry(
|
|
|
values.get("priority", 0), "知识优先级", minimum=-1000, maximum=1000
|
|
values.get("priority", 0), "知识优先级", minimum=-1000, maximum=1000
|
|
|
),
|
|
),
|
|
|
"enabled": bool(values.get("enabled", True)),
|
|
"enabled": bool(values.get("enabled", True)),
|
|
|
|
|
+ "source_candidate_id": clean_text(
|
|
|
|
|
+ values.get("source_candidate_id"), max_length=64
|
|
|
|
|
+ ),
|
|
|
|
|
+ "source_references": values.get("source_references")
|
|
|
|
|
+ if isinstance(values.get("source_references"), list)
|
|
|
|
|
+ else [],
|
|
|
"created_at": now,
|
|
"created_at": now,
|
|
|
"updated_at": now,
|
|
"updated_at": now,
|
|
|
}
|
|
}
|
|
@@ -467,6 +543,15 @@ async def update_knowledge_entry(entry_id: str, values: dict[str, Any]) -> dict[
|
|
|
)
|
|
)
|
|
|
if "enabled" in values:
|
|
if "enabled" in values:
|
|
|
update["enabled"] = bool(values.get("enabled"))
|
|
update["enabled"] = bool(values.get("enabled"))
|
|
|
|
|
+ if "source_candidate_id" in values:
|
|
|
|
|
+ update["source_candidate_id"] = clean_text(
|
|
|
|
|
+ values.get("source_candidate_id"), max_length=64
|
|
|
|
|
+ )
|
|
|
|
|
+ if "source_references" in values:
|
|
|
|
|
+ references = values.get("source_references")
|
|
|
|
|
+ if not isinstance(references, list):
|
|
|
|
|
+ raise AssistantDataError("invalid_knowledge", "知识来源引用必须是数组。")
|
|
|
|
|
+ update["source_references"] = references[:20]
|
|
|
if not update:
|
|
if not update:
|
|
|
raise AssistantDataError("unchanged", "没有可保存的知识条目字段。")
|
|
raise AssistantDataError("unchanged", "没有可保存的知识条目字段。")
|
|
|
update["updated_at"] = utc_now()
|
|
update["updated_at"] = utc_now()
|
|
@@ -510,6 +595,577 @@ async def list_knowledge_entries(
|
|
|
return [item async for item in cursor], total
|
|
return [item async for item in cursor], total
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
+def _knowledge_source_key(connection_id: str, source_type: str, chat_id: int) -> str:
|
|
|
|
|
+ return f"{connection_id}:{source_type}:{int(chat_id)}"
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def normalize_knowledge_source(
|
|
|
|
|
+ values: dict[str, Any], *, previous: dict[str, Any] | None = None
|
|
|
|
|
+) -> dict[str, Any]:
|
|
|
|
|
+ current = {**(previous or {}), **values}
|
|
|
|
|
+ source_type = str(current.get("source_type") or "")
|
|
|
|
|
+ if source_type not in VALID_SOURCE_TYPES:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "知识来源类型无效。")
|
|
|
|
|
+ try:
|
|
|
|
|
+ chat_id = int(current.get("chat_id") or 0)
|
|
|
|
|
+ except (TypeError, ValueError) as exc:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "来源聊天 ID 必须是整数。") from exc
|
|
|
|
|
+ if source_type == "business":
|
|
|
|
|
+ chat_id = 0
|
|
|
|
|
+ elif not chat_id:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "频道或群组来源必须填写聊天 ID。")
|
|
|
|
|
+ publication_mode = str(current.get("publication_mode") or "review")
|
|
|
|
|
+ if publication_mode not in VALID_PUBLICATION_MODES:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "知识发布方式无效。")
|
|
|
|
|
+ author_policy = str(current.get("author_policy") or "admins_or_allowlist")
|
|
|
|
|
+ if author_policy not in VALID_AUTHOR_POLICIES:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "来源作者策略无效。")
|
|
|
|
|
+ linked_chat_id = 0
|
|
|
|
|
+ try:
|
|
|
|
|
+ linked_chat_id = int(current.get("linked_chat_id") or 0)
|
|
|
|
|
+ except (TypeError, ValueError) as exc:
|
|
|
|
|
+ raise AssistantDataError("invalid_source", "关联讨论群 ID 必须是整数。") from exc
|
|
|
|
|
+ return {
|
|
|
|
|
+ "source_type": source_type,
|
|
|
|
|
+ "chat_id": chat_id,
|
|
|
|
|
+ "title": clean_text(
|
|
|
|
|
+ current.get("title") or ("Business 日常对话" if source_type == "business" else ""),
|
|
|
|
|
+ max_length=200,
|
|
|
|
|
+ required=True,
|
|
|
|
|
+ ),
|
|
|
|
|
+ "publication_mode": publication_mode,
|
|
|
|
|
+ "author_policy": author_policy,
|
|
|
|
|
+ "allowed_user_ids": _normalize_int_list(current.get("allowed_user_ids")),
|
|
|
|
|
+ "include_linked_chat": bool(current.get("include_linked_chat", False))
|
|
|
|
|
+ if source_type == "channel"
|
|
|
|
|
+ else False,
|
|
|
|
|
+ "linked_chat_id": linked_chat_id if source_type == "channel" else 0,
|
|
|
|
|
+ "linked_chat_title": clean_text(
|
|
|
|
|
+ current.get("linked_chat_title"), max_length=200
|
|
|
|
|
+ )
|
|
|
|
|
+ if source_type == "channel"
|
|
|
|
|
+ else "",
|
|
|
|
|
+ "backfill_days": _bounded_int(
|
|
|
|
|
+ current.get("backfill_days", 30), "历史回溯天数", minimum=0, maximum=365
|
|
|
|
|
+ ),
|
|
|
|
|
+ "backfill_limit": _bounded_int(
|
|
|
|
|
+ current.get("backfill_limit", 100), "历史回溯条数", minimum=1, maximum=500
|
|
|
|
|
+ ),
|
|
|
|
|
+ "enabled": bool(current.get("enabled", True)),
|
|
|
|
|
+ "access_status": clean_text(current.get("access_status"), max_length=64),
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def create_knowledge_source(
|
|
|
|
|
+ connection_id: str, values: dict[str, Any]
|
|
|
|
|
+) -> dict[str, Any]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ if not await get_business_connection(connection_id):
|
|
|
|
|
+ raise AssistantDataError("connection_not_found", "未找到 Business 连接。")
|
|
|
|
|
+ normalized = normalize_knowledge_source(values)
|
|
|
|
|
+ now = utc_now()
|
|
|
|
|
+ source = {
|
|
|
|
|
+ "bot_id": BOT_PROFILE_ID,
|
|
|
|
|
+ "source_id": uuid4().hex,
|
|
|
|
|
+ "source_key": _knowledge_source_key(
|
|
|
|
|
+ str(connection_id), normalized["source_type"], normalized["chat_id"]
|
|
|
|
|
+ ),
|
|
|
|
|
+ "connection_id": str(connection_id),
|
|
|
|
|
+ **normalized,
|
|
|
|
|
+ "sync_status": "idle",
|
|
|
|
|
+ "last_error": "",
|
|
|
|
|
+ "created_at": now,
|
|
|
|
|
+ "updated_at": now,
|
|
|
|
|
+ }
|
|
|
|
|
+ try:
|
|
|
|
|
+ await knowledge_sourcesdb.insert_one(source)
|
|
|
|
|
+ except DuplicateKeyError as exc:
|
|
|
|
|
+ raise AssistantDataError("source_exists", "该连接已配置相同的知识来源。") from exc
|
|
|
|
|
+ return source
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def get_knowledge_source(source_id: str) -> dict[str, Any] | None:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ return await knowledge_sourcesdb.find_one(_scope(source_id=str(source_id)))
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def update_knowledge_source(
|
|
|
|
|
+ source_id: str, values: dict[str, Any]
|
|
|
|
|
+) -> dict[str, Any]:
|
|
|
|
|
+ current = await get_knowledge_source(source_id)
|
|
|
|
|
+ if not current:
|
|
|
|
|
+ raise AssistantDataError("source_not_found", "未找到知识来源。")
|
|
|
|
|
+ for field in ("connection_id", "source_type", "chat_id"):
|
|
|
|
|
+ if field in values and str(values[field]) != str(current[field]):
|
|
|
|
|
+ raise AssistantDataError("immutable_source", "来源连接、类型和聊天 ID 不可修改。")
|
|
|
|
|
+ normalized = normalize_knowledge_source(values, previous=current)
|
|
|
|
|
+ await knowledge_sourcesdb.update_one(
|
|
|
|
|
+ _scope(source_id=str(source_id)),
|
|
|
|
|
+ {"$set": {**normalized, "updated_at": utc_now()}},
|
|
|
|
|
+ )
|
|
|
|
|
+ return await get_knowledge_source(source_id) or current
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def list_knowledge_sources(
|
|
|
|
|
+ connection_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ page: int = 1,
|
|
|
|
|
+ page_size: int = 20,
|
|
|
|
|
+) -> tuple[list[dict[str, Any]], int]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ filters = _scope(connection_id=str(connection_id))
|
|
|
|
|
+ page = max(1, int(page))
|
|
|
|
|
+ page_size = max(1, min(int(page_size), 100))
|
|
|
|
|
+ total = await knowledge_sourcesdb.count_documents(filters)
|
|
|
|
|
+ cursor = (
|
|
|
|
|
+ knowledge_sourcesdb.find(filters)
|
|
|
|
|
+ .sort("updated_at", DESCENDING)
|
|
|
|
|
+ .skip((page - 1) * page_size)
|
|
|
|
|
+ .limit(page_size)
|
|
|
|
|
+ )
|
|
|
|
|
+ return [item async for item in cursor], total
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def list_sources_for_chat(
|
|
|
|
|
+ chat_id: int, *, enabled_only: bool = True
|
|
|
|
|
+) -> list[dict[str, Any]]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ values: dict[str, Any] = {
|
|
|
|
|
+ "$or": [
|
|
|
|
|
+ {"chat_id": int(chat_id)},
|
|
|
|
|
+ {"include_linked_chat": True, "linked_chat_id": int(chat_id)},
|
|
|
|
|
+ ]
|
|
|
|
|
+ }
|
|
|
|
|
+ if enabled_only:
|
|
|
|
|
+ values["enabled"] = True
|
|
|
|
|
+ filters = _scope(values)
|
|
|
|
|
+ return [item async for item in knowledge_sourcesdb.find(filters)]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def list_enabled_sources_for_chat(chat_id: int) -> list[dict[str, Any]]:
|
|
|
|
|
+ return await list_sources_for_chat(chat_id, enabled_only=True)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def list_enabled_business_sources(connection_id: str) -> list[dict[str, Any]]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ return [
|
|
|
|
|
+ item
|
|
|
|
|
+ async for item in knowledge_sourcesdb.find(
|
|
|
|
|
+ _scope(
|
|
|
|
|
+ connection_id=str(connection_id),
|
|
|
|
|
+ source_type="business",
|
|
|
|
|
+ enabled=True,
|
|
|
|
|
+ )
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def list_business_sources(connection_id: str) -> list[dict[str, Any]]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ return [
|
|
|
|
|
+ item
|
|
|
|
|
+ async for item in knowledge_sourcesdb.find(
|
|
|
|
|
+ _scope(connection_id=str(connection_id), source_type="business")
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def update_knowledge_source_sync(
|
|
|
|
|
+ source_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ status: str,
|
|
|
|
|
+ error: str = "",
|
|
|
|
|
+ processed: int | None = None,
|
|
|
|
|
+) -> dict[str, Any]:
|
|
|
|
|
+ values: dict[str, Any] = {
|
|
|
|
|
+ "sync_status": clean_text(status, max_length=32, required=True),
|
|
|
|
|
+ "last_error": clean_text(error, max_length=1000),
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ if status == "running":
|
|
|
|
|
+ values["last_sync_started_at"] = utc_now()
|
|
|
|
|
+ if status in {"completed", "failed"}:
|
|
|
|
|
+ values["last_sync_finished_at"] = utc_now()
|
|
|
|
|
+ if processed is not None:
|
|
|
|
|
+ values["last_sync_processed"] = max(0, int(processed))
|
|
|
|
|
+ result = await knowledge_sourcesdb.update_one(
|
|
|
|
|
+ _scope(source_id=str(source_id)), {"$set": values}
|
|
|
|
|
+ )
|
|
|
|
|
+ if not result.matched_count:
|
|
|
|
|
+ raise AssistantDataError("source_not_found", "未找到知识来源。")
|
|
|
|
|
+ return await get_knowledge_source(source_id) or {}
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def touch_knowledge_source(source_id: str, *, error: str = "") -> None:
|
|
|
|
|
+ await knowledge_sourcesdb.update_one(
|
|
|
|
|
+ _scope(source_id=str(source_id)),
|
|
|
|
|
+ {
|
|
|
|
|
+ "$set": {
|
|
|
|
|
+ "last_event_at": utc_now(),
|
|
|
|
|
+ "last_error": clean_text(error, max_length=1000),
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ },
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def _disable_candidate_entries(filters: dict[str, Any]) -> None:
|
|
|
|
|
+ entry_ids = [
|
|
|
|
|
+ str(item.get("knowledge_entry_id") or "")
|
|
|
|
|
+ async for item in knowledge_candidatesdb.find(filters)
|
|
|
|
|
+ if item.get("knowledge_entry_id")
|
|
|
|
|
+ ]
|
|
|
|
|
+ if entry_ids:
|
|
|
|
|
+ await knowledgedb.update_many(
|
|
|
|
|
+ _scope({"entry_id": {"$in": entry_ids}}),
|
|
|
|
|
+ {"$set": {"enabled": False, "updated_at": utc_now()}},
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def delete_knowledge_source(source_id: str) -> None:
|
|
|
|
|
+ source = await get_knowledge_source(source_id)
|
|
|
|
|
+ if not source:
|
|
|
|
|
+ raise AssistantDataError("source_not_found", "未找到知识来源。")
|
|
|
|
|
+ candidate_filters = _scope(source_id=str(source_id))
|
|
|
|
|
+ await _disable_candidate_entries(candidate_filters)
|
|
|
|
|
+ await knowledge_candidatesdb.update_many(
|
|
|
|
|
+ candidate_filters,
|
|
|
|
|
+ {"$set": {"status": "stale", "stale_reason": "source_deleted", "updated_at": utc_now()}},
|
|
|
|
|
+ )
|
|
|
|
|
+ await source_eventsdb.update_many(
|
|
|
|
|
+ _scope(source_id=str(source_id)),
|
|
|
|
|
+ {"$set": {"deleted": True, "delete_reason": "source_deleted", "updated_at": utc_now()}},
|
|
|
|
|
+ )
|
|
|
|
|
+ await knowledge_sourcesdb.delete_one(_scope(source_id=str(source_id)))
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def source_event_id_for(source_id: str, event_key: str) -> str:
|
|
|
|
|
+ digest = hashlib.sha256(
|
|
|
|
|
+ f"{BOT_PROFILE_ID}:{source_id}:{event_key}".encode()
|
|
|
|
|
+ ).hexdigest()
|
|
|
|
|
+ return digest[:32]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def record_source_event(
|
|
|
|
|
+ source: dict[str, Any],
|
|
|
|
|
+ *,
|
|
|
|
|
+ event_key: str,
|
|
|
|
|
+ content: str,
|
|
|
|
|
+ metadata: dict[str, Any] | None = None,
|
|
|
|
|
+) -> tuple[dict[str, Any], bool]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ cleaned_content = clean_text(
|
|
|
|
|
+ content, max_length=20000, required=True, preserve_lines=True
|
|
|
|
|
+ )
|
|
|
|
|
+ cleaned_key = clean_text(event_key, max_length=300, required=True)
|
|
|
|
|
+ event_metadata = metadata if isinstance(metadata, dict) else {}
|
|
|
|
|
+ content_hash = hashlib.sha256(
|
|
|
|
|
+ json.dumps(
|
|
|
|
|
+ {"content": cleaned_content, "metadata": event_metadata},
|
|
|
|
|
+ ensure_ascii=False,
|
|
|
|
|
+ sort_keys=True,
|
|
|
|
|
+ default=str,
|
|
|
|
|
+ ).encode()
|
|
|
|
|
+ ).hexdigest()
|
|
|
|
|
+ filters = _scope(source_id=str(source["source_id"]), event_key=cleaned_key)
|
|
|
|
|
+ current = await source_eventsdb.find_one(filters)
|
|
|
|
|
+ if current and current.get("content_hash") == content_hash and not current.get("deleted"):
|
|
|
|
|
+ return current, False
|
|
|
|
|
+ now = utc_now()
|
|
|
|
|
+ event_id = str(current.get("event_id")) if current else source_event_id_for(
|
|
|
|
|
+ str(source["source_id"]), cleaned_key
|
|
|
|
|
+ )
|
|
|
|
|
+ version = int(current.get("version") or 0) + 1 if current else 1
|
|
|
|
|
+ document = {
|
|
|
|
|
+ "bot_id": BOT_PROFILE_ID,
|
|
|
|
|
+ "event_id": event_id,
|
|
|
|
|
+ "source_id": str(source["source_id"]),
|
|
|
|
|
+ "connection_id": str(source["connection_id"]),
|
|
|
|
|
+ "event_key": cleaned_key,
|
|
|
|
|
+ "content": cleaned_content,
|
|
|
|
|
+ "content_hash": content_hash,
|
|
|
|
|
+ "metadata": event_metadata,
|
|
|
|
|
+ "version": version,
|
|
|
|
|
+ "deleted": False,
|
|
|
|
|
+ "extraction_status": "pending",
|
|
|
|
|
+ "extraction_error": "",
|
|
|
|
|
+ "updated_at": now,
|
|
|
|
|
+ "expires_at": now + timedelta(days=30),
|
|
|
|
|
+ }
|
|
|
|
|
+ await source_eventsdb.update_one(
|
|
|
|
|
+ filters,
|
|
|
|
|
+ {"$set": document, "$setOnInsert": {"created_at": now}},
|
|
|
|
|
+ upsert=True,
|
|
|
|
|
+ )
|
|
|
|
|
+ return await source_eventsdb.find_one(filters) or document, True
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def get_source_event(event_id: str) -> dict[str, Any] | None:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ return await source_eventsdb.find_one(_scope(event_id=str(event_id)))
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def get_source_event_by_key(
|
|
|
|
|
+ source_id: str, event_key: str
|
|
|
|
|
+) -> dict[str, Any] | None:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ return await source_eventsdb.find_one(
|
|
|
|
|
+ _scope(source_id=str(source_id), event_key=str(event_key))
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def update_source_event_extraction(
|
|
|
|
|
+ event_id: str, *, status: str, error: str = "", item_count: int = 0
|
|
|
|
|
+) -> None:
|
|
|
|
|
+ await source_eventsdb.update_one(
|
|
|
|
|
+ _scope(event_id=str(event_id)),
|
|
|
|
|
+ {
|
|
|
|
|
+ "$set": {
|
|
|
|
|
+ "extraction_status": clean_text(status, max_length=32, required=True),
|
|
|
|
|
+ "extraction_error": clean_text(error, max_length=1000),
|
|
|
|
|
+ "extracted_item_count": max(0, int(item_count)),
|
|
|
|
|
+ "extracted_at": utc_now(),
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ },
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def invalidate_source_event_candidates(
|
|
|
|
|
+ event_id: str, *, reason: str
|
|
|
|
|
+) -> None:
|
|
|
|
|
+ filters = _scope(event_id=str(event_id))
|
|
|
|
|
+ await _disable_candidate_entries(filters)
|
|
|
|
|
+ await knowledge_candidatesdb.update_many(
|
|
|
|
|
+ filters,
|
|
|
|
|
+ {
|
|
|
|
|
+ "$set": {
|
|
|
|
|
+ "status": "stale",
|
|
|
|
|
+ "stale_reason": clean_text(reason, max_length=100),
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ },
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def mark_source_event_deleted(source_id: str, event_key: str) -> bool:
|
|
|
|
|
+ event = await source_eventsdb.find_one(
|
|
|
|
|
+ _scope(source_id=str(source_id), event_key=str(event_key))
|
|
|
|
|
+ )
|
|
|
|
|
+ if not event:
|
|
|
|
|
+ return False
|
|
|
|
|
+ await source_eventsdb.update_one(
|
|
|
|
|
+ _scope(event_id=str(event["event_id"])),
|
|
|
|
|
+ {"$set": {"deleted": True, "updated_at": utc_now()}},
|
|
|
|
|
+ )
|
|
|
|
|
+ await invalidate_source_event_candidates(
|
|
|
|
|
+ str(event["event_id"]), reason="source_message_deleted"
|
|
|
|
|
+ )
|
|
|
|
|
+ return True
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _normalize_candidate_item(item: dict[str, Any]) -> dict[str, Any]:
|
|
|
|
|
+ try:
|
|
|
|
|
+ confidence = float(item.get("confidence") or 0)
|
|
|
|
|
+ except (TypeError, ValueError):
|
|
|
|
|
+ confidence = 0
|
|
|
|
|
+ return {
|
|
|
|
|
+ "question": clean_text(item.get("question"), max_length=300, required=True),
|
|
|
|
|
+ "aliases": _normalize_string_list(item.get("aliases"), max_items=20, max_length=200),
|
|
|
|
|
+ "keywords": _normalize_string_list(item.get("keywords"), max_items=30, max_length=50),
|
|
|
|
|
+ "answer": clean_text(
|
|
|
|
|
+ item.get("answer"), max_length=4000, required=True, preserve_lines=True
|
|
|
|
|
+ ),
|
|
|
|
|
+ "tags": _normalize_string_list(item.get("tags"), max_items=20, max_length=30),
|
|
|
|
|
+ "confidence": max(0.0, min(confidence, 1.0)),
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def replace_event_candidates(
|
|
|
|
|
+ source: dict[str, Any],
|
|
|
|
|
+ event: dict[str, Any],
|
|
|
|
|
+ items: list[dict[str, Any]],
|
|
|
|
|
+ *,
|
|
|
|
|
+ auto_publish: bool,
|
|
|
|
|
+) -> list[dict[str, Any]]:
|
|
|
|
|
+ await invalidate_source_event_candidates(
|
|
|
|
|
+ str(event["event_id"]), reason="source_message_changed"
|
|
|
|
|
+ )
|
|
|
|
|
+ existing = {
|
|
|
|
|
+ int(item.get("item_index") or 0): item
|
|
|
|
|
+ async for item in knowledge_candidatesdb.find(
|
|
|
|
|
+ _scope(event_id=str(event["event_id"]))
|
|
|
|
|
+ )
|
|
|
|
|
+ }
|
|
|
|
|
+ results: list[dict[str, Any]] = []
|
|
|
|
|
+ metadata = event.get("metadata") if isinstance(event.get("metadata"), dict) else {}
|
|
|
|
|
+ for index, raw_item in enumerate(items[:5]):
|
|
|
|
|
+ normalized = _normalize_candidate_item(raw_item)
|
|
|
|
|
+ current = existing.get(index) or {}
|
|
|
|
|
+ now = utc_now()
|
|
|
|
|
+ candidate_id = str(current.get("candidate_id") or uuid4().hex)
|
|
|
|
|
+ candidate = {
|
|
|
|
|
+ "bot_id": BOT_PROFILE_ID,
|
|
|
|
|
+ "candidate_id": candidate_id,
|
|
|
|
|
+ "connection_id": str(source["connection_id"]),
|
|
|
|
|
+ "source_id": str(source["source_id"]),
|
|
|
|
|
+ "event_id": str(event["event_id"]),
|
|
|
|
|
+ "item_index": index,
|
|
|
|
|
+ **normalized,
|
|
|
|
|
+ "status": "pending",
|
|
|
|
|
+ "stale_reason": "",
|
|
|
|
|
+ "knowledge_entry_id": str(current.get("knowledge_entry_id") or ""),
|
|
|
|
|
+ "source_snapshot": {
|
|
|
|
|
+ "title": str(source.get("title") or ""),
|
|
|
|
|
+ "source_type": str(source.get("source_type") or ""),
|
|
|
|
|
+ "chat_id": int(metadata.get("chat_id") or source.get("chat_id") or 0),
|
|
|
|
|
+ "message_id": int(metadata.get("message_id") or 0),
|
|
|
|
|
+ "author_id": int(metadata.get("author_id") or 0),
|
|
|
|
|
+ },
|
|
|
|
|
+ "updated_at": now,
|
|
|
|
|
+ }
|
|
|
|
|
+ await knowledge_candidatesdb.update_one(
|
|
|
|
|
+ _scope(event_id=str(event["event_id"]), item_index=index),
|
|
|
|
|
+ {"$set": candidate, "$setOnInsert": {"created_at": now}},
|
|
|
|
|
+ upsert=True,
|
|
|
|
|
+ )
|
|
|
|
|
+ if auto_publish:
|
|
|
|
|
+ candidate = await publish_knowledge_candidate(candidate_id)
|
|
|
|
|
+ else:
|
|
|
|
|
+ candidate = await get_knowledge_candidate(candidate_id) or candidate
|
|
|
|
|
+ results.append(candidate)
|
|
|
|
|
+ await update_source_event_extraction(
|
|
|
|
|
+ str(event["event_id"]), status="completed", item_count=len(results)
|
|
|
|
|
+ )
|
|
|
|
|
+ return results
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def get_knowledge_candidate(candidate_id: str) -> dict[str, Any] | None:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ return await knowledge_candidatesdb.find_one(
|
|
|
|
|
+ _scope(candidate_id=str(candidate_id))
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def publish_knowledge_candidate(candidate_id: str) -> dict[str, Any]:
|
|
|
|
|
+ candidate = await get_knowledge_candidate(candidate_id)
|
|
|
|
|
+ if not candidate:
|
|
|
|
|
+ raise AssistantDataError("candidate_not_found", "未找到知识候选。")
|
|
|
|
|
+ if candidate.get("status") == "stale":
|
|
|
|
|
+ raise AssistantDataError("candidate_stale", "来源已变更或删除,不能发布该候选。")
|
|
|
|
|
+ if not await get_knowledge_source(str(candidate.get("source_id") or "")):
|
|
|
|
|
+ await knowledge_candidatesdb.update_one(
|
|
|
|
|
+ _scope(candidate_id=str(candidate_id)),
|
|
|
|
|
+ {
|
|
|
|
|
+ "$set": {
|
|
|
|
|
+ "status": "stale",
|
|
|
|
|
+ "stale_reason": "source_deleted",
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ },
|
|
|
|
|
+ )
|
|
|
|
|
+ raise AssistantDataError("source_not_found", "知识来源已删除,不能发布该候选。")
|
|
|
|
|
+ reference = {
|
|
|
|
|
+ **(candidate.get("source_snapshot") or {}),
|
|
|
|
|
+ "source_id": str(candidate.get("source_id") or ""),
|
|
|
|
|
+ "event_id": str(candidate.get("event_id") or ""),
|
|
|
|
|
+ }
|
|
|
|
|
+ values = {
|
|
|
|
|
+ key: candidate.get(key)
|
|
|
|
|
+ for key in ("question", "aliases", "keywords", "answer", "tags")
|
|
|
|
|
+ }
|
|
|
|
|
+ values.update(
|
|
|
|
|
+ {
|
|
|
|
|
+ "priority": round(float(candidate.get("confidence") or 0) * 100),
|
|
|
|
|
+ "enabled": True,
|
|
|
|
|
+ "source_candidate_id": str(candidate_id),
|
|
|
|
|
+ "source_references": [reference],
|
|
|
|
|
+ }
|
|
|
|
|
+ )
|
|
|
|
|
+ entry_id = str(candidate.get("knowledge_entry_id") or "")
|
|
|
|
|
+ if entry_id:
|
|
|
|
|
+ try:
|
|
|
|
|
+ entry = await update_knowledge_entry(entry_id, values)
|
|
|
|
|
+ except AssistantDataError as exc:
|
|
|
|
|
+ if exc.code != "knowledge_not_found":
|
|
|
|
|
+ raise
|
|
|
|
|
+ entry = await create_knowledge_entry(str(candidate["connection_id"]), values)
|
|
|
|
|
+ else:
|
|
|
|
|
+ entry = await create_knowledge_entry(str(candidate["connection_id"]), values)
|
|
|
|
|
+ await knowledge_candidatesdb.update_one(
|
|
|
|
|
+ _scope(candidate_id=str(candidate_id)),
|
|
|
|
|
+ {
|
|
|
|
|
+ "$set": {
|
|
|
|
|
+ "status": "published",
|
|
|
|
|
+ "knowledge_entry_id": str(entry["entry_id"]),
|
|
|
|
|
+ "published_at": utc_now(),
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ },
|
|
|
|
|
+ )
|
|
|
|
|
+ return await get_knowledge_candidate(candidate_id) or candidate
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def reject_knowledge_candidate(candidate_id: str) -> dict[str, Any]:
|
|
|
|
|
+ candidate = await get_knowledge_candidate(candidate_id)
|
|
|
|
|
+ if not candidate:
|
|
|
|
|
+ raise AssistantDataError("candidate_not_found", "未找到知识候选。")
|
|
|
|
|
+ entry_id = str(candidate.get("knowledge_entry_id") or "")
|
|
|
|
|
+ if entry_id:
|
|
|
|
|
+ await knowledgedb.update_one(
|
|
|
|
|
+ _scope(entry_id=entry_id),
|
|
|
|
|
+ {"$set": {"enabled": False, "updated_at": utc_now()}},
|
|
|
|
|
+ )
|
|
|
|
|
+ await knowledge_candidatesdb.update_one(
|
|
|
|
|
+ _scope(candidate_id=str(candidate_id)),
|
|
|
|
|
+ {
|
|
|
|
|
+ "$set": {
|
|
|
|
|
+ "status": "rejected",
|
|
|
|
|
+ "rejected_at": utc_now(),
|
|
|
|
|
+ "updated_at": utc_now(),
|
|
|
|
|
+ }
|
|
|
|
|
+ },
|
|
|
|
|
+ )
|
|
|
|
|
+ return await get_knowledge_candidate(candidate_id) or candidate
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def list_knowledge_candidates(
|
|
|
|
|
+ connection_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ status: str = "",
|
|
|
|
|
+ source_id: str = "",
|
|
|
|
|
+ query: str = "",
|
|
|
|
|
+ page: int = 1,
|
|
|
|
|
+ page_size: int = 20,
|
|
|
|
|
+) -> tuple[list[dict[str, Any]], int]:
|
|
|
|
|
+ await ensure_assistant_indexes()
|
|
|
|
|
+ filters: dict[str, Any] = _scope(connection_id=str(connection_id))
|
|
|
|
|
+ if status:
|
|
|
|
|
+ if status not in VALID_CANDIDATE_STATUSES:
|
|
|
|
|
+ raise AssistantDataError("invalid_status", "知识候选状态无效。")
|
|
|
|
|
+ filters["status"] = status
|
|
|
|
|
+ if source_id:
|
|
|
|
|
+ filters["source_id"] = str(source_id)
|
|
|
|
|
+ normalized_query = clean_text(query, max_length=100)
|
|
|
|
|
+ if normalized_query:
|
|
|
|
|
+ pattern = re.escape(normalized_query)
|
|
|
|
|
+ filters["$or"] = [
|
|
|
|
|
+ {"question": {"$regex": pattern, "$options": "i"}},
|
|
|
|
|
+ {"answer": {"$regex": pattern, "$options": "i"}},
|
|
|
|
|
+ {"keywords": {"$regex": pattern, "$options": "i"}},
|
|
|
|
|
+ ]
|
|
|
|
|
+ page = max(1, int(page))
|
|
|
|
|
+ page_size = max(1, min(int(page_size), 100))
|
|
|
|
|
+ total = await knowledge_candidatesdb.count_documents(filters)
|
|
|
|
|
+ cursor = (
|
|
|
|
|
+ knowledge_candidatesdb.find(filters)
|
|
|
|
|
+ .sort("updated_at", DESCENDING)
|
|
|
|
|
+ .skip((page - 1) * page_size)
|
|
|
|
|
+ .limit(page_size)
|
|
|
|
|
+ )
|
|
|
|
|
+ return [item async for item in cursor], total
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
def _search_text(value: str) -> str:
|
|
def _search_text(value: str) -> str:
|
|
|
return re.sub(r"[^\w\u3400-\u9fff]+", "", value.casefold())
|
|
return re.sub(r"[^\w\u3400-\u9fff]+", "", value.casefold())
|
|
|
|
|
|