3
This commit is contained in:
@@ -0,0 +1,129 @@
|
||||
"""认证路由:注册、登录、登出、用户信息。"""
|
||||
|
||||
from fastapi import APIRouter, Depends, Response
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.deps import get_current_user, get_db
|
||||
from app.core.session import get_cookie_params
|
||||
from app.models.user import User
|
||||
from app.schemas.user import (
|
||||
LoginRequest,
|
||||
MeResponse,
|
||||
RegisterRequest,
|
||||
StorageInfoResponse,
|
||||
UpdateMeRequest,
|
||||
UserResponse,
|
||||
)
|
||||
from app.services.auth_service import AuthService
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
|
||||
@router.post("/register", response_model=UserResponse, status_code=201)
|
||||
def register(body: RegisterRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
|
||||
"""注册新用户并自动登录。"""
|
||||
auth_service = AuthService(db)
|
||||
user, token = auth_service.register(body.username, body.email, body.password)
|
||||
response.set_cookie(value=token, **get_cookie_params())
|
||||
return _user_response(user)
|
||||
|
||||
|
||||
@router.post("/login", response_model=UserResponse)
|
||||
def login(body: LoginRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
|
||||
"""登录。"""
|
||||
auth_service = AuthService(db)
|
||||
user, token = auth_service.login(body.username_or_email, body.password)
|
||||
response.set_cookie(value=token, **get_cookie_params())
|
||||
return _user_response(user)
|
||||
|
||||
|
||||
@router.post("/logout", status_code=204)
|
||||
def logout(
|
||||
response: Response,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> None:
|
||||
"""登出:删除 session 并清除 Cookie。"""
|
||||
from app.core.session import SESSION_COOKIE_NAME
|
||||
|
||||
token = "" # 不需要实际 token,logout 内部通过 user_id 找 session
|
||||
auth_service = AuthService(db)
|
||||
# 直接删除所有该用户的 session(MVP 简化:单设备)
|
||||
from app.core.session import _store
|
||||
|
||||
to_delete = [t for t, e in _store.items() if e.user_id == user.id]
|
||||
for t in to_delete:
|
||||
auth_service.logout(t)
|
||||
|
||||
response.delete_cookie(SESSION_COOKIE_NAME, path="/")
|
||||
return None
|
||||
|
||||
|
||||
@router.get("/me", response_model=MeResponse)
|
||||
def get_me(user: User = Depends(get_current_user)) -> MeResponse:
|
||||
"""获取当前用户信息。"""
|
||||
from app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
return MeResponse(
|
||||
id=user.id,
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
status=user.status,
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
created_at=user.created_at,
|
||||
)
|
||||
|
||||
|
||||
@router.patch("/me", response_model=MeResponse)
|
||||
def update_me(
|
||||
body: UpdateMeRequest,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> MeResponse:
|
||||
"""修改密码。"""
|
||||
auth_service = AuthService(db)
|
||||
if body.password is not None:
|
||||
auth_service.change_password(user, body.password)
|
||||
from app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
return MeResponse(
|
||||
id=user.id,
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
status=user.status,
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
created_at=user.created_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/storage", response_model=StorageInfoResponse)
|
||||
def get_storage(user: User = Depends(get_current_user)) -> StorageInfoResponse:
|
||||
"""获取存储用量信息。"""
|
||||
from app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
return StorageInfoResponse(
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
storage_used_mb=round(user.storage_used / (1024 * 1024), 2),
|
||||
storage_quota_mb=round(settings.default_storage_quota / (1024 * 1024), 2),
|
||||
)
|
||||
|
||||
|
||||
def _user_response(user: User) -> UserResponse:
|
||||
from app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
return UserResponse(
|
||||
id=user.id,
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
status=user.status,
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
created_at=user.created_at,
|
||||
)
|
||||
@@ -0,0 +1,31 @@
|
||||
"""API 依赖注入。
|
||||
|
||||
- get_current_user: 从 Cookie 读 session → 查库 → 返回 User
|
||||
- 所有需登录的路由用 Depends(get_current_user)
|
||||
"""
|
||||
|
||||
from fastapi import Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.db import get_session as get_db_session
|
||||
from app.core.session import SESSION_COOKIE_NAME
|
||||
from app.models.user import User
|
||||
from app.services.auth_service import AuthService
|
||||
|
||||
|
||||
def get_db() -> Session:
|
||||
"""数据库会话依赖(同步 SQLAlchemy)。"""
|
||||
yield from get_db_session()
|
||||
|
||||
|
||||
def get_current_user(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> User:
|
||||
"""从 Cookie 中读取 session token,返回当前登录用户。
|
||||
|
||||
未登录或 session 过期抛出 AuthRequiredError(401)。
|
||||
"""
|
||||
token = request.cookies.get(SESSION_COOKIE_NAME)
|
||||
auth_service = AuthService(db)
|
||||
return auth_service.get_current_user(token)
|
||||
@@ -0,0 +1,140 @@
|
||||
"""文档 API 路由。"""
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, UploadFile, File, Form
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.deps import get_current_user, get_db
|
||||
from app.models.user import User
|
||||
from app.schemas.document import (
|
||||
DocumentListResponse,
|
||||
DocumentResponse,
|
||||
DocumentUpdateRequest,
|
||||
DocumentUploadResponse,
|
||||
)
|
||||
from app.services.doc_service import DocumentService
|
||||
|
||||
router = APIRouter(prefix="/documents", tags=["documents"])
|
||||
|
||||
|
||||
@router.post("/upload", response_model=DocumentUploadResponse, status_code=201)
|
||||
async def upload_document(
|
||||
kb_id: str = Form(...),
|
||||
file: UploadFile = File(...),
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> DocumentUploadResponse:
|
||||
"""上传文档到知识库。"""
|
||||
content = await file.read()
|
||||
svc = DocumentService(db)
|
||||
doc = svc.upload(
|
||||
user=user,
|
||||
kb_id=kb_id,
|
||||
filename=file.filename or "unknown",
|
||||
content=content,
|
||||
)
|
||||
return DocumentUploadResponse(
|
||||
id=doc.id,
|
||||
original_filename=doc.original_filename,
|
||||
file_size=doc.file_size,
|
||||
status=doc.status,
|
||||
message="文档上传成功,等待解析。",
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=DocumentListResponse)
|
||||
def list_documents(
|
||||
kb_id: str = Query(...),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(50, ge=1, le=200),
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> DocumentListResponse:
|
||||
svc = DocumentService(db)
|
||||
items, total = svc.list_by_knowledge_base(kb_id, user, page=page, page_size=page_size)
|
||||
return DocumentListResponse(
|
||||
items=[_to_response(doc) for doc in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{doc_id}", response_model=DocumentResponse)
|
||||
def get_document(
|
||||
doc_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> DocumentResponse:
|
||||
svc = DocumentService(db)
|
||||
doc = svc.get_or_404(doc_id, user)
|
||||
return _to_response(doc)
|
||||
|
||||
|
||||
@router.put("/{doc_id}", response_model=DocumentResponse)
|
||||
def update_document(
|
||||
doc_id: str,
|
||||
body: DocumentUpdateRequest,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> DocumentResponse:
|
||||
svc = DocumentService(db)
|
||||
doc = svc.get_or_404(doc_id, user)
|
||||
update_fields = {}
|
||||
if body.title is not None:
|
||||
update_fields["title"] = body.title
|
||||
if body.description is not None:
|
||||
update_fields["description"] = body.description
|
||||
if body.keywords is not None:
|
||||
update_fields["keywords"] = body.keywords
|
||||
if body.category_id is not None:
|
||||
update_fields["category_id"] = body.category_id
|
||||
if update_fields:
|
||||
from app.repositories.doc_repo import DocumentRepository
|
||||
|
||||
DocumentRepository(db).update(doc, **update_fields)
|
||||
db.commit()
|
||||
return _to_response(doc)
|
||||
|
||||
|
||||
@router.delete("/{doc_id}", status_code=204)
|
||||
def delete_document(
|
||||
doc_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> None:
|
||||
svc = DocumentService(db)
|
||||
svc.delete(doc_id, user)
|
||||
|
||||
|
||||
@router.post("/{doc_id}/reprocess", response_model=DocumentResponse)
|
||||
def reprocess_document(
|
||||
doc_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> DocumentResponse:
|
||||
"""重新解析文档。"""
|
||||
svc = DocumentService(db)
|
||||
doc = svc.get_or_404(doc_id, user)
|
||||
svc._process_document(doc)
|
||||
db.refresh(doc)
|
||||
return _to_response(doc)
|
||||
|
||||
|
||||
def _to_response(doc) -> DocumentResponse:
|
||||
return DocumentResponse(
|
||||
id=doc.id,
|
||||
knowledge_base_id=doc.knowledge_base_id,
|
||||
original_filename=doc.original_filename,
|
||||
file_size=doc.file_size,
|
||||
mime_type=doc.mime_type,
|
||||
file_ext=doc.file_ext,
|
||||
sha256=doc.sha256,
|
||||
title=doc.title,
|
||||
description=doc.description,
|
||||
keywords=doc.keywords,
|
||||
content_summary=doc.content_summary,
|
||||
status=doc.status,
|
||||
error_code=doc.error_code,
|
||||
doc_token_hint=doc.doc_token_hint,
|
||||
category_id=doc.category_id,
|
||||
created_at=doc.created_at,
|
||||
updated_at=doc.updated_at,
|
||||
)
|
||||
@@ -0,0 +1,140 @@
|
||||
"""知识库 API 路由。"""
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.deps import get_current_user, get_db
|
||||
from app.models.user import User
|
||||
from app.schemas.knowledge_base import (
|
||||
KbCreateRequest,
|
||||
KbListResponse,
|
||||
KbResponse,
|
||||
KbTokenResponse,
|
||||
KbUpdateRequest,
|
||||
)
|
||||
from app.services.kb_service import KnowledgeBaseService
|
||||
|
||||
router = APIRouter(prefix="/knowledge-bases", tags=["knowledge-bases"])
|
||||
|
||||
|
||||
@router.get("", response_model=KbListResponse)
|
||||
def list_knowledge_bases(
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=100),
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbListResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
items, total = svc.list_by_user(user, page=page, page_size=page_size)
|
||||
doc_count_map = db and svc._kb_repo.get_document_count_map(user.id)
|
||||
return KbListResponse(
|
||||
items=[_to_response(kb, doc_count_map.get(kb.id, 0)) for kb in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.post("", response_model=KbResponse, status_code=201)
|
||||
def create_knowledge_base(
|
||||
body: KbCreateRequest,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb, token = svc.create(user, body.name, body.description)
|
||||
return _to_response(kb, ai_url=f"/k/{token}")
|
||||
|
||||
|
||||
@router.get("/{kb_id}", response_model=KbResponse)
|
||||
def get_knowledge_base(
|
||||
kb_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb = svc.get_or_404(kb_id, user)
|
||||
doc_count = svc._kb_repo.count_documents(kb_id)
|
||||
return _to_response(kb, doc_count)
|
||||
|
||||
|
||||
@router.put("/{kb_id}", response_model=KbResponse)
|
||||
def update_knowledge_base(
|
||||
kb_id: str,
|
||||
body: KbUpdateRequest,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb = svc.update(kb_id, user, body.name, body.description)
|
||||
return _to_response(kb)
|
||||
|
||||
|
||||
@router.delete("/{kb_id}", status_code=204)
|
||||
def delete_knowledge_base(
|
||||
kb_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> None:
|
||||
svc = KnowledgeBaseService(db)
|
||||
svc.delete(kb_id, user)
|
||||
|
||||
|
||||
@router.post("/{kb_id}/regenerate-token", response_model=KbTokenResponse)
|
||||
def regenerate_token(
|
||||
kb_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbTokenResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb, token = svc.regenerate_token(kb_id, user)
|
||||
return KbTokenResponse(token=token, ai_url=f"/k/{token}", token_hint=kb.token_hint)
|
||||
|
||||
|
||||
@router.post("/{kb_id}/enable", response_model=KbResponse)
|
||||
def enable_knowledge_base(
|
||||
kb_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb = svc.set_enabled(kb_id, user, True)
|
||||
return _to_response(kb)
|
||||
|
||||
|
||||
@router.post("/{kb_id}/disable", response_model=KbResponse)
|
||||
def disable_knowledge_base(
|
||||
kb_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb = svc.set_enabled(kb_id, user, False)
|
||||
return _to_response(kb)
|
||||
|
||||
|
||||
@router.get("/{kb_id}/link", response_model=KbTokenResponse)
|
||||
def get_knowledge_base_link(
|
||||
kb_id: str,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbTokenResponse:
|
||||
"""获取完整 AI 链接(解密 token)。"""
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb = svc.get_or_404(kb_id, user)
|
||||
token = svc.get_full_token(kb)
|
||||
if token is None:
|
||||
return KbTokenResponse(token="", ai_url="", token_hint=kb.token_hint or "")
|
||||
return KbTokenResponse(token=token, ai_url=f"/k/{token}", token_hint=kb.token_hint or "")
|
||||
|
||||
|
||||
def _to_response(kb, doc_count: int = 0, ai_url: str | None = None) -> KbResponse:
|
||||
return KbResponse(
|
||||
id=kb.id,
|
||||
name=kb.name,
|
||||
description=kb.description,
|
||||
enabled=kb.enabled,
|
||||
token_hint=kb.token_hint,
|
||||
ai_url=ai_url,
|
||||
document_count=doc_count,
|
||||
created_at=kb.created_at,
|
||||
updated_at=kb.updated_at,
|
||||
)
|
||||
Reference in New Issue
Block a user