305 lines
11 KiB
Python
305 lines
11 KiB
Python
"""文档服务:上传、删除、配额管理。"""
|
||
|
||
import hashlib
|
||
from pathlib import Path
|
||
|
||
import filetype
|
||
|
||
from app.core.config import get_settings
|
||
from app.core.errors import (
|
||
FileTooLargeError,
|
||
FileTypeUnsupportedError,
|
||
NotFoundError,
|
||
StorageQuotaExceededError,
|
||
)
|
||
from app.core.security import decrypt_token, encrypt_token, generate_token, hash_token
|
||
from app.models.document import Document
|
||
from app.models.user import User
|
||
from app.repositories.doc_repo import DocumentRepository
|
||
from app.repositories.kb_repo import KnowledgeBaseRepository
|
||
from app.storage.local_storage import get_storage
|
||
from app.storage.object_keys import original_object_key
|
||
from sqlalchemy.orm import Session
|
||
|
||
# 允许的文件扩展名
|
||
ALLOWED_EXTENSIONS = {".docx", ".pdf"}
|
||
ALLOWED_MIMES = {
|
||
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||
"application/pdf",
|
||
}
|
||
|
||
|
||
class DocumentService:
|
||
def __init__(self, session: Session) -> None:
|
||
self._session = session
|
||
self._doc_repo = DocumentRepository(session)
|
||
self._kb_repo = KnowledgeBaseRepository(session)
|
||
|
||
def upload(
|
||
self,
|
||
user: User,
|
||
kb_id: str,
|
||
filename: str,
|
||
content: bytes,
|
||
category_id: str | None = None,
|
||
) -> Document:
|
||
"""上传文档。
|
||
|
||
流程:校验 KB 归属 → 文件大小 → 扩展名/MIME → 配额 → SHA256 → 存储 → 入库
|
||
"""
|
||
# 校验 KB 归属
|
||
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("知识库不存在。")
|
||
|
||
settings = get_settings()
|
||
|
||
# 文件大小
|
||
if len(content) > settings.default_max_file_size:
|
||
raise FileTooLargeError(
|
||
f"单文件大小超出限制(最大 {settings.default_max_file_size // (1024*1024)}MB)。"
|
||
)
|
||
|
||
# 扩展名
|
||
ext = Path(filename).suffix.lower()
|
||
if ext not in ALLOWED_EXTENSIONS:
|
||
raise FileTypeUnsupportedError(
|
||
f"不支持的文件类型 '{ext}'。当前支持:{', '.join(sorted(ALLOWED_EXTENSIONS))}"
|
||
)
|
||
|
||
# MIME 嗅探(前 8KB)——仅用于辅助校验,扩展名为主
|
||
kind = filetype.guess(content[:8192])
|
||
detected_mime = kind.mime if kind else ""
|
||
# .docx 底层是 ZIP,filetype 会识别为 application/zip,这是正常的
|
||
# 只在检测到明确不属于文档类型的 MIME 时才拒绝(如图片、视频等)
|
||
BLOCKED_MIMES = {"image/jpeg", "image/png", "video/mp4", "audio/mpeg", "application/x-executable"}
|
||
if detected_mime in BLOCKED_MIMES:
|
||
raise FileTypeUnsupportedError(f"文件内容类型不受支持:{detected_mime}")
|
||
|
||
# 配额检查(原子 SQL)
|
||
file_size = len(content)
|
||
self._check_quota(user, file_size)
|
||
|
||
# SHA256
|
||
sha256 = hashlib.sha256(content).hexdigest()
|
||
|
||
# 生成文档 token
|
||
doc_token = generate_token()
|
||
doc_token_hash = hash_token(doc_token)
|
||
doc_token_encrypted = encrypt_token(doc_token)
|
||
doc_token_hint = doc_token[-8:] if len(doc_token) >= 8 else doc_token
|
||
|
||
# 存储原始文件
|
||
storage_key = original_object_key(
|
||
user_id=user.id,
|
||
knowledge_base_id=kb_id,
|
||
document_id="pending", # 先存文件,入库后更新路径
|
||
original_filename=filename,
|
||
)
|
||
storage = get_storage()
|
||
storage.save(storage_key, content)
|
||
|
||
# 入库
|
||
doc = self._doc_repo.create(
|
||
knowledge_base_id=kb_id,
|
||
user_id=user.id,
|
||
original_filename=filename,
|
||
storage_path=storage_key,
|
||
file_size=file_size,
|
||
mime_type=detected_mime or "application/octet-stream",
|
||
file_ext=ext,
|
||
sha256=sha256,
|
||
doc_token_hash=doc_token_hash,
|
||
doc_token_encrypted=doc_token_encrypted,
|
||
doc_token_hint=doc_token_hint,
|
||
category_id=category_id,
|
||
)
|
||
|
||
# 更新存储路径中的 document_id
|
||
actual_key = original_object_key(
|
||
user_id=user.id,
|
||
knowledge_base_id=kb_id,
|
||
document_id=doc.id,
|
||
original_filename=filename,
|
||
)
|
||
# 移动文件到正确路径
|
||
if actual_key != storage_key:
|
||
storage.save(actual_key, storage.delete(storage_key) or content)
|
||
doc.storage_path = actual_key
|
||
|
||
# 扣减配额(原子 SQL)
|
||
self._deduct_quota(user, file_size)
|
||
|
||
self._session.commit()
|
||
|
||
# 同步解析文档(MVP:阻塞式)
|
||
self._process_document(doc)
|
||
|
||
return doc
|
||
|
||
def get_or_404(self, doc_id: str, user: User) -> Document:
|
||
doc = self._doc_repo.get_by_id(doc_id)
|
||
if doc is None or doc.user_id != user.id or doc.status == "DELETED":
|
||
raise NotFoundError("文档不存在。")
|
||
return doc
|
||
|
||
def list_by_knowledge_base(
|
||
self, kb_id: str, user: User, *, category_id: str | None = None, page: int = 1, page_size: int = 50
|
||
):
|
||
# 校验 KB 归属
|
||
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 self._doc_repo.list_by_knowledge_base(kb_id, category_id=category_id, page=page, page_size=page_size)
|
||
|
||
def delete(self, doc_id: str, user: User) -> None:
|
||
doc = self.get_or_404(doc_id, user)
|
||
file_size = doc.file_size
|
||
self._doc_repo.delete(doc)
|
||
self._delete_doc_files(doc)
|
||
# 回补配额
|
||
self._restore_quota(user, file_size)
|
||
self._session.commit()
|
||
|
||
def _delete_doc_files(self, doc: Document) -> None:
|
||
"""物理删除文档的原始文件与 Markdown 文件(软删后调用)。"""
|
||
from app.storage.local_storage import get_storage
|
||
|
||
storage = get_storage()
|
||
for key in (doc.storage_path, doc.markdown_path):
|
||
if not key:
|
||
continue
|
||
try:
|
||
storage.delete(key)
|
||
except Exception as exc: # noqa: BLE001
|
||
import logging
|
||
|
||
logging.getLogger(__name__).warning("清理文件失败 key=%s: %s", key, exc)
|
||
|
||
def _check_quota(self, user: User, file_size: int) -> None:
|
||
settings = get_settings()
|
||
if user.storage_used + file_size > settings.default_storage_quota:
|
||
raise StorageQuotaExceededError(
|
||
f"存储空间不足(已用 {user.storage_used // (1024*1024)}MB,"
|
||
f"上传 {file_size // (1024*1024)}MB,"
|
||
f"总配额 {settings.default_storage_quota // (1024*1024)}MB)。"
|
||
)
|
||
|
||
def _deduct_quota(self, user: User, file_size: int) -> None:
|
||
"""原子扣减配额。"""
|
||
from sqlalchemy import update
|
||
|
||
stmt = (
|
||
update(User)
|
||
.where(User.id == user.id, User.storage_used + file_size <= get_settings().default_storage_quota)
|
||
.values(storage_used=User.storage_used + file_size)
|
||
)
|
||
result = self._session.execute(stmt)
|
||
if result.rowcount == 0:
|
||
raise StorageQuotaExceededError("存储空间不足(并发上传导致)。")
|
||
# 更新本地对象
|
||
user.storage_used += file_size
|
||
|
||
def _restore_quota(self, user: User, file_size: int) -> None:
|
||
"""回补配额。"""
|
||
from sqlalchemy import update
|
||
|
||
stmt = (
|
||
update(User)
|
||
.where(User.id == user.id)
|
||
.values(storage_used=User.storage_used - file_size)
|
||
)
|
||
self._session.execute(stmt)
|
||
user.storage_used = max(0, user.storage_used - file_size)
|
||
|
||
def _process_document(self, doc: Document) -> None:
|
||
"""同步处理文档(MVP 阶段,阻塞式)。"""
|
||
from app.processors.local_processor import LocalDocumentProcessor
|
||
|
||
processor = LocalDocumentProcessor(self._session)
|
||
processor.process(doc.id)
|
||
|
||
def create_text(
|
||
self,
|
||
user: User,
|
||
kb_id: str,
|
||
title: str,
|
||
content: str,
|
||
content_format: str = "markdown",
|
||
category_id: str | None = None,
|
||
) -> Document:
|
||
"""创建文本内容文档(无需上传文件)。"""
|
||
# 校验 KB 归属
|
||
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("知识库不存在。")
|
||
|
||
# 生成文档 token
|
||
doc_token = generate_token()
|
||
doc_token_hash = hash_token(doc_token)
|
||
doc_token_encrypted = encrypt_token(doc_token)
|
||
doc_token_hint = doc_token[-8:] if len(doc_token) >= 8 else doc_token
|
||
|
||
# 计算内容大小
|
||
content_size = len(content.encode("utf-8"))
|
||
|
||
# 配额检查
|
||
self._check_quota(user, content_size)
|
||
|
||
# 从内容提取摘要和关键词
|
||
summary = self._extract_summary(content)
|
||
keywords = self._extract_keywords(content)
|
||
|
||
# 创建文档记录
|
||
doc = Document(
|
||
knowledge_base_id=kb_id,
|
||
user_id=user.id,
|
||
category_id=category_id,
|
||
original_filename=title,
|
||
storage_path=None, # 文本内容不存储文件
|
||
markdown_path=None,
|
||
content=content,
|
||
content_format=content_format,
|
||
file_size=content_size,
|
||
mime_type="text/markdown" if content_format == "markdown" else "text/plain",
|
||
file_ext=".md" if content_format == "markdown" else ".txt",
|
||
sha256=None,
|
||
doc_token_hash=doc_token_hash,
|
||
doc_token_encrypted=doc_token_encrypted,
|
||
doc_token_hint=doc_token_hint,
|
||
title=title,
|
||
content_summary=summary,
|
||
keywords=keywords,
|
||
status="READY", # 文本内容直接就绪
|
||
)
|
||
self._session.add(doc)
|
||
|
||
# 扣减配额
|
||
self._deduct_quota(user, content_size)
|
||
|
||
self._session.commit()
|
||
return doc
|
||
|
||
def _extract_summary(self, content: str, max_len: int = 200) -> str:
|
||
"""提取摘要。"""
|
||
import re
|
||
|
||
text = re.sub(r"[#*_`\[\]()>]", "", content)
|
||
text = re.sub(r"\s+", " ", text).strip()
|
||
if len(text) <= max_len:
|
||
return text
|
||
return text[:max_len] + "..."
|
||
|
||
def _extract_keywords(self, content: str, top_k: int = 10) -> str:
|
||
"""提取关键词。"""
|
||
try:
|
||
import jieba.analyse
|
||
import re
|
||
|
||
text = re.sub(r"[#*_`\[\]()>]", "", content)
|
||
text = re.sub(r"\s+", " ", text).strip()
|
||
keywords = jieba.analyse.extract_tags(text, topK=top_k)
|
||
return ",".join(keywords)
|
||
except Exception:
|
||
return ""
|