Files

150 lines
6.0 KiB
Python

"""知识库服务:CRUD + Token 管理 + 默认目录树。"""
from datetime import datetime, timedelta
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,
)
# AI 链接默认 30 分钟有效期
from datetime import datetime, timedelta
kb.token_expires_at = (datetime.now() + timedelta(minutes=30)).strftime("%Y-%m-%d %H:%M:%S")
self._seed_default_categories(kb.id)
self._session.commit()
return kb, token
def _seed_default_categories(self, kb_id: str) -> None:
"""创建默认目录树结构。"""
default_tree = [
("01 公司层", ["公司基本信息", "经营理念", "四大价值"]),
("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:
"""删除知识库 → 进回收站(3 天后自动彻底删除)。"""
kb = self.get_or_404(kb_id, user)
self._kb_repo.delete(kb)
kb.deleted_at = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
self._session.commit()
def regenerate_token(self, kb_id: str, user: User) -> tuple[KnowledgeBase, str]:
"""重新生成 Token。旧链接立即失效,新链接默认 30 分钟有效期。"""
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,
)
kb.token_expires_at = (datetime.now() + timedelta(minutes=30)).strftime("%Y-%m-%d %H:%M:%S")
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 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