3
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
"""认证服务:注册、登录、登出、Session 校验。"""
|
||||
|
||||
from app.core.errors import (
|
||||
AuthRequiredError,
|
||||
ConflictError,
|
||||
InvalidCredentialsError,
|
||||
)
|
||||
from app.core.security import hash_password, verify_password
|
||||
from app.core.session import create_session, delete_session, get_session_user
|
||||
from app.models.user import User
|
||||
from app.repositories.plan_repo import PlanRepository
|
||||
from app.repositories.user_repo import UserRepository
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
class AuthService:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
self._user_repo = UserRepository(session)
|
||||
self._plan_repo = PlanRepository(session)
|
||||
|
||||
def register(self, username: str, email: str, password: str) -> tuple[User, str]:
|
||||
"""注册新用户。返回 (user, session_token)。
|
||||
|
||||
Raises:
|
||||
ConflictError: 用户名或邮箱已存在
|
||||
"""
|
||||
username_taken, email_taken = self._user_repo.exists_username_or_email(username, email)
|
||||
if username_taken:
|
||||
raise ConflictError("用户名已被占用。", code="USERNAME_TAKEN")
|
||||
if email_taken:
|
||||
raise ConflictError("邮箱已被注册。", code="EMAIL_TAKEN")
|
||||
|
||||
plan = self._plan_repo.get_or_create_free()
|
||||
user = self._user_repo.create(
|
||||
username=username,
|
||||
email=email,
|
||||
password_hash=hash_password(password),
|
||||
plan_id=plan.id,
|
||||
)
|
||||
self._session.commit()
|
||||
|
||||
token = create_session(user.id)
|
||||
return user, token
|
||||
|
||||
def login(self, username_or_email: str, password: str) -> tuple[User, str]:
|
||||
"""登录。返回 (user, session_token)。
|
||||
|
||||
Raises:
|
||||
InvalidCredentialsError: 用户名/邮箱或密码错误
|
||||
"""
|
||||
user = self._user_repo.get_by_username_or_email(username_or_email)
|
||||
if user is None:
|
||||
raise InvalidCredentialsError()
|
||||
if not verify_password(password, user.password_hash):
|
||||
raise InvalidCredentialsError()
|
||||
if user.status != "active":
|
||||
raise InvalidCredentialsError("账户已被禁用。")
|
||||
|
||||
token = create_session(user.id)
|
||||
return user, token
|
||||
|
||||
def logout(self, session_token: str) -> None:
|
||||
"""登出:删除 session。"""
|
||||
delete_session(session_token)
|
||||
|
||||
def get_current_user(self, session_token: str | None) -> User:
|
||||
"""根据 session token 获取当前用户。
|
||||
|
||||
Raises:
|
||||
AuthRequiredError: 未登录或 session 已过期
|
||||
"""
|
||||
if not session_token:
|
||||
raise AuthRequiredError()
|
||||
user_id = get_session_user(session_token)
|
||||
if user_id is None:
|
||||
raise AuthRequiredError("登录已过期,请重新登录。")
|
||||
user = self._user_repo.get_by_id(user_id)
|
||||
if user is None or user.status != "active":
|
||||
raise AuthRequiredError()
|
||||
return user
|
||||
|
||||
def change_password(self, user: User, new_password: str) -> None:
|
||||
"""修改密码。"""
|
||||
self._user_repo.update_password(user, hash_password(new_password))
|
||||
self._session.commit()
|
||||
@@ -0,0 +1,202 @@
|
||||
"""文档服务:上传、删除、配额管理。"""
|
||||
|
||||
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,
|
||||
) -> 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,
|
||||
)
|
||||
|
||||
# 更新存储路径中的 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, *, 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, 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._restore_quota(user, file_size)
|
||||
self._session.commit()
|
||||
|
||||
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)
|
||||
@@ -0,0 +1,100 @@
|
||||
"""公共知识库服务(Phase 10/11/12 统一数据来源)。
|
||||
|
||||
HTML/MD/TXT/JSON/搜索全部通过此 Service 获取数据,不各自写查询逻辑。
|
||||
"""
|
||||
|
||||
from app.core.errors import NotFoundError
|
||||
from app.core.security import hash_token
|
||||
from app.models.document import Document
|
||||
from app.models.knowledge_base import KnowledgeBase
|
||||
from app.repositories.doc_repo import DocumentRepository
|
||||
from app.repositories.kb_repo import KnowledgeBaseRepository
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
class KbPublicService:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
self._kb_repo = KnowledgeBaseRepository(session)
|
||||
self._doc_repo = DocumentRepository(session)
|
||||
|
||||
def get_kb_by_token(self, token: str) -> KnowledgeBase:
|
||||
"""通过 token 获取知识库。不存在/禁用/删除 → 404。"""
|
||||
token_hash = hash_token(token)
|
||||
kb = self._kb_repo.get_by_token_hash(token_hash)
|
||||
if kb is None or not kb.enabled or kb.status == "DELETED":
|
||||
raise NotFoundError("知识库不存在。")
|
||||
return kb
|
||||
|
||||
def list_documents(
|
||||
self, kb: KnowledgeBase, *, page: int = 1, page_size: int = 50
|
||||
) -> tuple[list[Document], int]:
|
||||
"""获取知识库的文档列表(仅 READY 状态)。"""
|
||||
return self._doc_repo.list_by_knowledge_base(
|
||||
kb.id, page=page, page_size=page_size, status="READY"
|
||||
)
|
||||
|
||||
def get_document_by_token(self, kb: KnowledgeBase, doc_token: str) -> Document:
|
||||
"""通过 token 获取单个文档。"""
|
||||
doc_token_hash = hash_token(doc_token)
|
||||
doc = self._doc_repo.get_by_doc_token_hash(doc_token_hash)
|
||||
if doc is None or doc.knowledge_base_id != kb.id or doc.status != "READY":
|
||||
raise NotFoundError("文档不存在。")
|
||||
return doc
|
||||
|
||||
def get_document_markdown(self, doc: Document) -> str:
|
||||
"""读取文档的 Markdown 内容。"""
|
||||
from app.storage.local_storage import get_storage
|
||||
|
||||
if not doc.markdown_path:
|
||||
return ""
|
||||
storage = get_storage()
|
||||
content = storage.read(doc.markdown_path)
|
||||
return content.decode("utf-8")
|
||||
|
||||
def search_documents(
|
||||
self, kb: KnowledgeBase, query: str, *, page: int = 1, page_size: int = 20
|
||||
) -> tuple[list[dict], int]:
|
||||
"""关键词搜索文档(Phase 12:SQLite LIKE / FTS5)。"""
|
||||
# MVP 简化:使用 LIKE 搜索标题+描述+关键词
|
||||
from sqlalchemy import func, or_, select
|
||||
|
||||
conditions = [
|
||||
Document.knowledge_base_id == kb.id,
|
||||
Document.status == "READY",
|
||||
]
|
||||
|
||||
like_pattern = f"%{query}%"
|
||||
search_condition = or_(
|
||||
Document.title.like(like_pattern),
|
||||
Document.description.like(like_pattern),
|
||||
Document.keywords.like(like_pattern),
|
||||
Document.content_summary.like(like_pattern),
|
||||
)
|
||||
conditions.append(search_condition)
|
||||
|
||||
count_stmt = select(func.count()).select_from(Document).where(*conditions)
|
||||
total = self._session.scalar(count_stmt) or 0
|
||||
|
||||
stmt = (
|
||||
select(Document)
|
||||
.where(*conditions)
|
||||
.order_by(Document.created_at.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
docs = list(self._session.scalars(stmt).all())
|
||||
|
||||
results = []
|
||||
for doc in docs:
|
||||
results.append({
|
||||
"id": doc.id,
|
||||
"title": doc.title or doc.original_filename,
|
||||
"description": doc.description,
|
||||
"keywords": doc.keywords,
|
||||
"file_type": doc.file_ext,
|
||||
"updated_at": doc.updated_at,
|
||||
"url_hint": doc.doc_token_hint,
|
||||
})
|
||||
|
||||
return results, total
|
||||
@@ -0,0 +1,85 @@
|
||||
"""知识库服务:CRUD + Token 管理。"""
|
||||
|
||||
from app.core.errors import NotFoundError, PermissionDeniedError
|
||||
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
|
||||
|
||||
|
||||
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._session.commit()
|
||||
return kb, token
|
||||
|
||||
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._session.commit()
|
||||
|
||||
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 get_full_token(self, kb: KnowledgeBase) -> str | None:
|
||||
"""解密 token 原文(供后台显示完整链接)。"""
|
||||
if kb.token_encrypted:
|
||||
return self._token_svc.decrypt_token(kb.token_encrypted)
|
||||
return None
|
||||
@@ -0,0 +1,31 @@
|
||||
"""Token 服务:生成、哈希、加密、解密。
|
||||
|
||||
用于知识库和文档的 Secret URL token 管理。
|
||||
"""
|
||||
|
||||
from app.core.security import decrypt_token, encrypt_token, generate_token, hash_token
|
||||
|
||||
|
||||
class TokenService:
|
||||
@staticmethod
|
||||
def create_token_pair() -> tuple[str, str, str, str]:
|
||||
"""生成 token 并返回 (token, token_hash, token_encrypted, token_hint)。
|
||||
|
||||
- token: 原文(仅此一次返回给用户)
|
||||
- token_hash: SHA-256 哈希(存 DB,用于查询)
|
||||
- token_encrypted: Fernet 加密原文(存 DB,供后台显示完整链接)
|
||||
- token_hint: 末 8 位明文(存 DB,供后台识别)
|
||||
"""
|
||||
token = generate_token()
|
||||
token_hash = hash_token(token)
|
||||
token_encrypted = encrypt_token(token)
|
||||
token_hint = token[-8:] if len(token) >= 8 else token
|
||||
return token, token_hash, token_encrypted, token_hint
|
||||
|
||||
@staticmethod
|
||||
def hash_token(token: str) -> str:
|
||||
return hash_token(token)
|
||||
|
||||
@staticmethod
|
||||
def decrypt_token(encrypted: str) -> str:
|
||||
return decrypt_token(encrypted)
|
||||
Reference in New Issue
Block a user