"""知识库服务:CRUD + Token 管理 + 默认目录树。""" from app.core.errors import NotFoundError, PermissionDeniedError from app.core.logging import get_logger from app.models.document_category import DocumentCategory from app.models.knowledge_base import KnowledgeBase from app.models.user import User from app.repositories.kb_repo import KnowledgeBaseRepository from app.services.token_service import TokenService from sqlalchemy.orm import Session logger = get_logger(__name__) class KnowledgeBaseService: def __init__(self, session: Session) -> None: self._session = session self._kb_repo = KnowledgeBaseRepository(session) self._token_svc = TokenService() def create(self, user: User, name: str, description: str | None) -> tuple[KnowledgeBase, str]: """创建知识库。返回 (kb, full_token)。 full_token 仅此一次返回,用于构建完整 AI URL。 自动创建默认目录树结构。 """ token, token_hash, token_encrypted, token_hint = self._token_svc.create_token_pair() kb = self._kb_repo.create( user_id=user.id, name=name, description=description, token_hash=token_hash, token_encrypted=token_encrypted, token_hint=token_hint, ) self._seed_default_categories(kb.id) self._session.commit() return kb, token def _seed_default_categories(self, kb_id: str) -> None: """创建默认目录树结构。""" default_tree = [ ("01 公司层", ["公司基本信息", "经营理念", "四大价值", "A/M/B三态"]), ("02 战略层", ["公司战略", "客户战略", "AI战略", "产品战略"]), ("03 部门层", ["企划", "技术", "交付", "市场"]), ("04 岗位/AI角色", []), ("05 业务知识", ["价值创造", "价值传递", "价值交付", "价值支持"]), ("06 资产库", ["文案", "SOP", "模板", "案例", "Prompt"]), ] for i, (folder_name, children) in enumerate(default_tree): # 创建顶层文件夹 folder = DocumentCategory( knowledge_base_id=kb_id, parent_id=None, name=folder_name, path=f"/{folder_name}/", is_folder=True, sort_order=i, ) self._session.add(folder) self._session.flush() # 创建子分类 for j, child_name in enumerate(children): child = DocumentCategory( knowledge_base_id=kb_id, parent_id=folder.id, name=child_name, path=f"/{folder_name}/{child_name}/", is_folder=False, sort_order=j, ) self._session.add(child) def get_or_404(self, kb_id: str, user: User) -> KnowledgeBase: """获取知识库,校验所有权。不存在或无权 → 404。""" kb = self._kb_repo.get_by_id(kb_id) if kb is None or kb.user_id != user.id or kb.status == "DELETED": raise NotFoundError("知识库不存在。") return kb def list_by_user(self, user: User, *, page: int = 1, page_size: int = 20): """分页列表。""" return self._kb_repo.list_by_user(user.id, page=page, page_size=page_size) def update(self, kb_id: str, user: User, name: str | None, description: str | None) -> KnowledgeBase: kb = self.get_or_404(kb_id, user) update_fields = {} if name is not None: update_fields["name"] = name if description is not None: update_fields["description"] = description if update_fields: self._kb_repo.update(kb, **update_fields) self._session.commit() return kb def delete(self, kb_id: str, user: User) -> None: kb = self.get_or_404(kb_id, user) self._kb_repo.delete(kb) self._cleanup_kb_files(kb) self._session.commit() def _cleanup_kb_files(self, kb: KnowledgeBase) -> None: """物理删除知识库下所有文档的原始文件与 Markdown 文件(软删后调用)。 文件删除失败不阻塞删除流程(记录日志,可由后续清理兜底)。 """ from sqlalchemy import select from app.models.document import Document from app.storage.local_storage import get_storage stmt = select(Document).where(Document.knowledge_base_id == kb.id) docs = list(self._session.scalars(stmt).all()) storage = get_storage() removed = 0 for doc in docs: for key in (doc.storage_path, doc.markdown_path): if not key: continue try: storage.delete(key) removed += 1 except Exception as exc: # noqa: BLE001 logger.warning("清理文件失败 key=%s: %s", key, exc) if removed: logger.info("KB %s 软删,已物理清理 %d 个文件", kb.id, removed) def regenerate_token(self, kb_id: str, user: User) -> tuple[KnowledgeBase, str]: """重新生成 Token。旧链接立即失效。""" kb = self.get_or_404(kb_id, user) token, token_hash, token_encrypted, token_hint = self._token_svc.create_token_pair() self._kb_repo.update( kb, token_hash=token_hash, token_encrypted=token_encrypted, token_hint=token_hint, ) self._session.commit() return kb, token def set_enabled(self, kb_id: str, user: User, enabled: bool) -> KnowledgeBase: kb = self.get_or_404(kb_id, user) self._kb_repo.update(kb, enabled=enabled) self._session.commit() return kb def set_expiry(self, kb_id: str, user: User, expires_in_minutes: int | None) -> KnowledgeBase: """设置链接有效期。None = 长期有效;负数表示已过期(测试用)。""" from datetime import datetime, timedelta kb = self.get_or_404(kb_id, user) if expires_in_minutes is None: kb.token_expires_at = None else: kb.token_expires_at = (datetime.now() + timedelta(minutes=expires_in_minutes)).strftime( "%Y-%m-%d %H:%M:%S" ) self._session.commit() return kb def get_full_token(self, kb: KnowledgeBase) -> str | None: """解密 token 原文(供后台显示完整链接)。""" if kb.token_encrypted: return self._token_svc.decrypt_token(kb.token_encrypted) return None