细节优化
This commit is contained in:
@@ -0,0 +1,29 @@
|
||||
"""add_kb_token_expires_at
|
||||
|
||||
Revision ID: 5bb03575e2e7
|
||||
Revises: 4e40432cab9f
|
||||
Create Date: 2026-09-02 17:49:46.298836
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '5bb03575e2e7'
|
||||
down_revision: Union[str, None] = '4e40432cab9f'
|
||||
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('token_expires_at', sa.String(length=32), nullable=True, comment='链接过期时间,NULL=长期有效'))
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_column('knowledge_bases', 'token_expires_at')
|
||||
# ### end Alembic commands ###
|
||||
@@ -9,6 +9,7 @@ from app.schemas.knowledge_base import (
|
||||
KbCreateRequest,
|
||||
KbListResponse,
|
||||
KbResponse,
|
||||
KbSetExpiryRequest,
|
||||
KbTokenResponse,
|
||||
KbUpdateRequest,
|
||||
)
|
||||
@@ -86,6 +87,8 @@ def regenerate_token(
|
||||
) -> KbTokenResponse:
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb, token = svc.regenerate_token(kb_id, user)
|
||||
# 重新生成 = 全新链接,有效期重置为长期
|
||||
svc.set_expiry(kb_id, user, None)
|
||||
return KbTokenResponse(token=token, ai_url=f"/k/{token}", token_hint=kb.token_hint)
|
||||
|
||||
|
||||
@@ -117,13 +120,42 @@ def get_knowledge_base_link(
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbTokenResponse:
|
||||
"""获取完整 AI 链接(解密 token)。"""
|
||||
"""获取完整 AI 链接(解密 token)+ 有效期信息。"""
|
||||
from app.services.kb_public_service import KbPublicService
|
||||
|
||||
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 "")
|
||||
return KbTokenResponse(
|
||||
token="", ai_url="", token_hint=kb.token_hint or "",
|
||||
expires_at=None, is_expired=False,
|
||||
)
|
||||
return KbTokenResponse(
|
||||
token=token,
|
||||
ai_url=f"/k/{token}",
|
||||
token_hint=kb.token_hint or "",
|
||||
expires_at=kb.token_expires_at,
|
||||
is_expired=KbPublicService.is_expired(kb),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/{kb_id}/set-expiry", response_model=KbResponse)
|
||||
def set_link_expiry(
|
||||
kb_id: str,
|
||||
body: KbSetExpiryRequest,
|
||||
user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> KbResponse:
|
||||
"""设置 AI 链接有效期。expires_in_minutes 为 null 表示长期有效。"""
|
||||
svc = KnowledgeBaseService(db)
|
||||
kb = svc.set_expiry(kb_id, user, body.expires_in_minutes)
|
||||
doc_count = svc._kb_repo.count_documents(kb_id)
|
||||
return KbResponse(
|
||||
id=kb.id, name=kb.name, description=kb.description, enabled=kb.enabled,
|
||||
token_hint=kb.token_hint, document_count=doc_count,
|
||||
created_at=kb.created_at, updated_at=kb.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _to_response(kb, doc_count: int = 0, ai_url: str | None = None) -> KbResponse:
|
||||
|
||||
@@ -38,6 +38,10 @@ class Settings(BaseSettings):
|
||||
rate_limit_per_token_per_min: int = 60
|
||||
rate_limit_per_ip_per_min: int = 30
|
||||
|
||||
# --- 数据保留 ---
|
||||
# 访问日志保留天数,超过自动清理(0 = 永久保留,不推荐)
|
||||
access_log_retention_days: int = 90
|
||||
|
||||
# --- CORS ---
|
||||
frontend_origin: str = "http://localhost:5173"
|
||||
|
||||
|
||||
@@ -83,6 +83,13 @@ class RateLimitedError(AppError):
|
||||
message = "请求过于频繁,请稍后再试。"
|
||||
|
||||
|
||||
class LinkExpiredError(AppError):
|
||||
"""AI 链接已过期(HTTP 410 Gone)。"""
|
||||
status_code = status.HTTP_410_GONE
|
||||
code = "LINK_EXPIRED"
|
||||
message = "链接已失效,请重新生成。"
|
||||
|
||||
|
||||
# --- 上传与配额 ---
|
||||
|
||||
|
||||
|
||||
+53
-1
@@ -18,7 +18,7 @@ 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, get_session_factory
|
||||
from app.core.errors import register_exception_handlers
|
||||
from app.core.errors import LinkExpiredError, register_exception_handlers
|
||||
from app.core.logging import get_logger, setup_logging
|
||||
|
||||
logger = get_logger(__name__)
|
||||
@@ -40,6 +40,43 @@ def _seed_free_plan() -> None:
|
||||
logger.debug("Free plan already exists (id=%s)", plan.id)
|
||||
|
||||
|
||||
def _cleanup_access_logs() -> int:
|
||||
"""清理超过保留期的访问日志。返回删除条数。"""
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from sqlalchemy import delete
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.models.access_log import AccessLog
|
||||
|
||||
settings = get_settings()
|
||||
retention = settings.access_log_retention_days
|
||||
if retention <= 0:
|
||||
return 0
|
||||
|
||||
cutoff = (datetime.now() - timedelta(days=retention)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
factory = get_session_factory()
|
||||
with factory() as session:
|
||||
result = session.execute(delete(AccessLog).where(AccessLog.accessed_at < cutoff))
|
||||
session.commit()
|
||||
count = result.rowcount or 0
|
||||
if count:
|
||||
logger.info("已清理 %d 条过期访问日志(保留 %d 天)", count, retention)
|
||||
return count
|
||||
|
||||
|
||||
async def _periodic_log_cleanup() -> None:
|
||||
"""每 24 小时清理一次过期访问日志(在线程池执行同步 DB 操作)。"""
|
||||
import asyncio
|
||||
|
||||
while True:
|
||||
await asyncio.sleep(24 * 3600)
|
||||
try:
|
||||
await asyncio.to_thread(_cleanup_access_logs)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("定期清理访问日志失败")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
settings = get_settings()
|
||||
@@ -53,8 +90,18 @@ async def lifespan(app: FastAPI):
|
||||
# Seed:确保 free plan 存在
|
||||
_seed_free_plan()
|
||||
|
||||
# 启动时清理过期访问日志(不阻塞启动)
|
||||
import asyncio
|
||||
|
||||
cleanup_task = asyncio.create_task(_periodic_log_cleanup())
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.to_thread(_cleanup_access_logs), timeout=15)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("启动时清理访问日志未完成(首次部署属正常)")
|
||||
|
||||
yield
|
||||
|
||||
cleanup_task.cancel()
|
||||
dispose_engine()
|
||||
logger.info("Backend shutdown complete")
|
||||
|
||||
@@ -80,6 +127,11 @@ def create_app() -> FastAPI:
|
||||
|
||||
register_exception_handlers(app)
|
||||
|
||||
# 链接过期:HTML 返回友好失效页,JSON/MD/TXT 返回对应格式(HTTP 410)
|
||||
from app.public.routes import link_expired_handler
|
||||
|
||||
app.add_exception_handler(LinkExpiredError, link_expired_handler)
|
||||
|
||||
# 路由挂载
|
||||
app.include_router(health_router, prefix="/api", tags=["health"])
|
||||
app.include_router(auth_router, prefix="/api", tags=["auth"])
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""SQLAlchemy 2.0 基础模型类与 Mixin。"""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import String, text
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
@@ -13,8 +13,11 @@ def generate_uuid() -> str:
|
||||
|
||||
|
||||
def utcnow_iso() -> str:
|
||||
"""返回 UTC 当前时间的 ISO8601 字符串。"""
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
"""当前时间,格式:年-月-日 时:分:秒(容器时区,生产为 Asia/Shanghai)。
|
||||
|
||||
定长格式保证字符串排序 = 时间排序。
|
||||
"""
|
||||
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
|
||||
@@ -60,6 +60,11 @@ class KnowledgeBase(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
nullable=True,
|
||||
comment="token 末 8 位明文,供后台识别",
|
||||
)
|
||||
token_expires_at: Mapped[str | None] = mapped_column(
|
||||
String(32),
|
||||
nullable=True,
|
||||
comment="链接过期时间,NULL=长期有效",
|
||||
)
|
||||
|
||||
# 关系
|
||||
user = relationship("User", back_populates="knowledge_bases", lazy="selectin")
|
||||
|
||||
@@ -23,7 +23,7 @@ from fastapi import Path as PathParam
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.deps import get_db
|
||||
from app.core.errors import NotFoundError, RateLimitedError
|
||||
from app.core.errors import LinkExpiredError, NotFoundError, RateLimitedError
|
||||
from app.core.rate_limit import check_rate_limit
|
||||
from app.core.security import decrypt_token
|
||||
from app.models.document_category import DocumentCategory
|
||||
@@ -33,6 +33,51 @@ from app.services.kb_public_service import KbPublicService
|
||||
router = APIRouter(prefix="/k", tags=["public"])
|
||||
|
||||
|
||||
async def link_expired_handler(request: Request, exc: LinkExpiredError) -> Response:
|
||||
"""链接过期:HTML 路径返回友好失效页,JSON/MD/TXT 返回对应格式的错误信息。"""
|
||||
path = request.url.path
|
||||
message = "链接已失效,请联系分享者重新生成链接。"
|
||||
|
||||
if path.endswith(".json"):
|
||||
return Response(
|
||||
status_code=410,
|
||||
content=json.dumps({"code": "LINK_EXPIRED", "message": message, "detail": None}, ensure_ascii=False),
|
||||
media_type="application/json",
|
||||
)
|
||||
if path.endswith(".md") or path.endswith(".txt"):
|
||||
return PlainTextResponse(content=message, status_code=410, media_type="text/plain; charset=utf-8")
|
||||
|
||||
return HTMLResponse(
|
||||
status_code=410,
|
||||
content=f"""<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>链接已失效</title>
|
||||
<meta name="robots" content="noindex,nofollow,noarchive">
|
||||
<style>
|
||||
body {{ font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif;
|
||||
display: flex; justify-content: center; align-items: center; min-height: 100vh;
|
||||
margin: 0; background: #f5f7fa; color: #333; text-align: center; }}
|
||||
.box {{ background: #fff; padding: 50px 40px; border-radius: 12px;
|
||||
box-shadow: 0 2px 12px rgba(0,0,0,0.08); max-width: 420px; }}
|
||||
.icon {{ font-size: 52px; margin-bottom: 16px; }}
|
||||
h1 {{ font-size: 20px; margin: 0 0 12px; }}
|
||||
p {{ color: #888; font-size: 14px; line-height: 1.8; margin: 0; }}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="box">
|
||||
<div class="icon">⏰</div>
|
||||
<h1>链接无效或已失效</h1>
|
||||
<p>该知识库链接已过期。<br>请联系分享者重新生成链接后,获取新的访问地址。</p>
|
||||
</div>
|
||||
</body>
|
||||
</html>""",
|
||||
)
|
||||
|
||||
|
||||
def _log_access(db: Session, kb_id: str, path: str, request: Request, doc_id: str | None = None, req_type: str | None = None) -> None:
|
||||
"""记录访问日志(best-effort)。"""
|
||||
try:
|
||||
|
||||
@@ -33,7 +33,14 @@ class KbListResponse(BaseModel):
|
||||
|
||||
|
||||
class KbTokenResponse(BaseModel):
|
||||
"""Token 重置/创建时的完整 URL 响应。"""
|
||||
"""Token 重置/创建/查询时的完整 URL 响应。"""
|
||||
token: str
|
||||
ai_url: str
|
||||
token_hint: str
|
||||
token_hint: str
|
||||
expires_at: str | None = None
|
||||
is_expired: bool = False
|
||||
|
||||
|
||||
class KbSetExpiryRequest(BaseModel):
|
||||
"""设置链接有效期。expires_in_minutes=null 表示长期有效。"""
|
||||
expires_in_minutes: int | None = None
|
||||
@@ -1,6 +1,6 @@
|
||||
"""访问日志服务。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
|
||||
from app.models.access_log import AccessLog
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -24,7 +24,7 @@ class AccessLogService:
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
document_id=document_id,
|
||||
path=path,
|
||||
accessed_at=datetime.now(timezone.utc).isoformat(),
|
||||
accessed_at=datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
user_agent=user_agent[:500] if user_agent else None,
|
||||
request_type=request_type,
|
||||
)
|
||||
|
||||
@@ -156,10 +156,26 @@ class DocumentService:
|
||||
doc = self.get_or_404(doc_id, user)
|
||||
file_size = doc.file_size
|
||||
self._doc_repo.delete(doc)
|
||||
self._delete_doc_files(doc)
|
||||
# 回补配额
|
||||
self._restore_quota(user, file_size)
|
||||
self._session.commit()
|
||||
|
||||
def _delete_doc_files(self, doc: Document) -> None:
|
||||
"""物理删除文档的原始文件与 Markdown 文件(软删后调用)。"""
|
||||
from app.storage.local_storage import get_storage
|
||||
|
||||
storage = get_storage()
|
||||
for key in (doc.storage_path, doc.markdown_path):
|
||||
if not key:
|
||||
continue
|
||||
try:
|
||||
storage.delete(key)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).warning("清理文件失败 key=%s: %s", key, exc)
|
||||
|
||||
def _check_quota(self, user: User, file_size: int) -> None:
|
||||
settings = get_settings()
|
||||
if user.storage_used + file_size > settings.default_storage_quota:
|
||||
|
||||
@@ -4,7 +4,9 @@ HTML/MD/TXT/JSON/搜索全部通过此 Service 获取数据,不各自写查询
|
||||
支持目录树结构和分类过滤。
|
||||
"""
|
||||
|
||||
from app.core.errors import NotFoundError
|
||||
from datetime import datetime
|
||||
|
||||
from app.core.errors import LinkExpiredError, NotFoundError
|
||||
from app.core.security import hash_token
|
||||
from app.models.document import Document
|
||||
from app.models.document_category import DocumentCategory
|
||||
@@ -14,6 +16,8 @@ from app.repositories.kb_repo import KnowledgeBaseRepository
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
_TIME_FMT = "%Y-%m-%d %H:%M:%S"
|
||||
|
||||
|
||||
class KbPublicService:
|
||||
def __init__(self, session: Session) -> None:
|
||||
@@ -22,13 +26,35 @@ class KbPublicService:
|
||||
self._doc_repo = DocumentRepository(session)
|
||||
|
||||
def get_kb_by_token(self, token: str) -> KnowledgeBase:
|
||||
"""通过 token 获取知识库。不存在/禁用/删除 → 404。"""
|
||||
"""通过 token 获取知识库。
|
||||
|
||||
不存在/禁用/删除 → 404;存在但已过期 → 410(LinkExpiredError)。
|
||||
"""
|
||||
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("知识库不存在。")
|
||||
|
||||
if kb.token_expires_at:
|
||||
try:
|
||||
expires = datetime.strptime(kb.token_expires_at, _TIME_FMT)
|
||||
except ValueError:
|
||||
expires = None
|
||||
if expires is not None and datetime.now() > expires:
|
||||
raise LinkExpiredError()
|
||||
return kb
|
||||
|
||||
@staticmethod
|
||||
def is_expired(kb: KnowledgeBase) -> bool:
|
||||
"""判断知识库链接是否已过期(管理端展示用)。"""
|
||||
if not kb.token_expires_at:
|
||||
return False
|
||||
try:
|
||||
expires = datetime.strptime(kb.token_expires_at, _TIME_FMT)
|
||||
except ValueError:
|
||||
return False
|
||||
return datetime.now() > expires
|
||||
|
||||
def get_category_tree(self, kb: KnowledgeBase) -> list[dict]:
|
||||
"""获取目录树(含文档数量)。"""
|
||||
stmt = (
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""知识库服务:CRUD + Token 管理 + 默认目录树。"""
|
||||
|
||||
from app.core.errors import NotFoundError, PermissionDeniedError
|
||||
from app.core.logging import get_logger
|
||||
from app.models.document_category import DocumentCategory
|
||||
from app.models.knowledge_base import KnowledgeBase
|
||||
from app.models.user import User
|
||||
@@ -8,6 +9,8 @@ from app.repositories.kb_repo import KnowledgeBaseRepository
|
||||
from app.services.token_service import TokenService
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class KnowledgeBaseService:
|
||||
def __init__(self, session: Session) -> None:
|
||||
@@ -96,8 +99,35 @@ class KnowledgeBaseService:
|
||||
def delete(self, kb_id: str, user: User) -> None:
|
||||
kb = self.get_or_404(kb_id, user)
|
||||
self._kb_repo.delete(kb)
|
||||
self._cleanup_kb_files(kb)
|
||||
self._session.commit()
|
||||
|
||||
def _cleanup_kb_files(self, kb: KnowledgeBase) -> None:
|
||||
"""物理删除知识库下所有文档的原始文件与 Markdown 文件(软删后调用)。
|
||||
|
||||
文件删除失败不阻塞删除流程(记录日志,可由后续清理兜底)。
|
||||
"""
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.models.document import Document
|
||||
from app.storage.local_storage import get_storage
|
||||
|
||||
stmt = select(Document).where(Document.knowledge_base_id == kb.id)
|
||||
docs = list(self._session.scalars(stmt).all())
|
||||
storage = get_storage()
|
||||
removed = 0
|
||||
for doc in docs:
|
||||
for key in (doc.storage_path, doc.markdown_path):
|
||||
if not key:
|
||||
continue
|
||||
try:
|
||||
storage.delete(key)
|
||||
removed += 1
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("清理文件失败 key=%s: %s", key, exc)
|
||||
if removed:
|
||||
logger.info("KB %s 软删,已物理清理 %d 个文件", kb.id, removed)
|
||||
|
||||
def regenerate_token(self, kb_id: str, user: User) -> tuple[KnowledgeBase, str]:
|
||||
"""重新生成 Token。旧链接立即失效。"""
|
||||
kb = self.get_or_404(kb_id, user)
|
||||
@@ -117,6 +147,20 @@ class KnowledgeBaseService:
|
||||
self._session.commit()
|
||||
return kb
|
||||
|
||||
def set_expiry(self, kb_id: str, user: User, expires_in_minutes: int | None) -> KnowledgeBase:
|
||||
"""设置链接有效期。None = 长期有效;负数表示已过期(测试用)。"""
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
kb = self.get_or_404(kb_id, user)
|
||||
if expires_in_minutes is None:
|
||||
kb.token_expires_at = None
|
||||
else:
|
||||
kb.token_expires_at = (datetime.now() + timedelta(minutes=expires_in_minutes)).strftime(
|
||||
"%Y-%m-%d %H:%M:%S"
|
||||
)
|
||||
self._session.commit()
|
||||
return kb
|
||||
|
||||
def get_full_token(self, kb: KnowledgeBase) -> str | None:
|
||||
"""解密 token 原文(供后台显示完整链接)。"""
|
||||
if kb.token_encrypted:
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""链接有效期功能测试。"""
|
||||
|
||||
import io
|
||||
import zipfile
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.main import app
|
||||
|
||||
|
||||
def _client() -> TestClient:
|
||||
return TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
|
||||
def _setup_kb(client: TestClient) -> tuple[str, str]:
|
||||
"""注册+创建知识库,返回 (kb_id, token)。"""
|
||||
client.post("/api/auth/register", json={
|
||||
"username": "expiry_user", "email": "expiry@example.com", "password": "password123",
|
||||
})
|
||||
resp = client.post("/api/knowledge-bases", json={"name": "Expiry Test KB"})
|
||||
kb_id = resp.json()["id"]
|
||||
resp = client.get(f"/api/knowledge-bases/{kb_id}/link")
|
||||
return kb_id, resp.json()["token"]
|
||||
|
||||
|
||||
def test_default_permanent() -> None:
|
||||
with _client() as client:
|
||||
kb_id, token = _setup_kb(client)
|
||||
resp = client.get(f"/k/{token}")
|
||||
assert resp.status_code == 200
|
||||
# link 接口返回长期有效
|
||||
resp = client.get(f"/api/knowledge-bases/{kb_id}/link")
|
||||
assert resp.json()["expires_at"] is None
|
||||
assert resp.json()["is_expired"] is False
|
||||
|
||||
|
||||
def test_set_expiry_future() -> None:
|
||||
with _client() as client:
|
||||
kb_id, token = _setup_kb(client)
|
||||
resp = client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": 10})
|
||||
assert resp.status_code == 200
|
||||
# 未过期仍可访问
|
||||
assert client.get(f"/k/{token}").status_code == 200
|
||||
# link 接口返回过期时间
|
||||
resp = client.get(f"/api/knowledge-bases/{kb_id}/link")
|
||||
assert resp.json()["expires_at"] is not None
|
||||
assert resp.json()["is_expired"] is False
|
||||
|
||||
|
||||
def test_expired_link_html() -> None:
|
||||
with _client() as client:
|
||||
kb_id, token = _setup_kb(client)
|
||||
# 设为过去时间(负数分钟)→ 立即过期
|
||||
client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": -1})
|
||||
resp = client.get(f"/k/{token}")
|
||||
assert resp.status_code == 410
|
||||
assert "已失效" in resp.text
|
||||
assert "noindex" in resp.text
|
||||
|
||||
|
||||
def test_expired_link_json_md_txt() -> None:
|
||||
with _client() as client:
|
||||
kb_id, token = _setup_kb(client)
|
||||
client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": -1})
|
||||
|
||||
resp = client.get(f"/k/{token}.json")
|
||||
assert resp.status_code == 410
|
||||
assert resp.json()["code"] == "LINK_EXPIRED"
|
||||
|
||||
resp = client.get(f"/k/{token}.md")
|
||||
assert resp.status_code == 410
|
||||
assert "已失效" in resp.text
|
||||
|
||||
resp = client.get(f"/k/{token}.txt")
|
||||
assert resp.status_code == 410
|
||||
|
||||
|
||||
def test_expired_doc_page() -> None:
|
||||
"""过期后文档页也不可访问。"""
|
||||
with _client() as client:
|
||||
kb_id, token = _setup_kb(client)
|
||||
client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": -1})
|
||||
resp = client.get(f"/k/{token}/doc/whatever-token")
|
||||
assert resp.status_code == 410
|
||||
|
||||
|
||||
def test_restore_permanent() -> None:
|
||||
"""过期后重新设为长期有效可恢复访问。"""
|
||||
with _client() as client:
|
||||
kb_id, token = _setup_kb(client)
|
||||
client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": -1})
|
||||
assert client.get(f"/k/{token}").status_code == 410
|
||||
# 恢复
|
||||
client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": None})
|
||||
assert client.get(f"/k/{token}").status_code == 200
|
||||
|
||||
|
||||
def test_regenerate_resets_expiry() -> None:
|
||||
with _client() as client:
|
||||
kb_id, old_token = _setup_kb(client)
|
||||
client.post(f"/api/knowledge-bases/{kb_id}/set-expiry",
|
||||
json={"expires_in_minutes": -1})
|
||||
# 重新生成 → 新链接长期有效
|
||||
resp = client.post(f"/api/knowledge-bases/{kb_id}/regenerate-token")
|
||||
new_token = resp.json()["token"]
|
||||
assert new_token != old_token
|
||||
assert client.get(f"/k/{new_token}").status_code == 200
|
||||
resp = client.get(f"/api/knowledge-bases/{kb_id}/link")
|
||||
assert resp.json()["expires_at"] is None
|
||||
|
||||
|
||||
def _make_docx_bytes() -> bytes:
|
||||
buf = io.BytesIO()
|
||||
with zipfile.ZipFile(buf, "w") as zf:
|
||||
zf.writestr("[Content_Types].xml", '<?xml version="1.0"?><Types></Types>')
|
||||
zf.writestr("_rels/.rels", '<?xml version="1.0"?><Relationships></Relationships>')
|
||||
zf.writestr("word/document.xml", '<?xml version="1.0"?><w:document xmlns:w="http://schemas.openxmlformats.org/wordprocessingml/2006/main"><w:body><w:p><w:r><w:t>Test</w:t></w:r></w:p></w:body></w:document>')
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def test_doc_delete_removes_files() -> None:
|
||||
"""删除文档后物理文件应被清理。"""
|
||||
with _client() as client:
|
||||
client.post("/api/auth/register", json={
|
||||
"username": "cleanup_user", "email": "cleanup@example.com", "password": "password123",
|
||||
})
|
||||
resp = client.post("/api/knowledge-bases", json={"name": "Cleanup KB"})
|
||||
kb_id = resp.json()["id"]
|
||||
|
||||
files = {"file": ("t.docx", _make_docx_bytes(),
|
||||
"application/vnd.openxmlformats-officedocument.wordprocessingml.document")}
|
||||
client.post("/api/documents/upload", data={"kb_id": kb_id}, files=files)
|
||||
|
||||
resp = client.get(f"/api/documents?kb_id={kb_id}")
|
||||
doc = resp.json()["items"][0]
|
||||
storage_path = doc["storage_path"] if "storage_path" in doc else None
|
||||
doc_id = doc["id"]
|
||||
|
||||
# 删除文档
|
||||
resp = client.delete(f"/api/documents/{doc_id}")
|
||||
assert resp.status_code == 204
|
||||
|
||||
# 验证物理文件已删除
|
||||
from app.core.config import get_settings
|
||||
from pathlib import Path
|
||||
# 通过数据库查询 storage_path(响应里没有)
|
||||
from app.core.db import get_session_factory
|
||||
from app.models.document import Document
|
||||
factory = get_session_factory()
|
||||
with factory() as session:
|
||||
d = session.get(Document, doc_id)
|
||||
storage_path = d.storage_path
|
||||
if storage_path:
|
||||
full = Path(get_settings().storage_root_path) / storage_path
|
||||
assert not full.exists(), f"文件未被清理: {full}"
|
||||
|
||||
|
||||
def test_kb_delete_removes_doc_files() -> None:
|
||||
"""删除知识库后其下所有文档文件应被清理。"""
|
||||
with _client() as client:
|
||||
client.post("/api/auth/register", json={
|
||||
"username": "kb_cleanup_user", "email": "kbcleanup@example.com", "password": "password123",
|
||||
})
|
||||
resp = client.post("/api/knowledge-bases", json={"name": "KB Cleanup KB"})
|
||||
kb_id = resp.json()["id"]
|
||||
|
||||
files = {"file": ("t.docx", _make_docx_bytes(),
|
||||
"application/vnd.openxmlformats-officedocument.wordprocessingml.document")}
|
||||
client.post("/api/documents/upload", data={"kb_id": kb_id}, files=files)
|
||||
|
||||
from app.core.db import get_session_factory
|
||||
from app.models.document import Document
|
||||
from app.core.config import get_settings
|
||||
from pathlib import Path
|
||||
|
||||
factory = get_session_factory()
|
||||
with factory() as session:
|
||||
docs = list(session.query(Document).filter_by(knowledge_base_id=kb_id).all())
|
||||
paths = [d.storage_path for d in docs if d.storage_path]
|
||||
|
||||
# 删除知识库
|
||||
resp = client.delete(f"/api/knowledge-bases/{kb_id}")
|
||||
assert resp.status_code == 204
|
||||
|
||||
root = Path(get_settings().storage_root_path)
|
||||
for p in paths:
|
||||
assert not (root / p).exists(), f"文件未被清理: {p}"
|
||||
Reference in New Issue
Block a user