diff --git a/backend/alembic/versions/4fcdb390ae18_add_kb_status.py b/backend/alembic/versions/4fcdb390ae18_add_kb_status.py new file mode 100644 index 0000000..a5d3686 --- /dev/null +++ b/backend/alembic/versions/4fcdb390ae18_add_kb_status.py @@ -0,0 +1,29 @@ +"""add_kb_status + +Revision ID: 4fcdb390ae18 +Revises: 5c27c8e31fde +Create Date: 2026-09-01 12:41:55.919052 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '4fcdb390ae18' +down_revision: Union[str, None] = '5c27c8e31fde' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('knowledge_bases', sa.Column('status', sa.String(length=16), nullable=False, comment='状态 (active/deleted)')) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('knowledge_bases', 'status') + # ### end Alembic commands ### diff --git a/backend/app/api/auth.py b/backend/app/api/auth.py new file mode 100644 index 0000000..7570d5b --- /dev/null +++ b/backend/app/api/auth.py @@ -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, + ) \ No newline at end of file diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py new file mode 100644 index 0000000..2f06ccb --- /dev/null +++ b/backend/app/api/deps.py @@ -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) \ No newline at end of file diff --git a/backend/app/api/documents.py b/backend/app/api/documents.py new file mode 100644 index 0000000..0f10e06 --- /dev/null +++ b/backend/app/api/documents.py @@ -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, + ) diff --git a/backend/app/api/knowledge_bases.py b/backend/app/api/knowledge_bases.py new file mode 100644 index 0000000..c6688fa --- /dev/null +++ b/backend/app/api/knowledge_bases.py @@ -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, + ) \ No newline at end of file diff --git a/backend/app/core/errors.py b/backend/app/core/errors.py index 3978b3b..75f1e13 100644 --- a/backend/app/core/errors.py +++ b/backend/app/core/errors.py @@ -119,7 +119,7 @@ async def validation_error_handler(_: Request, exc: RequestValidationError) -> J content={ "code": "VALIDATION_ERROR", "message": "请求参数不正确。", - "detail": exc.errors(include_url=False, include_input=False), + "detail": exc.errors(), }, ) diff --git a/backend/app/core/security.py b/backend/app/core/security.py new file mode 100644 index 0000000..8faaa5d --- /dev/null +++ b/backend/app/core/security.py @@ -0,0 +1,59 @@ +"""密码哈希与 Token 工具。 + +密码:Argon2id(argon2-cffi) +Token:secrets.token_urlsafe + SHA-256 hash + Fernet 加密 +""" + +import hashlib +import secrets + +from argon2 import PasswordHasher +from argon2.exceptions import VerifyMismatchError +from cryptography.fernet import Fernet + +from app.core.config import get_settings + +_ph = PasswordHasher() + + +def hash_password(password: str) -> str: + """Argon2id 哈希密码。""" + return _ph.hash(password) + + +def verify_password(password: str, password_hash: str) -> bool: + """验证密码。""" + try: + return _ph.verify(password_hash, password) + except VerifyMismatchError: + return False + + +def generate_token() -> str: + """生成 128bit 熵 URL-safe token(22 字符)。""" + return secrets.token_urlsafe(16) + + +def hash_token(token: str) -> str: + """SHA-256 哈希 token(用于数据库存储和查询)。""" + return hashlib.sha256(token.encode()).hexdigest() + + +def get_fernet() -> Fernet: + """从 SECRET_KEY 派生 Fernet 实例(用于加密存储 token 原文)。""" + settings = get_settings() + # 将 SECRET_KEY 转为 32 字节 Fernet key(简单派生,MVP 够用) + key = hashlib.sha256(settings.secret_key.encode()).digest() + import base64 + + return Fernet(base64.urlsafe_b64encode(key)) + + +def encrypt_token(token: str) -> str: + """Fernet 加密 token 原文。""" + return get_fernet().encrypt(token.encode()).decode() + + +def decrypt_token(encrypted: str) -> str: + """Fernet 解密 token 原文。""" + return get_fernet().decrypt(encrypted.encode()).decode() \ No newline at end of file diff --git a/backend/app/main.py b/backend/app/main.py index 10eaa8a..d936741 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -1,7 +1,7 @@ """AI Knowledge Link — FastAPI 应用工厂。 -Phase 1 只包含:配置校验、日志、异常处理器、健康检查、lifespan 资源管理。 -后续 Phase 逐步挂载:auth、knowledge-bases、documents、public /k/ 路由。 +Phase 1-3: 配置校验、日志、异常处理器、健康检查、lifespan 资源管理、认证路由。 +后续 Phase 逐步挂载:knowledge-bases、documents、public /k/ 路由。 """ from contextlib import asynccontextmanager @@ -10,15 +10,35 @@ from pathlib import Path from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware +from app.api.auth import router as auth_router +from app.api.documents import router as doc_router from app.api.health import router as health_router +from app.api.knowledge_bases import router as kb_router +from app.public.routes import router as public_router from app.core.config import get_settings -from app.core.db import dispose_engine +from app.core.db import dispose_engine, get_session_factory from app.core.errors import register_exception_handlers from app.core.logging import get_logger, setup_logging logger = get_logger(__name__) +def _seed_free_plan() -> None: + """启动时确保 free plan 存在(幂等)。""" + from app.repositories.plan_repo import PlanRepository + + factory = get_session_factory() + with factory() as session: + repo = PlanRepository(session) + plan = repo.get_by_code("free") + if plan is None: + repo.get_or_create_free() + session.commit() + logger.info("Seeded free plan (100MB / 20MB-per-file)") + else: + logger.debug("Free plan already exists (id=%s)", plan.id) + + @asynccontextmanager async def lifespan(app: FastAPI): settings = get_settings() @@ -29,6 +49,9 @@ async def lifespan(app: FastAPI): data_dir.mkdir(parents=True, exist_ok=True) logger.info("Starting backend (env=%s, data=%s)", settings.environment, data_dir) + # Seed:确保 free plan 存在 + _seed_free_plan() + yield dispose_engine() @@ -58,9 +81,10 @@ def create_app() -> FastAPI: # 路由挂载 app.include_router(health_router, prefix="/api", tags=["health"]) - # Phase 3+: app.include_router(auth_router, prefix="/api/auth", tags=["auth"]) - # Phase 4+: app.include_router(kb_router, prefix="/api/knowledge-bases", tags=["knowledge-bases"]) - # Phase 10+: app.include_router(public_router, tags=["public"]) + app.include_router(auth_router, prefix="/api", tags=["auth"]) + app.include_router(kb_router, prefix="/api", tags=["knowledge-bases"]) + app.include_router(doc_router, prefix="/api", tags=["documents"]) + app.include_router(public_router, tags=["public"]) return app diff --git a/backend/app/models/knowledge_base.py b/backend/app/models/knowledge_base.py index 9f14a2b..bc18d62 100644 --- a/backend/app/models/knowledge_base.py +++ b/backend/app/models/knowledge_base.py @@ -37,6 +37,12 @@ class KnowledgeBase(UUIDPrimaryKeyMixin, TimestampMixin, Base): nullable=False, comment="是否启用", ) + status: Mapped[str] = mapped_column( + String(16), + default="active", + nullable=False, + comment="状态 (active/deleted)", + ) token_hash: Mapped[str] = mapped_column( String(64), unique=True, diff --git a/backend/app/processors/base.py b/backend/app/processors/base.py new file mode 100644 index 0000000..7bff94d --- /dev/null +++ b/backend/app/processors/base.py @@ -0,0 +1,12 @@ +"""文档处理器抽象(扩展接口 3:同步 → 异步任务)。""" + +from typing import Protocol, runtime_checkable + + +@runtime_checkable +class DocumentProcessor(Protocol): + """文档处理统一接口。业务代码通过此接口处理文档,不直接写解析逻辑。""" + + def process(self, document_id: str) -> None: + """处理文档:解析 → 提取元数据 → 保存 Markdown。""" + ... diff --git a/backend/app/processors/local_processor.py b/backend/app/processors/local_processor.py new file mode 100644 index 0000000..82be072 --- /dev/null +++ b/backend/app/processors/local_processor.py @@ -0,0 +1,164 @@ +"""本地文档处理器(MVP:同步处理)。""" + +from pathlib import Path + +from app.core.logging import get_logger +from app.core.security import decrypt_token, encrypt_token, generate_token, hash_token +from app.models.document import Document +from app.processors.parsers.docx_parser import parse_docx_with_python_docx +from app.processors.parsers.markitdown_parser import parse_with_markitdown +from app.processors.parsers.pdf_parser import has_text_layer, parse_pdf_with_pymupdf +from app.repositories.doc_repo import DocumentRepository +from app.storage.local_storage import get_storage +from app.storage.object_keys import markdown_object_key +from sqlalchemy.orm import Session + +logger = get_logger(__name__) + + +class LocalDocumentProcessor: + """本地文档处理器:在同一进程中完成解析。""" + + def __init__(self, session: Session) -> None: + self._session = session + self._doc_repo = DocumentRepository(session) + + def process(self, document_id: str) -> None: + """处理文档:解析 → 提取元数据 → 保存 Markdown → 更新 DB。""" + doc = self._doc_repo.get_by_id(document_id) + if doc is None: + logger.error("Document not found: %s", document_id) + return + + # 更新状态为 PROCESSING + doc.status = "PROCESSING" + self._session.commit() + + try: + # 读取原始文件 + storage = get_storage() + file_content = storage.read(doc.storage_path) + + # 解析 + markdown = self._parse_file(file_content, doc.file_ext, doc.original_filename) + + if markdown is None: + # 解析失败 + if doc.file_ext == ".pdf" and not has_text_layer(file_content): + doc.error_code = "SCANNED_PDF_NO_TEXT_LAYER" + doc.status = "FAILED" + else: + doc.error_code = "PARSING_FAILED" + doc.status = "FAILED" + self._session.commit() + logger.warning("Document parsing failed: %s (%s)", doc.id, doc.error_code) + return + + # 清洗 Markdown + markdown = self._clean_markdown(markdown) + + # 保存 Markdown 文件 + md_key = markdown_object_key( + user_id=doc.user_id, + knowledge_base_id=doc.knowledge_base_id, + document_id=doc.id, + ) + storage.save(md_key, markdown.encode("utf-8")) + doc.markdown_path = md_key + + # 提取元数据 + title = self._extract_title(markdown, doc.original_filename) + summary = self._extract_summary(markdown) + keywords = self._extract_keywords(markdown) + + doc.title = title + doc.content_summary = summary + doc.keywords = keywords + doc.status = "READY" + self._session.commit() + + logger.info("Document processed: %s → READY", doc.id) + + except Exception as e: + logger.error("Document processing error: %s - %s", doc.id, e, exc_info=True) + doc.status = "FAILED" + doc.error_code = "PARSING_FAILED" + self._session.commit() + + def _parse_file(self, content: bytes, ext: str, filename: str) -> str | None: + """按扩展名路由到对应解析器。""" + # 优先 MarkItDown + markdown = parse_with_markitdown(content, filename) + if markdown and markdown.strip(): + return markdown + + # Fallback + if ext == ".pdf": + return parse_pdf_with_pymupdf(content) + elif ext == ".docx": + return parse_docx_with_python_docx(content) + + return None + + def _clean_markdown(self, markdown: str) -> str: + """清洗 Markdown:压缩空行、规整标题层级。""" + import re + + # 压缩连续空行为最多 2 个 + markdown = re.sub(r"\n{3,}", "\n\n", markdown) + # 截断超长行(> 1000 字符的行) + lines = markdown.split("\n") + cleaned = [] + for line in lines: + if len(line) > 1000: + line = line[:1000] + "..." + cleaned.append(line) + return "\n".join(cleaned) + + def _extract_title(self, markdown: str, fallback: str) -> str: + """提取标题:H1 → 文件名(去扩展名)。""" + import re + + # 找第一个 H1 + match = re.search(r"^#\s+(.+)$", markdown, re.MULTILINE) + if match: + return match.group(1).strip() + + # Fallback:文件名去扩展名 + return Path(fallback).stem + + def _extract_summary(self, markdown: str, max_len: int = 200) -> str: + """抽取式摘要:取正文前 ~200 字纯文本。""" + import re + + # 去掉 Markdown 标记 + text = re.sub(r"[#*_`\[\]()>]", "", markdown) + text = re.sub(r"\s+", " ", text).strip() + + if len(text) <= max_len: + return text + + # 在 max_len 附近找句号/逗号断句 + cut = text[:max_len] + for sep in ["。", "!", "?", ";", ".", "!", "?", ";"]: + idx = cut.rfind(sep) + if idx > max_len * 0.5: + return cut[: idx + 1] + return cut + "..." + + def _extract_keywords(self, markdown: str, top_k: int = 10) -> str: + """提取关键词(jieba TF-IDF)。""" + try: + import jieba.analyse + + # 去掉 Markdown 标记 + import re + + text = re.sub(r"[#*_`\[\]()>]", "", markdown) + text = re.sub(r"\s+", " ", text).strip() + + keywords = jieba.analyse.extract_tags(text, topK=top_k) + return ",".join(keywords) + except Exception as e: + logger.warning("关键词提取失败: %s", e) + return "" diff --git a/backend/app/processors/parsers/docx_parser.py b/backend/app/processors/parsers/docx_parser.py new file mode 100644 index 0000000..a2e5b2e --- /dev/null +++ b/backend/app/processors/parsers/docx_parser.py @@ -0,0 +1,54 @@ +"""DOCX 解析器(python-docx fallback)。""" + +from app.core.logging import get_logger + +logger = get_logger(__name__) + + +def parse_docx_with_python_docx(file_content: bytes) -> str | None: + """用 python-docx 解析 DOCX,返回 Markdown 或 None(失败时)。""" + try: + import io + + from docx import Document + + doc = Document(io.BytesIO(file_content)) + md_parts = [] + + for para in doc.paragraphs: + text = para.text.strip() + if not text: + continue + + style = para.style.name.lower() + if style.startswith("heading 1"): + md_parts.append(f"# {text}") + elif style.startswith("heading 2"): + md_parts.append(f"## {text}") + elif style.startswith("heading 3"): + md_parts.append(f"### {text}") + elif style.startswith("heading"): + md_parts.append(f"#### {text}") + elif style.startswith("list"): + md_parts.append(f"- {text}") + else: + md_parts.append(text) + + # 处理表格 + for table in doc.tables: + rows = [] + for row in table.rows: + cells = [cell.text.strip() for cell in row.cells] + rows.append("| " + " | ".join(cells) + " |") + if rows: + # 添加表头分隔行 + if len(rows) > 1: + sep = "| " + " | ".join(["---"] * len(table.columns)) + " |" + rows.insert(1, sep) + md_parts.append("\n".join(rows)) + + markdown = "\n\n".join(md_parts) + return markdown if markdown.strip() else None + except Exception as e: + logger.warning("python-docx 解析失败: %s", e) + return None diff --git a/backend/app/processors/parsers/markitdown_parser.py b/backend/app/processors/parsers/markitdown_parser.py new file mode 100644 index 0000000..ca69b8c --- /dev/null +++ b/backend/app/processors/parsers/markitdown_parser.py @@ -0,0 +1,30 @@ +"""MarkItDown 解析器(优先)。""" + +from app.core.logging import get_logger + +logger = get_logger(__name__) + + +def parse_with_markitdown(file_content: bytes, filename: str) -> str | None: + """用 MarkItDown 解析文档,返回 Markdown 或 None(失败时)。""" + try: + from markitdown import MarkItDown + + md = MarkItDown() + # MarkItDown 需要文件路径或文件对象,用临时文件 + import tempfile + from pathlib import Path + + suffix = Path(filename).suffix.lower() + with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp: + tmp.write(file_content) + tmp_path = tmp.name + + try: + result = md.convert(tmp_path) + return result.text_content if result and result.text_content else None + finally: + Path(tmp_path).unlink(missing_ok=True) + except Exception as e: + logger.warning("MarkItDown 解析失败 (%s): %s", filename, e) + return None diff --git a/backend/app/processors/parsers/pdf_parser.py b/backend/app/processors/parsers/pdf_parser.py new file mode 100644 index 0000000..233890c --- /dev/null +++ b/backend/app/processors/parsers/pdf_parser.py @@ -0,0 +1,47 @@ +"""PDF 解析器(PyMuPDF fallback)。""" + +from app.core.logging import get_logger + +logger = get_logger(__name__) + + +def parse_pdf_with_pymupdf(file_content: bytes) -> str | None: + """用 PyMuPDF 解析 PDF,返回 Markdown 或 None(失败时)。""" + try: + import pymupdf # PyMuPDF >= 1.24 + + doc = pymupdf.open(stream=file_content, filetype="pdf") + pages = [] + for page_num in range(len(doc)): + page = doc.load_page(page_num) + text = page.get_text() + if text.strip(): + pages.append(text) + doc.close() + + if not pages: + return None + + # 组合成 Markdown(每页一个段落) + markdown = "\n\n".join(pages) + return markdown + except Exception as e: + logger.warning("PyMuPDF 解析失败: %s", e) + return None + + +def has_text_layer(file_content: bytes) -> bool: + """检查 PDF 是否有文本层。""" + try: + import pymupdf + + doc = pymupdf.open(stream=file_content, filetype="pdf") + for page_num in range(min(3, len(doc))): # 检查前 3 页 + page = doc.load_page(page_num) + if page.get_text().strip(): + doc.close() + return True + doc.close() + return False + except Exception: + return False diff --git a/backend/app/public/routes.py b/backend/app/public/routes.py new file mode 100644 index 0000000..d7c9d5f --- /dev/null +++ b/backend/app/public/routes.py @@ -0,0 +1,420 @@ +"""公共 AI 页面路由(/k/**)。 + +规则: +- 零 JS、零 Cookie、零登录、SSR 输出、标准 HTML +- +- +- 限流:内存 TokenBucket +- 统一 404 防存在性探测 +""" + +import json + +from fastapi import APIRouter, Depends, Query, Request, Response +from fastapi.responses import HTMLResponse, PlainTextResponse +from sqlalchemy.orm import Session + +from app.api.deps import get_db +from app.core.errors import NotFoundError, RateLimitedError +from app.core.rate_limit import check_rate_limit +from app.core.security import decrypt_token +from app.services.kb_public_service import KbPublicService + +router = APIRouter(prefix="/k", tags=["public"]) + + +def _rate_limit(request: Request, token: str) -> None: + """限流检查。""" + ip = request.client.host if request.client else "unknown" + if not check_rate_limit(token_key=token[:16], ip_key=ip): + raise RateLimitedError() + + +def _robots_meta() -> str: + return '' + + +def _referrer_meta() -> str: + return '' + + +# --- 知识库入口(后缀路由必须先于无后缀路由注册)--- + + +@router.get("/{token}.md") +def kb_index_markdown( + token: str, + request: Request, + db: Session = Depends(get_db), +) -> PlainTextResponse: + """知识库首页(Markdown)。""" + _rate_limit(request, token) + svc = KbPublicService(db) + kb = svc.get_kb_by_token(token) + docs, _ = svc.list_documents(kb) + + lines = [f"# {kb.name}", ""] + if kb.description: + lines.append(kb.description) + lines.append("") + + lines.append("## 文档列表") + lines.append("") + + for doc in docs: + title = doc.title or doc.original_filename + lines.append(f"### {title}") + lines.append(f"- 类型:{doc.file_ext}") + if doc.description: + lines.append(f"- 描述:{doc.description}") + if doc.keywords: + lines.append(f"- 关键词:{doc.keywords}") + if doc.content_summary: + lines.append(f"- 摘要:{doc.content_summary}") + lines.append(f"- 更新时间:{doc.updated_at}") + lines.append("") + + return PlainTextResponse(content="\n".join(lines), media_type="text/markdown") + + +@router.get("/{token}.txt") +def kb_index_text( + token: str, + request: Request, + db: Session = Depends(get_db), +) -> PlainTextResponse: + """知识库首页(纯文本)。""" + _rate_limit(request, token) + svc = KbPublicService(db) + kb = svc.get_kb_by_token(token) + docs, _ = svc.list_documents(kb) + + lines = [kb.name, "=" * len(kb.name), ""] + if kb.description: + lines.append(kb.description) + lines.append("") + + lines.append("文档列表:") + lines.append("") + + for i, doc in enumerate(docs, 1): + title = doc.title or doc.original_filename + lines.append(f"{i}. {title}") + if doc.description: + lines.append(f" 描述:{doc.description}") + if doc.keywords: + lines.append(f" 关键词:{doc.keywords}") + lines.append("") + + return PlainTextResponse(content="\n".join(lines), media_type="text/plain") + + +@router.get("/{token}.json") +def kb_index_json( + token: str, + request: Request, + db: Session = Depends(get_db), +) -> Response: + """知识库首页(JSON)。""" + _rate_limit(request, token) + svc = KbPublicService(db) + kb = svc.get_kb_by_token(token) + docs, _ = svc.list_documents(kb) + + doc_list = [] + for doc in docs: + doc_list.append({ + "title": doc.title or doc.original_filename, + "file_type": doc.file_ext, + "description": doc.description, + "summary": doc.content_summary, + "keywords": doc.keywords.split(",") if doc.keywords else [], + "updated_at": doc.updated_at, + }) + + data = { + "name": kb.name, + "description": kb.description, + "document_count": len(doc_list), + "documents": doc_list, + } + + return Response( + content=json.dumps(data, ensure_ascii=False, indent=2), + media_type="application/json", + ) + + +@router.get("/{token}") +def kb_index_html( + token: str, + request: Request, + page: int = Query(1, ge=1), + db: Session = Depends(get_db), +) -> HTMLResponse: + """知识库首页(HTML)。""" + _rate_limit(request, token) + svc = KbPublicService(db) + kb = svc.get_kb_by_token(token) + docs, total = svc.list_documents(kb, page=page, page_size=50) + + doc_rows = "" + for doc in docs: + doc_url = f"/k/{token}/doc/{decrypt_token(doc.doc_token_encrypted) if doc.doc_token_encrypted else ''}" + keywords_html = f'关键词:{doc.keywords}' if doc.keywords else "" + summary_html = f'
{doc.content_summary or ""}
' if doc.content_summary else "" + doc_rows += f""" +第 {page} / {total_pages} 页
' + + html = f""" + + + +" + kb.description + "
" if kb.description else ""} +暂无文档。
"} + {pagination} + + +""" + return HTMLResponse(content=html) + + +# --- 单文档访问 --- + + +@router.get("/{token}/doc/{doc_token}") +def doc_page_html( + token: str, + doc_token: str, + request: Request, + db: Session = Depends(get_db), +) -> HTMLResponse: + """文档页面(HTML)。""" + _rate_limit(request, token) + svc = KbPublicService(db) + kb = svc.get_kb_by_token(token) + doc = svc.get_document_by_token(kb, doc_token) + markdown_content = svc.get_document_markdown(doc) + + # Markdown → HTML(简单转换) + html_content = _markdown_to_html(markdown_content) + + keywords_html = f"关键词:{doc.keywords}
" if doc.keywords else "" + + html = f""" + + + +" + (item.get('description') or '') + "
"} +共找到 {total} 个结果
+ {result_items if result_items else "未找到相关文档。
"} + + +""" + return HTMLResponse(content=html) + + +@router.get("/{token}/search.json") +def search_json( + token: str, + q: str = Query(..., min_length=1), + page: int = Query(1, ge=1), + request: Request = None, + db: Session = Depends(get_db), +) -> Response: + """搜索文档(JSON)。""" + _rate_limit(request, token) + svc = KbPublicService(db) + kb = svc.get_kb_by_token(token) + results, total = svc.search_documents(kb, q, page=page) + + data = { + "query": q, + "total": total, + "results": results, + } + return Response( + content=json.dumps(data, ensure_ascii=False, indent=2), + media_type="application/json", + ) + + +def _markdown_to_html(markdown: str) -> str: + """简单 Markdown → HTML 转换(安全处理)。""" + import re + + # 转义 HTML 特殊字符 + html = markdown.replace("&", "&").replace("<", "<").replace(">", ">") + + # 标题 + html = re.sub(r"^#### (.+)$", r"{m.group(0)[3:-3]}", html)
+
+ # 行内代码
+ html = re.sub(r"`([^`]+)`", r"\1", html)
+
+ # 段落(双换行 → ) + html = re.sub(r"\n\n+", "
", html) + html = f"
{html}
" + + return html diff --git a/backend/app/repositories/doc_repo.py b/backend/app/repositories/doc_repo.py new file mode 100644 index 0000000..73aa5a5 --- /dev/null +++ b/backend/app/repositories/doc_repo.py @@ -0,0 +1,118 @@ +"""Document 文档 Repository。""" + +from sqlalchemy import func, select +from sqlalchemy.orm import Session + +from app.models.document import Document + + +class DocumentRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def get_by_id(self, doc_id: str) -> Document | None: + return self._session.get(Document, doc_id) + + def get_by_doc_token_hash(self, token_hash: str) -> Document | None: + stmt = select(Document).where(Document.doc_token_hash == token_hash) + return self._session.scalars(stmt).first() + + def list_by_knowledge_base( + self, + kb_id: str, + *, + page: int = 1, + page_size: int = 50, + status: str | None = None, + ) -> tuple[list[Document], int]: + """分页获取知识库的文档列表。""" + conditions = [ + Document.knowledge_base_id == kb_id, + Document.status != "DELETED", + ] + if status: + conditions.append(Document.status == status) + + 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) + ) + items = list(self._session.scalars(stmt).all()) + return items, total + + def list_by_user( + self, user_id: str, *, page: int = 1, page_size: int = 20 + ) -> tuple[list[Document], int]: + """分页获取用户的文档列表。""" + conditions = [Document.user_id == user_id, Document.status != "DELETED"] + + 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) + ) + items = list(self._session.scalars(stmt).all()) + return items, total + + def create( + self, + *, + knowledge_base_id: str, + user_id: str, + original_filename: str, + storage_path: str, + file_size: int, + mime_type: str, + file_ext: str, + sha256: str, + doc_token_hash: str, + doc_token_encrypted: str, + doc_token_hint: str, + category_id: str | None = None, + ) -> Document: + doc = Document( + knowledge_base_id=knowledge_base_id, + user_id=user_id, + original_filename=original_filename, + storage_path=storage_path, + file_size=file_size, + mime_type=mime_type, + file_ext=file_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, + status="PENDING", + ) + self._session.add(doc) + self._session.flush() + return doc + + def update(self, doc: Document, **fields) -> None: + for key, value in fields.items(): + setattr(doc, key, value) + self._session.flush() + + def delete(self, doc: Document) -> None: + doc.status = "DELETED" + self._session.flush() + + def count_by_user(self, user_id: str) -> int: + stmt = ( + select(func.count()) + .select_from(Document) + .where(Document.user_id == user_id, Document.status != "DELETED") + ) + return self._session.scalar(stmt) or 0 diff --git a/backend/app/repositories/kb_repo.py b/backend/app/repositories/kb_repo.py new file mode 100644 index 0000000..6c72474 --- /dev/null +++ b/backend/app/repositories/kb_repo.py @@ -0,0 +1,104 @@ +"""KnowledgeBase 知识库 Repository。""" + +from sqlalchemy import func, select +from sqlalchemy.orm import Session + +from app.models.document import Document +from app.models.knowledge_base import KnowledgeBase + + +class KnowledgeBaseRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def get_by_id(self, kb_id: str) -> KnowledgeBase | None: + return self._session.get(KnowledgeBase, kb_id) + + def get_by_token_hash(self, token_hash: str) -> KnowledgeBase | None: + stmt = select(KnowledgeBase).where(KnowledgeBase.token_hash == token_hash) + return self._session.scalars(stmt).first() + + def list_by_user( + self, user_id: str, *, page: int = 1, page_size: int = 20 + ) -> tuple[list[KnowledgeBase], int]: + """分页获取用户的知识库列表。返回 (items, total)。""" + count_stmt = ( + select(func.count()) + .select_from(KnowledgeBase) + .where( + KnowledgeBase.user_id == user_id, + KnowledgeBase.status != "DELETED", + ) + ) + total = self._session.scalar(count_stmt) or 0 + + stmt = ( + select(KnowledgeBase) + .where( + KnowledgeBase.user_id == user_id, + KnowledgeBase.status != "DELETED", + ) + .order_by(KnowledgeBase.created_at.desc()) + .offset((page - 1) * page_size) + .limit(page_size) + ) + items = list(self._session.scalars(stmt).all()) + return items, total + + def create( + self, + *, + user_id: str, + name: str, + description: str | None, + token_hash: str, + token_encrypted: str, + token_hint: str, + ) -> KnowledgeBase: + kb = KnowledgeBase( + user_id=user_id, + name=name, + description=description, + enabled=True, + status="active", + token_hash=token_hash, + token_encrypted=token_encrypted, + token_hint=token_hint, + ) + self._session.add(kb) + self._session.flush() + return kb + + def update(self, kb: KnowledgeBase, **fields) -> None: + for key, value in fields.items(): + if value is not None: + setattr(kb, key, value) + self._session.flush() + + def delete(self, kb: KnowledgeBase) -> None: + """软删除。""" + kb.status = "DELETED" + self._session.flush() + + def count_documents(self, kb_id: str) -> int: + stmt = ( + select(func.count()) + .select_from(Document) + .where( + Document.knowledge_base_id == kb_id, + Document.status != "DELETED", + ) + ) + return self._session.scalar(stmt) or 0 + + def get_document_count_map(self, user_id: str) -> dict[str, int]: + """批量获取用户所有知识库的文档数量(避免 N+1)。""" + stmt = ( + select(Document.knowledge_base_id, func.count()) + .where( + Document.user_id == user_id, + Document.status != "DELETED", + ) + .group_by(Document.knowledge_base_id) + ) + return dict(self._session.execute(stmt).all()) \ No newline at end of file diff --git a/backend/app/repositories/plan_repo.py b/backend/app/repositories/plan_repo.py new file mode 100644 index 0000000..35cd7b2 --- /dev/null +++ b/backend/app/repositories/plan_repo.py @@ -0,0 +1,33 @@ +"""Plan 套餐 Repository。""" + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.models.plan import Plan + + +class PlanRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def get_by_id(self, plan_id: str) -> Plan | None: + return self._session.get(Plan, plan_id) + + def get_by_code(self, code: str) -> Plan | None: + stmt = select(Plan).where(Plan.code == code) + return self._session.scalars(stmt).first() + + def get_or_create_free(self) -> Plan: + """获取 free plan,不存在则创建(seed 逻辑)。""" + plan = self.get_by_code("free") + if plan is None: + plan = Plan( + code="free", + name="免费版", + storage_quota=104_857_600, # 100 MB + max_file_size=20_971_520, # 20 MB + is_active=True, + ) + self._session.add(plan) + self._session.flush() + return plan \ No newline at end of file diff --git a/backend/app/repositories/user_repo.py b/backend/app/repositories/user_repo.py new file mode 100644 index 0000000..f135e1f --- /dev/null +++ b/backend/app/repositories/user_repo.py @@ -0,0 +1,64 @@ +"""User 用户 Repository。""" + +from sqlalchemy import select, or_ +from sqlalchemy.orm import Session + +from app.models.user import User + + +class UserRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def get_by_id(self, user_id: str) -> User | None: + return self._session.get(User, user_id) + + def get_by_username(self, username: str) -> User | None: + stmt = select(User).where(User.username == username) + return self._session.scalars(stmt).first() + + def get_by_email(self, email: str) -> User | None: + stmt = select(User).where(User.email == email) + return self._session.scalars(stmt).first() + + def get_by_username_or_email(self, value: str) -> User | None: + """登录用:按用户名或邮箱查找。""" + stmt = select(User).where( + or_(User.username == value, User.email == value.lower()) + ) + return self._session.scalars(stmt).first() + + def exists_username_or_email(self, username: str, email: str) -> tuple[bool, bool]: + """检查用户名/邮箱是否已存在。返回 (username_taken, email_taken)。""" + stmt = select(User).where( + or_(User.username == username, User.email == email.lower()) + ) + existing = list(self._session.scalars(stmt).all()) + return ( + any(u.username == username for u in existing), + any(u.email == email.lower() for u in existing), + ) + + def create( + self, + *, + username: str, + email: str, + password_hash: str, + plan_id: str, + ) -> User: + user = User( + username=username, + email=email.lower(), + password_hash=password_hash, + status="active", + plan_id=plan_id, + storage_used=0, + ) + self._session.add(user) + self._session.flush() + return user + + def update_password(self, user: User, password_hash: str) -> None: + user.password_hash = password_hash + self._session.flush() \ No newline at end of file diff --git a/backend/app/schemas/__init__.py b/backend/app/schemas/__init__.py index e69de29..0253526 100644 --- a/backend/app/schemas/__init__.py +++ b/backend/app/schemas/__init__.py @@ -0,0 +1,19 @@ +"""Pydantic v2 Schema 包。""" + +from app.schemas.user import ( + LoginRequest, + MeResponse, + RegisterRequest, + StorageInfoResponse, + UpdateMeRequest, + UserResponse, +) + +__all__ = [ + "LoginRequest", + "MeResponse", + "RegisterRequest", + "StorageInfoResponse", + "UpdateMeRequest", + "UserResponse", +] \ No newline at end of file diff --git a/backend/app/schemas/document.py b/backend/app/schemas/document.py new file mode 100644 index 0000000..5f775ec --- /dev/null +++ b/backend/app/schemas/document.py @@ -0,0 +1,45 @@ +"""文档相关 Pydantic v2 Schema。""" + +from pydantic import BaseModel, Field + + +class DocumentResponse(BaseModel): + id: str + knowledge_base_id: str + original_filename: str + file_size: int + mime_type: str + file_ext: str + sha256: str + title: str | None + description: str | None + keywords: str | None + content_summary: str | None + status: str + error_code: str | None + doc_token_hint: str | None + category_id: str | None + created_at: str + updated_at: str + + model_config = {"from_attributes": True} + + +class DocumentListResponse(BaseModel): + items: list[DocumentResponse] + total: int + + +class DocumentUpdateRequest(BaseModel): + title: str | None = Field(default=None, max_length=512) + description: str | None = Field(default=None, max_length=2000) + keywords: str | None = Field(default=None, max_length=500) + category_id: str | None = None + + +class DocumentUploadResponse(BaseModel): + id: str + original_filename: str + file_size: int + status: str + message: str diff --git a/backend/app/schemas/knowledge_base.py b/backend/app/schemas/knowledge_base.py new file mode 100644 index 0000000..a80488b --- /dev/null +++ b/backend/app/schemas/knowledge_base.py @@ -0,0 +1,39 @@ +"""知识库相关 Pydantic v2 Schema。""" + +from pydantic import BaseModel, Field + + +class KbCreateRequest(BaseModel): + name: str = Field(min_length=1, max_length=255) + description: str | None = Field(default=None, max_length=2000) + + +class KbUpdateRequest(BaseModel): + name: str | None = Field(default=None, min_length=1, max_length=255) + description: str | None = None + + +class KbResponse(BaseModel): + id: str + name: str + description: str | None + enabled: bool + token_hint: str | None + ai_url: str | None = None # 仅创建/重置时返回完整 URL + document_count: int = 0 + created_at: str + updated_at: str + + model_config = {"from_attributes": True} + + +class KbListResponse(BaseModel): + items: list[KbResponse] + total: int + + +class KbTokenResponse(BaseModel): + """Token 重置/创建时的完整 URL 响应。""" + token: str + ai_url: str + token_hint: str \ No newline at end of file diff --git a/backend/app/schemas/user.py b/backend/app/schemas/user.py new file mode 100644 index 0000000..a2b9c47 --- /dev/null +++ b/backend/app/schemas/user.py @@ -0,0 +1,82 @@ +"""用户相关 Pydantic v2 Schema(请求/响应)。""" + +import re + +from pydantic import BaseModel, EmailStr, Field, field_validator + +_USERNAME_RE = re.compile(r"^[A-Za-z0-9_\-一-鿿]{2,32}$") + + +class RegisterRequest(BaseModel): + username: str = Field(min_length=2, max_length=32) + email: str = Field(min_length=5, max_length=255) + password: str = Field(min_length=8, max_length=128) + + @field_validator("username") + @classmethod + def validate_username(cls, v: str) -> str: + if not _USERNAME_RE.match(v): + raise ValueError("用户名只能包含字母、数字、下划线、短横线和中文,长度 2-32。") + return v + + @field_validator("email") + @classmethod + def validate_email(cls, v: str) -> str: + # 简单格式校验(不引入 email-validator 重依赖) + if "@" not in v or "." not in v.split("@")[-1]: + raise ValueError("邮箱格式不正确。") + return v.lower().strip() + + @field_validator("password") + @classmethod + def validate_password(cls, v: str) -> str: + if len(v) < 8: + raise ValueError("密码长度不能少于 8 位。") + return v + + +class LoginRequest(BaseModel): + username_or_email: str = Field(min_length=2, max_length=255) + password: str = Field(min_length=1, max_length=128) + + +class UserResponse(BaseModel): + id: str + username: str + email: str + status: str + storage_used: int + storage_quota: int + created_at: str + + model_config = {"from_attributes": True} + + +class MeResponse(BaseModel): + id: str + username: str + email: str + status: str + storage_used: int + storage_quota: int + created_at: str + + model_config = {"from_attributes": True} + + +class StorageInfoResponse(BaseModel): + storage_used: int + storage_quota: int + storage_used_mb: float + storage_quota_mb: float + + +class UpdateMeRequest(BaseModel): + password: str | None = Field(default=None, min_length=8, max_length=128) + + @field_validator("password") + @classmethod + def validate_password(cls, v: str | None) -> str | None: + if v is not None and len(v) < 8: + raise ValueError("密码长度不能少于 8 位。") + return v \ No newline at end of file diff --git a/backend/app/services/auth_service.py b/backend/app/services/auth_service.py new file mode 100644 index 0000000..fa5e07b --- /dev/null +++ b/backend/app/services/auth_service.py @@ -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() \ No newline at end of file diff --git a/backend/app/services/doc_service.py b/backend/app/services/doc_service.py new file mode 100644 index 0000000..3d48043 --- /dev/null +++ b/backend/app/services/doc_service.py @@ -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) diff --git a/backend/app/services/kb_public_service.py b/backend/app/services/kb_public_service.py new file mode 100644 index 0000000..90b2766 --- /dev/null +++ b/backend/app/services/kb_public_service.py @@ -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 diff --git a/backend/app/services/kb_service.py b/backend/app/services/kb_service.py new file mode 100644 index 0000000..a2a5529 --- /dev/null +++ b/backend/app/services/kb_service.py @@ -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 \ No newline at end of file diff --git a/backend/app/services/token_service.py b/backend/app/services/token_service.py new file mode 100644 index 0000000..dcea51f --- /dev/null +++ b/backend/app/services/token_service.py @@ -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) \ No newline at end of file diff --git a/backend/requirements.txt b/backend/requirements.txt index aa6f2aa..0a6ee10 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -15,6 +15,7 @@ alembic>=1.13 # --- 安全 --- argon2-cffi>=23.1 cryptography>=42.0 +python-multipart>=0.0.7 # --- 文档解析 --- # MIME 嗅探:纯 Python,Windows 无需 libmagic DLL diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py new file mode 100644 index 0000000..8c126b1 --- /dev/null +++ b/backend/tests/conftest.py @@ -0,0 +1,44 @@ +"""测试配置:每个测试用例使用隔离的内存数据库。""" + +import pytest +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker + +from app.core import db as db_module +from app.models import Base + + +@pytest.fixture(autouse=True) +def _isolate_db(tmp_path, monkeypatch): + """每个测试用例:创建独立内存 SQLite → 建表 → 替换全局引擎 → 测试结束自动清理。 + + 这确保测试之间完全隔离,不共享任何数据。 + """ + test_db_url = f"sqlite:///{tmp_path}/test.db" + engine = create_engine(test_db_url, connect_args={"check_same_thread": False}) + + # 启用 WAL + 外键 + from sqlalchemy import event, text + + @event.listens_for(engine, "connect") + def _set_pragma(dbapi_conn, _): + cursor = dbapi_conn.cursor() + cursor.execute("PRAGMA journal_mode=WAL") + cursor.execute("PRAGMA foreign_keys=ON") + cursor.close() + + Base.metadata.create_all(engine) + test_factory = sessionmaker(bind=engine, expire_on_commit=False) + + # 替换全局引擎和 session 工厂 + monkeypatch.setattr(db_module, "_engine", engine) + monkeypatch.setattr(db_module, "_session_factory", test_factory) + + # 清空内存 session 存储 + from app.core import session as session_module + + session_module._store.clear() + + yield + + engine.dispose() \ No newline at end of file diff --git a/backend/tests/test_auth.py b/backend/tests/test_auth.py new file mode 100644 index 0000000..eca18b7 --- /dev/null +++ b/backend/tests/test_auth.py @@ -0,0 +1,208 @@ +"""Phase 3 认证测试:注册、登录、登出、鉴权、IDOR 防护。""" + +from fastapi.testclient import TestClient + +from app.main import app + + +def _client() -> TestClient: + return TestClient(app, raise_server_exceptions=False) + + +# --- 注册 --- + +def test_register_success() -> None: + with _client() as client: + resp = client.post("/api/auth/register", json={ + "username": "alice", + "email": "alice@example.com", + "password": "password123", + }) + assert resp.status_code == 201 + body = resp.json() + assert body["username"] == "alice" + assert body["email"] == "alice@example.com" + assert body["storage_quota"] == 104_857_600 + # Cookie 已签发 + assert "session_id" in resp.cookies + + +def test_register_duplicate_username() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "bob", + "email": "bob@example.com", + "password": "password123", + }) + resp = client.post("/api/auth/register", json={ + "username": "bob", + "email": "bob2@example.com", + "password": "password123", + }) + assert resp.status_code == 409 + assert resp.json()["code"] == "USERNAME_TAKEN" + + +def test_register_duplicate_email() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "carol", + "email": "carol@example.com", + "password": "password123", + }) + resp = client.post("/api/auth/register", json={ + "username": "carol2", + "email": "carol@example.com", + "password": "password123", + }) + assert resp.status_code == 409 + assert resp.json()["code"] == "EMAIL_TAKEN" + + +def test_register_short_password() -> None: + with _client() as client: + resp = client.post("/api/auth/register", json={ + "username": "dave", + "email": "dave@example.com", + "password": "short", + }) + assert resp.status_code == 422 + + +# --- 登录 --- + +def test_login_by_username() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "eve", + "email": "eve@example.com", + "password": "password123", + }) + resp = client.post("/api/auth/login", json={ + "username_or_email": "eve", + "password": "password123", + }) + assert resp.status_code == 200 + assert resp.json()["username"] == "eve" + assert "session_id" in resp.cookies + + +def test_login_by_email() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "frank", + "email": "frank@example.com", + "password": "password123", + }) + resp = client.post("/api/auth/login", json={ + "username_or_email": "frank@example.com", + "password": "password123", + }) + assert resp.status_code == 200 + assert resp.json()["username"] == "frank" + + +def test_login_wrong_password() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "grace", + "email": "grace@example.com", + "password": "password123", + }) + resp = client.post("/api/auth/login", json={ + "username_or_email": "grace", + "password": "wrongpassword", + }) + assert resp.status_code == 401 + assert resp.json()["code"] == "AUTH_INVALID_CREDENTIALS" + + +def test_login_nonexistent_user() -> None: + with _client() as client: + resp = client.post("/api/auth/login", json={ + "username_or_email": "nobody", + "password": "password123", + }) + assert resp.status_code == 401 + + +# --- 登出 --- + +def test_logout() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "heidi", + "email": "heidi@example.com", + "password": "password123", + }) + resp = client.post("/api/auth/logout") + assert resp.status_code == 204 + + +# --- /api/me --- + +def test_me_authenticated() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "ivan", + "email": "ivan@example.com", + "password": "password123", + }) + resp = client.get("/api/auth/me") + assert resp.status_code == 200 + body = resp.json() + assert body["username"] == "ivan" + assert body["storage_used"] == 0 + + +def test_me_unauthenticated() -> None: + with _client() as client: + resp = client.get("/api/auth/me") + assert resp.status_code == 401 + assert resp.json()["code"] == "AUTH_REQUIRED" + + +# --- IDOR 防护 --- + +def test_user_a_cannot_see_user_b() -> None: + """用户A 登录后,session 只能访问自己的 /me。""" + with _client() as client: + client.post("/api/auth/register", json={ + "username": "user_a", + "email": "a@example.com", + "password": "password123", + }) + # user_a 已登录(Cookie 自动携带) + me_resp = client.get("/api/auth/me") + assert me_resp.status_code == 200 + assert me_resp.json()["username"] == "user_a" + + +def test_storage_endpoint() -> None: + with _client() as client: + client.post("/api/auth/register", json={ + "username": "storage_user", + "email": "storage@example.com", + "password": "password123", + }) + resp = client.get("/api/auth/storage") + assert resp.status_code == 200 + body = resp.json() + assert body["storage_used"] == 0 + assert body["storage_quota"] == 104_857_600 + assert body["storage_quota_mb"] == 100.0 + + +# --- free plan seed --- + +def test_free_plan_auto_created() -> None: + """free plan 应在应用启动时自动创建。""" + with _client() as client: + client.post("/api/auth/register", json={ + "username": "plan_user", + "email": "plan@example.com", + "password": "password123", + }) + resp = client.get("/api/auth/me") + assert resp.status_code == 200 + assert resp.json()["storage_quota"] == 104_857_600 diff --git a/backend/tests/test_documents.py b/backend/tests/test_documents.py new file mode 100644 index 0000000..6c3c0d8 --- /dev/null +++ b/backend/tests/test_documents.py @@ -0,0 +1,229 @@ +"""Phase 6 文档上传测试。""" + +import io + +from fastapi.testclient import TestClient + +from app.main import app + + +def _client() -> TestClient: + return TestClient(app, raise_server_exceptions=False) + + +def _register_and_login(client: TestClient, username: str = "testuser") -> None: + client.post("/api/auth/register", json={ + "username": username, + "email": f"{username}@example.com", + "password": "password123", + }) + + +def _create_kb(client: TestClient, name: str = "测试知识库") -> str: + resp = client.post("/api/knowledge-bases", json={"name": name}) + return resp.json()["id"] + + +def _make_docx_bytes() -> bytes: + """创建一个最小的 .docx 文件字节(实际是 ZIP 格式)。""" + import zipfile + + buf = io.BytesIO() + with zipfile.ZipFile(buf, "w") as zf: + zf.writestr("[Content_Types].xml", '