diff --git a/.env.example b/.env.example index 78879d5..6e08560 100644 --- a/.env.example +++ b/.env.example @@ -72,3 +72,29 @@ CACHE_TTL=300 # --- 分块参数 --- # chunk 超长二次切分阈值(字符数) CHUNK_MAX_CHARS=800 + +# --- 认证(JWT)--- +# JWT 签名密钥,生产环境必须设置为足够长的随机串;留空则启动时自动生成(仅开发用) +JWT_SECRET_KEY= +# JWT 签名算法 +JWT_ALGORITHM=HS256 +# token 有效期(分钟),默认 24 小时 +JWT_EXPIRE_MINUTES=1440 +# 是否开放 POST /api/v1/auth/register 注册接口 +AUTH_REGISTER_ENABLED=true +# 启动时自动创建的默认管理员用户名 +DEFAULT_ADMIN_USERNAME=admin +# 默认管理员密码,留空则不创建默认管理员 +DEFAULT_ADMIN_PASSWORD= + +# --- 检索结果 AI 总结 --- +# 参与总结的最大 hit 条数(控制 prompt 长度) +RESULT_SUMMARY_MAX_HITS=5 + +# --- 文件上传 --- +# 原始文件保存目录(相对路径以工作目录为基) +UPLOAD_DIR=./uploads +# 单文件大小上限(MB) +UPLOAD_MAX_SIZE_MB=20 +# 允许上传的扩展名(逗号分隔,含点号,全小写) +UPLOAD_ALLOWED_EXTENSIONS=.txt,.md,.html,.htm,.pdf,.docx diff --git a/.trae/specs/add-file-upload-support/checklist.md b/.trae/specs/add-file-upload-support/checklist.md new file mode 100644 index 0000000..c14c13f --- /dev/null +++ b/.trae/specs/add-file-upload-support/checklist.md @@ -0,0 +1,72 @@ +# Checklist + +## 多格式文件文本提取 + +- [x] `app/core/file_parser.py` 存在并实现 `parse_file(filename, content)` 与 `supported_extensions()` +- [x] 支持 .txt/.md(UTF-8 解码 errors=replace) +- [x] 支持 .html/.htm(html.parser 剥离 script/style 后提取可见文本) +- [x] 支持 .pdf(pypdf 逐页 extract_text 拼接) +- [x] 支持 .docx(python-docx 段落文本拼接) +- [x] 未识别扩展名抛 `ValueError("不支持的文件类型: {ext}")` +- [x] 解析异常包装为 `ValueError("文件解析失败: {detail}")` 并保留原异常链 +- [x] `tests/test_file_parser.py` 13 例覆盖各格式 + 异常路径 + +## 文件上传入库端点 + +- [x] `POST /api/v1/documents/upload` 端点存在,接收 multipart/form-data(file + title/source/metadata) +- [x] 扩展名校验基于 `settings.upload_allowed_extensions`,失败抛 code=1001 含扩展名 +- [x] 大小校验基于 `settings.upload_max_size_mb`,失败抛 code=1001 含「大小上限」 +- [x] 调用 `parse_file` 提取文本,失败抛 code=1001 +- [x] 提取文本为空抛 code=1001 含「无法从文件提取文本」 +- [x] 落盘到 `settings.upload_dir`,失败仅 warning 不阻塞入库 +- [x] 落盘元数据 raw_file_path/original_filename/original_size_bytes 写入 DocumentInput.metadata +- [x] 用户传入 metadata JSON 解析合并(落盘元数据优先级更高,用 setdefault) +- [x] 默认 title 取文件名 stem,默认 source 取 `file:{原文件名}` +- [x] 复用现有 IngestTaskManager.submit 提交异步入库流水线 +- [x] 返回 202 + `{code:0, data:{task_id, status:"pending", saved_path}}` +- [x] `tests/test_document_upload_api.py` 10 例覆盖合法上传/不支持扩展名/超大/空文本/损坏 PDF/落盘失败降级/metadata JSON/默认 title/source + +## 配置项 + +- [x] `app/config.py` Settings 新增 `upload_dir: str = "./uploads"` +- [x] `app/config.py` Settings 新增 `upload_max_size_mb: int = 20` +- [x] `app/config.py` Settings 新增 `upload_allowed_extensions: str = ".txt,.md,.html,.htm,.pdf,.docx"` + +## 依赖声明 + +- [x] `pyproject.toml` dependencies 含 `python-multipart>=0.0.20` +- [x] `pyproject.toml` dependencies 含 `pypdf>=5.1.0` +- [x] `pyproject.toml` dependencies 含 `python-docx>=1.1.2` + +## 应用启动初始化(Task 4.1) + +- [x] `app/main.py` lifespan 在 `ensure_default_admin` 之后调用 `Path(settings.upload_dir).mkdir(parents=True, exist_ok=True)` +- [x] mkdir 失败仅 warning 不阻塞启动 + +## 环境变量模板(Task 4.2) + +- [x] `.env.example` 末尾含 `--- 文件上传 ---` 段落 +- [x] 含 `UPLOAD_DIR=./uploads` +- [x] 含 `UPLOAD_MAX_SIZE_MB=20` +- [x] 含 `UPLOAD_ALLOWED_EXTENSIONS=.txt,.md,.html,.htm,.pdf,.docx` + +## 管理后台 UI(Task 5) + +- [x] `admin.html` `#section-ingest` 区块含「文件上传」子表单 +- [x] file input 含 `accept=".txt,.md,.html,.htm,.pdf,.docx"` +- [x] 可选 title 输入框 +- [x] 可选 source 输入框 +- [x] 提交按钮 id=`btn-upload-submit` +- [x] submit 监听构造 FormData 调用 `/api/v1/documents/upload` +- [x] 复用 `renderIngestHeader` + `startIngestPolling` 展示任务进度 +- [x] 提交期间禁用按钮并显示「上传中…」 +- [x] 完成(成功/失败)后恢复按钮文案为「上传入库」 + +## 文档同步(Task 6) + +- [x] `CLAUDE.md` API 清单含 `POST /api/v1/documents/upload` 行 +- [x] `CLAUDE.md` 「文档入库与三级总结」节简述文件上传通道 + +## 全量回归(Task 7) + +- [x] `uv run pytest` 全绿(289 passed,含新增 23 测试 + 既有 266 测试套件) diff --git a/.trae/specs/add-file-upload-support/spec.md b/.trae/specs/add-file-upload-support/spec.md new file mode 100644 index 0000000..ea4e28d --- /dev/null +++ b/.trae/specs/add-file-upload-support/spec.md @@ -0,0 +1,90 @@ +# 多格式文件上传入库 Spec + +## Why + +QMDSearch 现有入库接口 `POST /api/v1/documents` 仅接受 JSON 文本字段,用户必须先将文件内容贴成 text 才能入库;对 .md/.html/.pdf/.docx 等常见知识载体缺少直接通道。引入文件上传端点可减少入库摩擦、保留原始文件元数据、让管理后台一站式完成「上传→解析→入库→轮询任务」全流程。 + +## What Changes + +- 新增 `app/core/file_parser.py`:按扩展名分发解析器(.txt/.md/.html/.htm/.pdf/.docx),统一异常包装。 +- 新增 `POST /api/v1/documents/upload` 端点:multipart/form-data 接收文件 + 可选 title/source/metadata,校验扩展名/大小→提取文本→落盘→复用现有 IngestTaskManager 提交异步入库流水线,返回 202 + task_id + saved_path。 +- 配置项:`upload_dir` / `upload_max_size_mb` / `upload_allowed_extensions` 注入到 `Settings`。 +- 应用启动 lifespan 中 `mkdir -p upload_dir`(失败仅 warning 不阻塞启动)。 +- 依赖:`pyproject.toml` 增加 `python-multipart` / `pypdf` / `python-docx`。 +- 环境变量模板 `.env.example` 增加 `--- 文件上传 ---` 段落。 +- 管理后台 `admin.html` 在入库区块新增「文件上传」子表单:file input + 可选 title/source + 提交按钮,构造 FormData 调用 upload 端点,复用 `ingestPollTimer` 轮询 UI;提交期间禁用按钮。 +- `CLAUDE.md` API 清单加入 `POST /api/v1/documents/upload`,并在「文档入库与三级总结」节简述文件上传通道。 +- 单元测试:`tests/test_file_parser.py`(13 例覆盖各格式与异常路径)、`tests/test_document_upload_api.py`(10 例覆盖合法上传/不支持扩展名/超大/空文本/损坏 PDF/落盘失败降级)。 + +## Impact + +- Affected specs: 无(首次新增能力,不修改既有 spec) +- Affected code: + - 新增:`app/core/file_parser.py`、`tests/test_file_parser.py`、`tests/test_document_upload_api.py` + - 修改:`app/api/v1/document.py`、`app/config.py`、`app/main.py`、`app/static/admin.html`、`pyproject.toml`、`.env.example`、`CLAUDE.md` + +## ADDED Requirements + +### Requirement: 多格式文件文本提取 + +系统 SHALL 提供 `parse_file(filename, content)` 函数,按扩展名分发到对应解析器:.txt/.md 走 UTF-8 解码(errors=replace 兜底);.html/.htm 走标准库 html.parser 剥离 script/style 后提取可见文本;.pdf 走 pypdf 逐页 extract_text 拼接;.docx 走 python-docx 段落文本拼接(不含表格/页眉页脚)。未识别扩展名 SHALL 抛 `ValueError("不支持的文件类型: {ext}")`;解析异常 SHALL 包装为 `ValueError("文件解析失败: {detail}")` 并保留原异常链。 + +#### Scenario: 支持的扩展名提取成功 +- **WHEN** 调用 `parse_file("notes.md", b"# Hello\n\nWorld")` +- **THEN** 返回字符串 `"# Hello\n\nWorld"` + +#### Scenario: 不支持的扩展名 +- **WHEN** 调用 `parse_file("data.xlsx", b"binary")` +- **THEN** 抛 `ValueError`,message 含 "不支持的文件类型: .xlsx" + +#### Scenario: 损坏 PDF +- **WHEN** 调用 `parse_file("bad.pdf", b"not a real pdf")` +- **THEN** 抛 `ValueError`,message 含 "文件解析失败" + +### Requirement: 文件上传入库端点 + +系统 SHALL 提供 `POST /api/v1/documents/upload` 端点,接收 multipart/form-data(file + 可选 title/source/metadata JSON 字符串),返回 `202` + `{code:0, data:{task_id, status:"pending", saved_path}}`。端点流程 SHALL 依次:扩展名校验(基于 `settings.upload_allowed_extensions`)→ 大小校验(`settings.upload_max_size_mb`)→ `parse_file` 提取文本→ 落盘到 `settings.upload_dir`(失败仅 warning 不阻塞)→ 合并用户 metadata(落盘元数据 raw_file_path/original_filename/original_size_bytes 优先级更高)→ 默认 title 取文件名 stem、默认 source 取 `file:{原文件名}` → 提交 IngestTaskManager → 返回 202。 + +#### Scenario: 合法 .md 上传 +- **WHEN** POST /api/v1/documents/upload 携带 file=notes.md(内容 "# Hello\n\nWorld")、title="我的笔记"、source="manual" +- **THEN** 返回 202,data.task_id 非空、data.status=="pending"、data.saved_path 指向已存在文件;任务管理器收到 DocumentInput,text 与文件解码一致、metadata.original_filename=="notes.md" + +#### Scenario: 不支持扩展名 +- **WHEN** POST 上传 data.xlsx +- **THEN** 返回 code=1001、message 含 ".xlsx",任务管理器未收到任何提交 + +#### Scenario: 超大文件 +- **WHEN** upload_max_size_mb=1 且上传 2MB txt +- **THEN** 返回 code=1001、message 含 "大小上限" + +#### Scenario: 解析后为空文本 +- **WHEN** 上传 blank.txt 内容仅空白 +- **THEN** 返回 code=1001、message 含 "无法从文件提取文本" + +#### Scenario: 落盘失败降级 +- **WHEN** upload_dir 指向已存在的文件(mkdir 失败) +- **THEN** 返回 202,saved_path=="",DocumentInput.metadata 不含 raw_file_path/original_filename + +### Requirement: 应用启动初始化 upload_dir + +应用 lifespan SHALL 在 `ensure_default_admin` 之后尝试 `Path(settings.upload_dir).mkdir(parents=True, exist_ok=True)`,失败仅 warning 不阻塞启动。 + +### Requirement: 管理后台文件上传 UI + +`admin.html` 入库区块 SHALL 在现有「文本入库」表单旁新增「文件上传」子表单,包含:file input(accept 当前支持的扩展名)+ 可选标题输入框 + 可选来源输入框 + 提交按钮。提交时构造 FormData 调用 `/api/v1/documents/upload`,复用现有 `ingestPollTimer` 与 ingest-result 渲染逻辑展示任务进度与最终结果。提交期间 SHALL 禁用提交按钮并显示「上传中…」文案。 + +#### Scenario: 用户上传文件 +- **WHEN** 用户选择 notes.md 文件、填标题、点提交 +- **THEN** 按钮禁用显示「上传中…」,FormData 含 file/title,请求返回 202 后开始轮询 task 状态直到 done/failed/timeout + +### Requirement: 配置项与环境变量 + +`Settings` SHALL 新增三项配置:`upload_dir: str = "./uploads"`、`upload_max_size_mb: int = 20`、`upload_allowed_extensions: str = ".txt,.md,.html,.htm,.pdf,.docx"`。`.env.example` SHALL 增加 `--- 文件上传 ---` 段落包含对应环境变量 `UPLOAD_DIR` / `UPLOAD_MAX_SIZE_MB` / `UPLOAD_ALLOWED_EXTENSIONS`。 + +### Requirement: 依赖声明 + +`pyproject.toml` dependencies SHALL 增加:`python-multipart>=0.0.20`(FastAPI 文件上传依赖)、`pypdf>=5.1.0`(PDF 解析)、`python-docx>=1.1.2`(DOCX 解析)。 + +### Requirement: 文档同步 + +`CLAUDE.md` API 清单 SHALL 加入 `POST /api/v1/documents/upload` 行;「文档入库与三级总结」节 SHALL 简述文件上传通道(multipart 入口→parse_file 提取→复用现有 IngestTaskManager)。 diff --git a/.trae/specs/add-file-upload-support/tasks.md b/.trae/specs/add-file-upload-support/tasks.md new file mode 100644 index 0000000..3476155 --- /dev/null +++ b/.trae/specs/add-file-upload-support/tasks.md @@ -0,0 +1,37 @@ +# Tasks + +## 已完成(前期工作) + +- [x] Task 1: 新增 `app/core/file_parser.py` 多格式文本提取 + - [x] SubTask 1.1: 实现 _parse_text/_parse_html/_parse_pdf/_parse_docx 解析器与 _PARSERS 注册表 + - [x] SubTask 1.2: 实现 `parse_file(filename, content)` 入口与异常包装(不支持/解析失败) + - [x] SubTask 1.3: 实现 `supported_extensions()` 辅助函数 +- [x] Task 2: 单元测试 `tests/test_file_parser.py`(13 例覆盖各格式 + 异常路径) +- [x] Task 3: 新增 `POST /api/v1/documents/upload` 端点 + 测试 + - [x] SubTask 3.1: 在 `app/api/v1/document.py` 增加 upload 端点(扩展名/大小/解析/落盘/metadata/默认 title-source/提交 IngestTaskManager) + - [x] SubTask 3.2: 增加 `_allowed_extensions()` 辅助函数 + - [x] SubTask 3.3: 在 `app/config.py` 增加 `upload_dir` / `upload_max_size_mb` / `upload_allowed_extensions` + - [x] SubTask 3.4: `tests/test_document_upload_api.py` 10 例测试(合法上传/不支持扩展名/超大/空文本/损坏 PDF/落盘失败降级/metadata JSON/默认 title/source) + - [x] SubTask 3.5: `pyproject.toml` dependencies 增加 `python-multipart>=0.0.20` / `pypdf>=5.1.0` / `python-docx>=1.1.2` + +## 剩余任务 + +- [x] Task 4: 应用启动初始化 + 环境变量模板 + - [x] SubTask 4.1: 在 `app/main.py` lifespan 的 `ensure_default_admin` 之后增加 `Path(settings.upload_dir).mkdir(parents=True, exist_ok=True)`,异常仅 warning 不阻塞启动 + - [x] SubTask 4.2: 在 `.env.example` 末尾增加 `--- 文件上传 ---` 段落,包含 `UPLOAD_DIR` / `UPLOAD_MAX_SIZE_MB` / `UPLOAD_ALLOWED_EXTENSIONS` 三项 +- [x] Task 5: 管理后台 `admin.html` 文件上传 UI + - [x] SubTask 5.1: 在 `#section-ingest` 现有 `#ingest-form` 之后新增「文件上传」子表单,含 file input(accept=".txt,.md,.html,.htm,.pdf,.docx")+ 可选 title + 可选 source + 提交按钮 `#btn-upload-submit` + - [x] SubTask 5.2: 新增 `#upload-form` submit 监听:构造 FormData(file/title/source)→ 调用 `/api/v1/documents/upload` → 复用 `renderIngestHeader` + `startIngestPolling` 展示任务进度 + - [x] SubTask 5.3: 提交期间禁用 `#btn-upload-submit` 显示「上传中…」,请求完成(成功/失败)后恢复按钮文案 +- [x] Task 6: 文档同步 `CLAUDE.md` + - [x] SubTask 6.1: API 清单表格在 `POST /api/v1/documents` 行之后增加 `POST /api/v1/documents/upload` 行(说明:multipart 文件上传入库,202 异步) + - [x] SubTask 6.2: 「文档入库与三级总结」节增加简述:文件上传通道(multipart 入口→parse_file 提取→复用现有 IngestTaskManager→同一任务查询端点) +- [x] Task 7: 全量 pytest 验证回归 + - [x] SubTask 7.1: 执行 `uv run pytest` 全绿(289 passed,含新增 23 测试 + 既有 266 测试套件) + +# Task Dependencies + +- Task 4 独立于 Task 5/6,可并行 +- Task 5(admin.html UI)依赖 Task 3 已完成(端点已存在)✓ +- Task 6 独立 +- Task 7 依赖 Task 4/5/6 全部完成 diff --git a/CLAUDE.md b/CLAUDE.md index a815d23..59f5ce8 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -85,6 +85,7 @@ QMDSearch/ | GET | `/api/v1/health` | 健康检查 | | POST | `/api/v1/search` | 分层检索 | | POST | `/api/v1/documents` | 文档入库(202 异步入库,返回 task_id) | +| POST | `/api/v1/documents/upload` | multipart 文件上传入库(202 异步,支持 .txt/.md/.html/.htm/.pdf/.docx) | | GET | `/api/v1/documents/tasks/{task_id}` | 入库任务状态查询(done 附 result,failed 附 error) | | GET | `/api/v1/knowledge/categories` | 知识分类类目集 | | GET | `/api/v1/knowledge/stats` | 统计(四层点数 + 类目分布 + uncategorized 数) | @@ -131,6 +132,8 @@ QMDSearch/ 不足三级 → 2.5级回退 ``` +**文件上传通道**: 除 JSON 文本入库外,`POST /api/v1/documents/upload` 提供 multipart 文件上传入口,支持 .txt/.md/.html/.htm/.pdf/.docx;服务端调用 `app/core/file_parser.py` 按扩展名提取纯文本后,复用上述同一入库流水线;原始文件落盘到 `settings.upload_dir`(落盘失败仅告警不阻塞入库),任务状态同样经 `GET /api/v1/documents/tasks/{task_id}` 查询。 + 1. **文本提取**: 从文件中提取纯文本内容 2. **三级总结**: 调用 Ollama 依次生成 L1/L2/L3 总结;其中 L2 大纲优先使用文档原生标题树(不调 LLM),无结构文本(标题数 < 2)回退 LLM 生成 3. **分类判定**: 根据 L1 总结将文档分配到 taxonomy 类目,输出主类目 + 附加标签 + 置信度;LLM 输出解析失败或置信度低于阈值时归 uncategorized 兜底(低置信时候选类目名保留进 tags 供软召回) diff --git a/app/api/v1/auth.py b/app/api/v1/auth.py new file mode 100644 index 0000000..5129ed6 --- /dev/null +++ b/app/api/v1/auth.py @@ -0,0 +1,62 @@ +"""认证 API:POST /auth/login、POST /auth/register、GET /auth/me""" + +from typing import Any + +import structlog +from fastapi import APIRouter, Depends + +from app.api.response import ApiError, ok +from app.config import settings +from app.core.auth import ( + ERR_REGISTER_DISABLED, + UserStore, + create_access_token, + get_current_user, + get_user_store, +) +from app.models.auth import ( + AuthUser, + LoginRequest, + RegisterRequest, + TokenResponse, +) + +logger = structlog.get_logger() + +router = APIRouter(prefix="/api/v1", tags=["auth"]) + + +def _build_token_response(user, store: UserStore) -> dict[str, Any]: + """签发 token 并构造统一响应 data""" + token, expires_in = create_access_token(user.username, user.role) + auth_user = AuthUser(username=user.username, role=user.role, created_at=user.created_at) + resp = TokenResponse(access_token=token, expires_in=expires_in, user=auth_user) + return resp.model_dump(mode="json") + + +@router.post("/auth/login") +async def login(req: LoginRequest) -> dict[str, Any]: + """用户名密码登录,返回 JWT access_token""" + store = get_user_store() + user = await store.authenticate(req.username, req.password) + logger.info("用户登录成功", username=user.username) + return ok(_build_token_response(user, store)) + + +@router.post("/auth/register") +async def register(req: RegisterRequest) -> dict[str, Any]: + """注册新用户(role=user),注册后自动签发 token + + 受 settings.auth_register_enabled 控制,关闭时返回 ERR_REGISTER_DISABLED。 + """ + if not settings.auth_register_enabled: + raise ApiError(ERR_REGISTER_DISABLED, "注册已关闭") + store = get_user_store() + user = await store.create(req.username, req.password, role="user") + return ok(_build_token_response(user, store)) + + +@router.get("/auth/me") +async def me(user: AuthUser = Depends(get_current_user)) -> dict[str, Any]: + """返回当前登录用户信息""" + return ok(user.model_dump(mode="json")) diff --git a/app/api/v1/document.py b/app/api/v1/document.py index 2f0822a..01c927e 100644 --- a/app/api/v1/document.py +++ b/app/api/v1/document.py @@ -1,13 +1,19 @@ -"""文档 API:POST /api/v1/documents 异步入库 + 任务查询 + 文档管理(列表/详情/删除)""" +"""文档 API:POST /api/v1/documents 异步入库 + 任务查询 + 文档管理(列表/详情/删除)+ 文件上传入库""" +import json +import uuid +from datetime import UTC, datetime +from pathlib import Path from typing import Any import structlog -from fastapi import APIRouter, Query +from fastapi import APIRouter, Depends, File, Form, Query, UploadFile from fastapi.responses import JSONResponse from app.api.response import ApiError, ok -from app.config import Settings +from app.config import Settings, settings +from app.core.auth import AuthUser, get_current_user, require_admin +from app.core.file_parser import parse_file from app.core.ingest_tasks import IngestTaskManager from app.core.ingestion import Ingester from app.models.document import DocumentInput @@ -55,8 +61,13 @@ def _get_task_manager() -> IngestTaskManager: return _task_manager +def _allowed_extensions() -> set[str]: + """解析 settings.upload_allowed_extensions 逗号分隔字符串为扩展名集合(全小写、含点号)""" + return {ext.strip().lower() for ext in settings.upload_allowed_extensions.split(",") if ext.strip()} + + @router.post("/documents") -async def ingest_document(doc: DocumentInput) -> JSONResponse: +async def ingest_document(doc: DocumentInput, user: AuthUser = Depends(get_current_user)) -> JSONResponse: """文档入库入口:登记异步任务并返回 202 + task_id,入库结果经任务查询端点获取""" if not doc.text.strip(): raise ApiError(1001, "文档内容不能为空") @@ -64,8 +75,92 @@ async def ingest_document(doc: DocumentInput) -> JSONResponse: return JSONResponse(status_code=202, content=ok({"task_id": task_id, "status": "pending"})) +@router.post("/documents/upload") +async def upload_document( + file: UploadFile = File(..., description="上传的文件(.txt/.md/.html/.htm/.pdf/.docx)"), + title: str = Form(default="", description="可选标题,默认取原文件名去扩展"), + source: str = Form(default="", description="可选来源标识,默认 file:{原文件名}"), + metadata: str = Form(default="", description='可选元数据 JSON 字符串,如 \'{"author":"x"}\''), + user: AuthUser = Depends(get_current_user), +) -> JSONResponse: + """文件上传入库入口:校验 → 提取文本 → 落盘 → 提交异步入库流水线 + + 与 POST /documents 共用同一 IngestTaskManager;任务状态经 + GET /api/v1/documents/tasks/{task_id} 查询。 + """ + original_filename = file.filename or "unnamed" + ext = Path(original_filename).suffix.lower() + + # 1. 扩展名校验 + allowed = _allowed_extensions() + if ext not in allowed: + raise ApiError(1001, f"不支持的文件类型: {ext or '(无扩展名)'}") + + # 2. 读取字节并校验大小 + content = await file.read() + max_bytes = settings.upload_max_size_mb * 1024 * 1024 + if len(content) > max_bytes: + raise ApiError(1001, f"文件超过大小上限: {settings.upload_max_size_mb}MB") + + # 3. 提取文本 + try: + text = parse_file(original_filename, content) + except ValueError as exc: + raise ApiError(1001, str(exc)) from exc + if not text.strip(): + raise ApiError(1001, "无法从文件提取文本") + + # 4. 落盘(按 YYYY/MM 日期分片;失败仅 warning,不阻塞入库) + doc_id = uuid.uuid4().hex + saved_path = "" + metadata_dict: dict[str, str] = {} + try: + upload_dir = Path(settings.upload_dir).resolve() + shard_subdir = datetime.now(UTC).strftime("%Y/%m") + target_dir = upload_dir / shard_subdir + target_dir.mkdir(parents=True, exist_ok=True) + target = target_dir / f"{doc_id}_{original_filename}" + target.write_bytes(content) + saved_path = str(target) + metadata_dict.update( + { + "raw_file_path": saved_path, + "original_filename": original_filename, + "original_size_bytes": str(len(content)), + } + ) + logger.info("上传文件已落盘", doc_id=doc_id, saved_path=saved_path, size=len(content)) + except Exception: + logger.warning("上传文件落盘失败,仅做文本入库", doc_id=doc_id, filename=original_filename, exc_info=True) + + # 5. 合并用户传入的 metadata(落盘元数据优先级更高,不与用户键冲突) + if metadata.strip(): + try: + user_meta = json.loads(metadata) + if isinstance(user_meta, dict): + for k, v in user_meta.items(): + metadata_dict.setdefault(str(k), str(v)) + except (json.JSONDecodeError, TypeError): + logger.warning("metadata 不是合法 JSON,已忽略", raw=metadata) + + # 6. 默认 title / source + if not title.strip(): + title = Path(original_filename).stem + if not source.strip(): + source = f"file:{original_filename}" + + # 7. 提交入库流水线 + doc_input = DocumentInput(text=text, title=title, source=source, metadata=metadata_dict) + task_id = await _get_task_manager().submit(doc_input) + + return JSONResponse( + status_code=202, + content=ok({"task_id": task_id, "status": "pending", "saved_path": saved_path}), + ) + + @router.get("/documents/tasks/{task_id}") -async def get_ingest_task(task_id: str) -> dict[str, Any]: +async def get_ingest_task(task_id: str, user: AuthUser = Depends(get_current_user)) -> dict[str, Any]: """查询入库任务状态:含 task_id/status/created_at/updated_at,done 附 result,failed 附 error""" task = await _get_task_manager().get(task_id) if task is None: @@ -77,6 +172,7 @@ async def get_ingest_task(task_id: str) -> dict[str, Any]: async def list_documents( limit: int = Query(default=20, ge=1, le=100), offset: str | None = None, + user: AuthUser = Depends(get_current_user), ) -> dict[str, Any]: """分页列出文档(L1 摘要),返回 items 与下一页游标 next_offset""" try: @@ -88,7 +184,7 @@ async def list_documents( @router.get("/documents/{doc_id}") -async def get_document(doc_id: str) -> dict[str, Any]: +async def get_document(doc_id: str, user: AuthUser = Depends(get_current_user)) -> dict[str, Any]: """获取文档详情:L1 记录 + L2/L3 节点 + chunks 数量""" try: detail = await _get_qdrant().get_doc_detail(doc_id) @@ -101,7 +197,7 @@ async def get_document(doc_id: str) -> dict[str, Any]: @router.delete("/documents/{doc_id}") -async def delete_document(doc_id: str) -> dict[str, Any]: +async def delete_document(doc_id: str, user: AuthUser = Depends(require_admin)) -> dict[str, Any]: """删除文档:四层集合中该 doc_id 的所有点;幂等,不存在也返回成功(删除数全 0)""" try: deleted = await _get_qdrant().delete_by_doc_id(doc_id) diff --git a/app/api/v1/knowledge.py b/app/api/v1/knowledge.py index b9611d2..fe49c8b 100644 --- a/app/api/v1/knowledge.py +++ b/app/api/v1/knowledge.py @@ -4,10 +4,11 @@ from functools import lru_cache from typing import Any import structlog -from fastapi import APIRouter +from fastapi import APIRouter, Depends from app.api.response import ApiError, ok from app.config import settings +from app.core.auth import AuthUser, get_current_user from app.models.knowledge import UNCATEGORIZED, TaxonomyCategory, load_taxonomy from app.services.qdrant import ALL_COLLECTIONS, COLLECTION_L1, QdrantService @@ -36,14 +37,14 @@ def _get_taxonomy() -> list[TaxonomyCategory]: @router.get("/knowledge/categories") -async def list_categories() -> dict[str, Any]: +async def list_categories(user: AuthUser = Depends(get_current_user)) -> dict[str, Any]: """返回完整知识分类类目集""" categories = _get_taxonomy() return ok({"categories": [c.model_dump() for c in categories], "count": len(categories)}) @router.get("/knowledge/stats") -async def knowledge_stats() -> dict[str, Any]: +async def knowledge_stats(user: AuthUser = Depends(get_current_user)) -> dict[str, Any]: """返回四层集合规模与 L1 类目分布统计""" service = _get_qdrant() try: diff --git a/app/api/v1/search.py b/app/api/v1/search.py index e514d96..7eb2ff1 100644 --- a/app/api/v1/search.py +++ b/app/api/v1/search.py @@ -4,9 +4,10 @@ from hashlib import sha256 from typing import Any import structlog -from fastapi import APIRouter +from fastapi import APIRouter, Depends from app.api.response import ApiError, ok +from app.core.auth import AuthUser, get_current_user from app.core.retriever import Retriever from app.models.search import SearchRequest from app.services.redis import get_cache @@ -27,13 +28,13 @@ def _get_retriever() -> Retriever: def _cache_key(request: SearchRequest) -> str: - """检索缓存键:query + top_k 的短哈希(top_k 影响结果集,需参与键计算)""" - digest = sha256((request.query + "|" + str(request.top_k)).encode()).hexdigest()[:16] + """检索缓存键:query + top_k + summarize 的短哈希(三者均影响结果集,需参与键计算)""" + digest = sha256((request.query + "|" + str(request.top_k) + "|" + str(request.summarize)).encode()).hexdigest()[:16] return f"search:{digest}" @router.post("/search") -async def search(request: SearchRequest) -> dict[str, Any]: +async def search(request: SearchRequest, user: AuthUser = Depends(get_current_user)) -> dict[str, Any]: """分层检索入口,返回统一包装的 SearchResponse 先查 Redis 缓存:命中直接返回缓存的响应;未命中走检索流程并回写缓存。 diff --git a/app/config.py b/app/config.py index 774d7ab..bbc765e 100644 --- a/app/config.py +++ b/app/config.py @@ -54,6 +54,27 @@ class Settings(BaseSettings): # 分块参数 chunk_max_chars: int = 800 # chunk 超长二次切分阈值 + # 认证(JWT) + jwt_secret_key: str = "" # JWT 签名密钥,为空时启动自动生成(仅开发,生产必填) + jwt_algorithm: str = "HS256" + jwt_expire_minutes: int = 1440 # token 有效期(分钟),默认 24 小时 + auth_register_enabled: bool = True # 是否开放 POST /auth/register + default_admin_username: str = "admin" # 启动时自动创建的默认管理员用户名 + default_admin_password: str = "" # 默认管理员密码,为空则不创建默认管理员 + + # 文件上传 + upload_dir: str = "./uploads" # 原始文件保存目录(相对路径以工作目录为基) + upload_max_size_mb: int = 20 # 单文件大小上限(MB) + upload_allowed_extensions: str = ".txt,.md,.html,.htm,.pdf,.docx" # 允许上传的扩展名(逗号分隔) + + # PDF OCR(图片型/扫描件降级,pypdf extract_text 为空时触发) + pdf_ocr_enabled: bool = True # 是否启用 OCR 降级(关闭则扫描件按"无法提取文本"拒绝入库) + pdf_ocr_max_pages: int = 30 # 单文件 OCR 页数上限,超过仅前 N 页 + pdf_ocr_dpi: int = 200 # 渲染 DPI(越高越准但越慢,72~300 合理) + + # 检索结果 AI 总结 + result_summary_max_hits: int = 5 # 参与总结的最大 hit 条数(控制 prompt 长度) + model_config = {"env_prefix": "", "case_sensitive": False} diff --git a/app/core/auth.py b/app/core/auth.py new file mode 100644 index 0000000..a6a63c0 --- /dev/null +++ b/app/core/auth.py @@ -0,0 +1,213 @@ +"""JWT 认证核心:密码哈希、token 签发/验签、用户存储、FastAPI 鉴权依赖 + +用户持久化在 Redis(key: auth:user:{username})。登录/注册需 Redis 可用; +鉴权(get_current_user)以 JWT 自包含信息为主,Redis 不可用时降级可用。 +""" + +from datetime import UTC, datetime, timedelta + +import bcrypt +import jwt +import structlog +from fastapi import Depends +from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer +from redis import asyncio as redis_async + +from app.api.response import ApiError +from app.config import settings +from app.models.auth import AuthUser, StoredUser + +logger = structlog.get_logger() + +# 错误码 +ERR_UNAUTHORIZED = 1003 # 未认证(需要登录) +ERR_TOKEN_INVALID = 1005 # token 无效或已过期 +ERR_FORBIDDEN = 1006 # 权限不足 +ERR_USER_EXISTS = 1007 # 用户名已存在 +ERR_BAD_CREDENTIALS = 1008 # 用户名或密码错误 +ERR_REGISTER_DISABLED = 1009 # 注册已关闭 + +# Bearer token 提取器(auto_error=False,统一由 ApiError 处理) +_bearer = HTTPBearer(auto_error=False) + +# 模块级 JWT secret(settings 为空时自动生成,仅开发用) +_auto_secret: str | None = None + +# 模块级用户存储单例 +_user_store: "UserStore | None" = None + + +def _get_secret() -> str: + """获取 JWT 签名密钥,settings 为空时生成一次性随机密钥(仅开发)""" + global _auto_secret + if settings.jwt_secret_key: + return settings.jwt_secret_key + if _auto_secret is None: + import secrets + + _auto_secret = secrets.token_urlsafe(48) + logger.warning( + "JWT_SECRET_KEY 未配置,已自动生成随机密钥(仅开发用,重启后旧 token 失效)" + ) + return _auto_secret + + +def hash_password(password: str) -> str: + """密码 bcrypt 哈希(返回 utf-8 字符串形式的哈希)""" + return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8") + + +def verify_password(password: str, hashed: str) -> bool: + """校验明文密码与哈希是否匹配,异常或不匹配均返回 False""" + try: + return bcrypt.checkpw(password.encode("utf-8"), hashed.encode("utf-8")) + except Exception: + return False + + +def create_access_token(username: str, role: str) -> tuple[str, int]: + """签发 JWT,返回 (token, expires_in_seconds)""" + now = datetime.now(UTC) + expire = now + timedelta(minutes=settings.jwt_expire_minutes) + payload = { + "sub": username, + "role": role, + "iat": int(now.timestamp()), + "exp": int(expire.timestamp()), + } + token = jwt.encode(payload, _get_secret(), algorithm=settings.jwt_algorithm) + return token, settings.jwt_expire_minutes * 60 + + +def decode_token(token: str) -> dict: + """验签并返回 payload,失败抛 ApiError(ERR_TOKEN_INVALID)""" + try: + return jwt.decode(token, _get_secret(), algorithms=[settings.jwt_algorithm]) + except jwt.ExpiredSignatureError as exc: + raise ApiError(ERR_TOKEN_INVALID, "token 已过期") from exc + except jwt.InvalidTokenError as exc: + raise ApiError(ERR_TOKEN_INVALID, "token 无效") from exc + + +class UserStore: + """Redis 用户存储""" + + def __init__(self, redis_url: str | None = None) -> None: + self._redis_url = redis_url or settings.redis_url + self._client: redis_async.Redis | None = None + + def _get_client(self) -> redis_async.Redis: + if self._client is None: + self._client = redis_async.from_url(self._redis_url, decode_responses=True) + return self._client + + @staticmethod + def _key(username: str) -> str: + return f"auth:user:{username}" + + async def get(self, username: str) -> StoredUser | None: + """读取用户,Redis 异常或数据损坏均返回 None""" + try: + raw = await self._get_client().get(self._key(username)) + except Exception: + logger.error("Redis 读取用户失败", username=username, exc_info=True) + return None + if raw is None: + return None + try: + return StoredUser.model_validate_json(raw) + except Exception: + logger.error("用户数据损坏", username=username, exc_info=True) + return None + + async def exists(self, username: str) -> bool: + try: + return bool(await self._get_client().exists(self._key(username))) + except Exception: + logger.error("Redis 检查用户存在性失败", username=username, exc_info=True) + return False + + async def create( + self, username: str, password: str, role: str = "user" + ) -> StoredUser: + """创建用户,已存在抛 ApiError(ERR_USER_EXISTS),Redis 写失败抛 ApiError(2000)""" + if await self.exists(username): + raise ApiError(ERR_USER_EXISTS, f"用户名已存在: {username}") + user = StoredUser( + username=username, + role=role, + created_at=datetime.now(UTC), + hashed_password=hash_password(password), + ) + try: + await self._get_client().set(self._key(username), user.model_dump_json()) + except Exception as exc: + logger.error("Redis 写入用户失败", username=username, exc_info=True) + raise ApiError(2000, f"用户创建失败: {exc}") from exc + logger.info("用户创建成功", username=username, role=role) + return user + + async def authenticate(self, username: str, password: str) -> StoredUser: + """验证用户名密码,失败抛 ApiError(ERR_BAD_CREDENTIALS)""" + user = await self.get(username) + if user is None or not verify_password(password, user.hashed_password): + raise ApiError(ERR_BAD_CREDENTIALS, "用户名或密码错误") + return user + + +def get_user_store() -> UserStore: + """全局 UserStore 单例""" + global _user_store + if _user_store is None: + _user_store = UserStore() + return _user_store + + +async def ensure_default_admin() -> None: + """启动时按配置创建默认管理员(密码为空则跳过)""" + pwd = settings.default_admin_password + if not pwd: + return + store = get_user_store() + try: + if await store.exists(settings.default_admin_username): + return + await store.create(settings.default_admin_username, pwd, role="admin") + logger.info("默认管理员已创建", username=settings.default_admin_username) + except ApiError: + raise + except Exception: + logger.warning("默认管理员创建失败,跳过", exc_info=True) + + +async def get_current_user( + credentials: HTTPAuthorizationCredentials | None = Depends(_bearer), +) -> AuthUser: + """FastAPI 鉴权依赖:解析 Bearer token,返回 AuthUser + + 优先查 Redis 获取最新用户信息(支持删除/改角色生效); + Redis 不可用或用户未命中时降级为 JWT payload 构造 AuthUser,保证鉴权可用。 + """ + if credentials is None or (credentials.scheme or "").lower() != "bearer": + raise ApiError(ERR_UNAUTHORIZED, "未提供认证凭证,请先登录") + payload = decode_token(credentials.credentials) + username = payload.get("sub") + role = payload.get("role", "user") + if not username: + raise ApiError(ERR_TOKEN_INVALID, "token 缺少用户标识") + + stored = await get_user_store().get(username) + if stored is not None: + return AuthUser( + username=stored.username, role=stored.role, created_at=stored.created_at + ) + # Redis 不可用或用户不存在(可能已删除):降级用 JWT payload + logger.warning("用户存储查询未命中,降级使用 JWT payload", username=username) + return AuthUser(username=username, role=role, created_at=datetime.now(UTC)) + + +async def require_admin(user: AuthUser = Depends(get_current_user)) -> AuthUser: + """要求管理员角色""" + if user.role != "admin": + raise ApiError(ERR_FORBIDDEN, "需要管理员权限") + return user diff --git a/app/core/file_parser.py b/app/core/file_parser.py new file mode 100644 index 0000000..796c7e4 --- /dev/null +++ b/app/core/file_parser.py @@ -0,0 +1,204 @@ +"""多格式文件文本提取:按扩展名分发到对应解析器 + +支持的扩展名: +- .txt / .md:UTF-8 解码(errors="replace" 兜底) +- .html / .htm:标准库 html.parser 剥离标签提取可见文本 +- .pdf:pypdf 逐页 extract_text 拼接;文本层为空(扫描件/图片型)时 + 自动降级为 OCR(pypdfium2 渲染 + rapidocr-onnxruntime 识别) +- .docx:python-docx 段落文本拼接(不含表格/页眉页脚) + +未识别扩展名抛 ValueError("不支持的文件类型: {ext}"); +解析异常统一包装为 ValueError("文件解析失败: {detail}"),原异常链式保留。 +""" + +from __future__ import annotations + +import io +from collections.abc import Callable +from html.parser import HTMLParser +from pathlib import Path +from typing import Any + +import structlog + +from app.config import settings + +logger = structlog.get_logger() + +# 模块级懒加载 OCR 引擎单例(首次调用时初始化,避免无扫描件场景白白下载模型) +_ocr_engine: Any | None = None +_ocr_unavailable: bool = False # 标记 OCR 依赖不可用,后续直接跳过避免重复尝试 + + +class _VisibleTextExtractor(HTMLParser): + """HTMLParser 子类:累积可见文本,跳过 script/style 内容""" + + def __init__(self) -> None: + super().__init__(convert_charrefs=True) + self._parts: list[str] = [] + self._skip_depth = 0 # 在 script/style 标签内时 > 0 + + def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None: + if tag.lower() in {"script", "style"}: + self._skip_depth += 1 + + def handle_endtag(self, tag: str) -> None: + if tag.lower() in {"script", "style"} and self._skip_depth > 0: + self._skip_depth -= 1 + + def handle_data(self, data: str) -> None: + if self._skip_depth == 0: + self._parts.append(data) + + def get_text(self) -> str: + # 块级标签间用空格连接,再折叠多余空白 + text = " ".join(self._parts) + return " ".join(text.split()) + + +def _decode_html(content: bytes) -> str: + """HTML 内容解码:优先 utf-8(带 BOM),失败回退 latin-1""" + try: + return content.decode("utf-8-sig") + except UnicodeDecodeError: + return content.decode("latin-1", errors="replace") + + +def _parse_text(content: bytes) -> str: + """UTF-8 解码(errors=replace 兜底),保留原字符""" + return content.decode("utf-8", errors="replace") + + +def _parse_html(content: bytes) -> str: + """HTML 剥离标签,保留可见文本""" + parser = _VisibleTextExtractor() + parser.feed(_decode_html(content)) + parser.close() + return parser.get_text() + + +def _parse_pdf(content: bytes) -> str: + """PDF 解析:优先 pypdf extract_text;文本层为空(扫描件)时降级 OCR + + OCR 流程:pypdfium2 渲染每页为 PIL Image → rapidocr-onnxruntime 识别 → + 拼接每页识别出的文本。受 settings.pdf_ocr_* 控制:开关、最大页数、DPI。 + OCR 依赖未安装或运行异常时降级返回空字符串(由上游 upload 端点拒绝入库)。 + """ + from pypdf import PdfReader + + reader = PdfReader(io.BytesIO(content)) + parts: list[str] = [] + for page in reader.pages: + text = page.extract_text() or "" + if text: + parts.append(text) + text_layer = "\n".join(parts).strip() + + # 文本层非空:直接返回 + if text_layer: + return text_layer + + # 文本层为空 → 尝试 OCR 降级 + if not settings.pdf_ocr_enabled: + return "" + + ocr_text = _ocr_pdf(content, max_pages=settings.pdf_ocr_max_pages, dpi=settings.pdf_ocr_dpi) + return ocr_text + + +def _ocr_pdf(content: bytes, max_pages: int, dpi: int) -> str: + """对扫描件 PDF 跑 OCR:渲染每页 → 识别 → 拼接 + + 返回空字符串的场景:依赖未安装 / 渲染或识别异常 / 无识别结果。 + 任何异常仅告警不抛出,由上游按"无法提取文本"处理。 + """ + global _ocr_engine, _ocr_unavailable + + if _ocr_unavailable: + return "" + + # 1. 懒加载 OCR 引擎 + if _ocr_engine is None: + try: + from rapidocr_onnxruntime import RapidOCR + + _ocr_engine = RapidOCR() + logger.info("PDF OCR 引擎已初始化", dpi=dpi) + except Exception: + _ocr_unavailable = True + logger.warning("OCR 依赖不可用,扫描件 PDF 将无法提取文本", exc_info=True) + return "" + + # 2. 渲染并识别 + try: + import pypdfium2 as pdfium + + scale = max(1.0, dpi / 72.0) + pdf = pdfium.PdfDocument(io.BytesIO(content)) + total = min(len(pdf), max(1, max_pages)) + page_texts: list[str] = [] + for i in range(total): + page = pdf[i] + pil_image = page.render(scale=scale).to_pil() + result, _ = _ocr_engine(pil_image) + if result: + # result: [[box, text, score], ...],按行拼接 + lines = [item[1] for item in result if item and len(item) >= 2 and item[1]] + if lines: + page_texts.append("\n".join(lines)) + pdf.close() + return "\n".join(page_texts).strip() + except Exception as exc: + logger.warning("PDF OCR 失败,降级返回空文本", error=str(exc), exc_info=True) + return "" + + +def _parse_docx(content: bytes) -> str: + """DOCX 段落文本拼接(不含表格/页眉页脚)""" + from docx import Document # type: ignore[import-untyped] + + document = Document(io.BytesIO(content)) + parts = [p.text for p in document.paragraphs if p.text and p.text.strip()] + return "\n".join(parts).strip() + + +# 扩展名 → 解析函数映射(启动时构建,避免每次请求重复构造) +_PARSERS: dict[str, Callable[[bytes], str]] = { + ".txt": _parse_text, + ".md": _parse_text, + ".html": _parse_html, + ".htm": _parse_html, + ".pdf": _parse_pdf, + ".docx": _parse_docx, +} + + +def parse_file(filename: str, content: bytes) -> str: + """按扩展名分发解析器提取文本 + + Args: + filename: 文件名(用于判定扩展名) + content: 文件二进制内容 + + Returns: + 提取的纯文本 + + Raises: + ValueError: 未识别扩展名 → "不支持的文件类型: {ext}" + 解析失败 → "文件解析失败: {detail}"(保留原异常链) + """ + ext = Path(filename).suffix.lower() + parser = _PARSERS.get(ext) + if parser is None: + raise ValueError(f"不支持的文件类型: {ext or '(无扩展名)'}") + try: + return parser(content) + except ValueError: + raise + except Exception as exc: + raise ValueError(f"文件解析失败: {exc}") from exc + + +def supported_extensions() -> set[str]: + """返回当前支持的扩展名集合(含点号,全小写)""" + return set(_PARSERS.keys()) diff --git a/app/core/ingest_tasks.py b/app/core/ingest_tasks.py index 73a837a..efad77b 100644 --- a/app/core/ingest_tasks.py +++ b/app/core/ingest_tasks.py @@ -9,6 +9,7 @@ Redis 不可用或写入失败仅记录 warning,不影响任务执行。 """ import asyncio +import hashlib import uuid from datetime import UTC, datetime from enum import StrEnum @@ -25,6 +26,8 @@ logger = structlog.get_logger() # Redis 任务状态 key 前缀 REDIS_KEY_PREFIX = "ingest_task:" +# Redis 文本去重 key 前缀(value 为已入库文档的 IngestionResult JSON) +DEDUP_KEY_PREFIX = "dedup:sha256:" class IngestTaskStatus(StrEnum): @@ -62,9 +65,39 @@ class IngestTaskManager: self._mirror_tasks: set[asyncio.Task[None]] = set() async def submit(self, doc: DocumentInput) -> str: - """登记入库任务并后台执行,立即返回 task_id""" + """登记入库任务并后台执行,立即返回 task_id + + 文本去重:基于 doc.text 的 sha256 在 Redis 中查重;命中则直接复用旧 + IngestionResult(仅置 deduplicated=True),不重跑流水线;未命中走原 + 异步入库流程,完成后写入去重记录供后续命中复用。Redis 不可用时跳过 + 去重,按原流程执行,不影响主流程。 + """ task_id = uuid.uuid4().hex now = _utc_now_iso() + text_hash = hashlib.sha256(doc.text.encode("utf-8")).hexdigest() + + # 1. 去重命中:直接置 done,复用旧结果,不调 _run + dedup_record = await self._lookup_dedup(text_hash) + if dedup_record is not None: + result_dict = dict(dedup_record) + result_dict["deduplicated"] = True + self._tasks[task_id] = { + "task_id": task_id, + "status": IngestTaskStatus.DONE, + "created_at": now, + "updated_at": now, + "result": result_dict, + "error": None, + } + self._schedule_mirror(task_id) + logger.info( + "入库任务命中去重,复用既有文档", + task_id=task_id, + document_id=result_dict.get("document_id"), + ) + return task_id + + # 2. 未命中:登记 pending 并后台跑流水线 self._tasks[task_id] = { "task_id": task_id, "status": IngestTaskStatus.PENDING, @@ -74,7 +107,7 @@ class IngestTaskManager: "error": None, } self._schedule_mirror(task_id) - background = asyncio.create_task(self._run(task_id, doc)) + background = asyncio.create_task(self._run(task_id, doc, text_hash)) self._background_tasks.add(background) background.add_done_callback(self._background_tasks.discard) logger.info("入库任务已登记", task_id=task_id, title=doc.title) @@ -110,8 +143,8 @@ class IngestTaskManager: raise TimeoutError(f"入库任务 {task_id} 在 {timeout}s 内未进入终态") await asyncio.sleep(0.01) - async def _run(self, task_id: str, doc: DocumentInput) -> None: - """后台执行入库:并发限流 + 阶段状态推进 + 结果/错误落账""" + async def _run(self, task_id: str, doc: DocumentInput, text_hash: str) -> None: + """后台执行入库:并发限流 + 阶段状态推进 + 结果/错误落账 + 去重记录写入""" async with self._semaphore: try: result = await self._ingester.ingest(doc, progress_cb=lambda stage: self._on_progress(task_id, stage)) @@ -128,14 +161,39 @@ class IngestTaskManager: except Exception as exc: self._finish_failed(task_id, {"stage": "unknown", "message": str(exc), "partial_summary": None}) else: + result_dict = result.model_dump(mode="json") self._tasks[task_id].update( status=IngestTaskStatus.DONE, updated_at=_utc_now_iso(), - result=result.model_dump(mode="json"), + result=result_dict, ) self._schedule_mirror(task_id) + await self._record_dedup(text_hash, result_dict) logger.info("入库任务完成", task_id=task_id) + async def _lookup_dedup(self, text_hash: str) -> dict[str, Any] | None: + """查询文本去重记录;Redis 不可用或异常时降级为未命中""" + if self._redis is None: + return None + try: + return await self._redis.get_json(f"{DEDUP_KEY_PREFIX}{text_hash}") + except Exception: + logger.warning("去重记录查询失败,降级为未命中", text_hash=text_hash, exc_info=True) + return None + + async def _record_dedup(self, text_hash: str, result_dict: dict[str, Any]) -> None: + """写入文本去重记录(含完整 IngestionResult),供后续命中复用;失败仅告警""" + if self._redis is None: + return + try: + await self._redis.set_json( + f"{DEDUP_KEY_PREFIX}{text_hash}", + result_dict, + ttl=self._settings.ingest_task_ttl_done, + ) + except Exception: + logger.warning("去重记录写入失败", text_hash=text_hash, exc_info=True) + def _finish_failed(self, task_id: str, error: dict[str, Any]) -> None: """将任务置为 failed 并记录错误信息""" self._tasks[task_id].update(status=IngestTaskStatus.FAILED, updated_at=_utc_now_iso(), error=error) diff --git a/app/core/query_parser.py b/app/core/query_parser.py index 5647ae1..29c0ed3 100644 --- a/app/core/query_parser.py +++ b/app/core/query_parser.py @@ -40,6 +40,9 @@ class ParsedQuery(BaseModel): rewrite: str = Field(description="rewrite 后的 query") keywords: list[str] = Field(default_factory=list, description="提取的关键词") categories: list[CategoryHit] = Field(default_factory=list, description="命中类目列表") + entities: list[str] = Field(default_factory=list, description="提取的实体(人名/产品/技术等)") + intent: str = Field(default="", description="查询意图分类(如 查询/对比/操作/定义/排查)") + time_range: str = Field(default="", description="时间范围(如 '2023年'、'最近一个月'),空表示无") parse_failed: bool = Field(default=False, description="LLM 输出解析是否失败") @@ -136,11 +139,28 @@ class QueryParser: logger.warning("query 解析失败:JSON 必填字段缺失或类型错误", query=query, output=raw[:200]) return ParsedQuery(raw_query=query, rewrite=query, parse_failed=True) + # 扩展字段:实体/意图/时间范围(缺失或类型错误时降级为默认值,不触发 parse_failed) + entities = data.get("entities", []) + if not isinstance(entities, list): + logger.warning("query 解析 entities 非列表,置空", query=query) + entities = [] + intent = data.get("intent", "") + if not isinstance(intent, str): + logger.warning("query 解析 intent 非字符串,置空", query=query) + intent = "" + time_range = data.get("time_range", "") + if not isinstance(time_range, str): + logger.warning("query 解析 time_range 非字符串,置空", query=query) + time_range = "" + return ParsedQuery( raw_query=query, rewrite=rewrite, keywords=[str(k) for k in keywords], categories=self._validate_categories(categories, query), + entities=[str(e) for e in entities], + intent=intent, + time_range=time_range, ) async def parse_and_route(self, query: str) -> RouteDecision: @@ -178,17 +198,22 @@ class QueryParser: category_lines = [f"- {c.name}: {c.description}" for c in self.taxonomy if c.name != UNCATEGORIZED] category_block = "\n".join(category_lines) return ( - "你是搜索查询分析助手。请分析用户 query,完成三件事:\n" + "你是搜索查询分析助手。请分析用户 query,完成以下事项:\n" "1. 判断 query 意图命中以下哪些知识类目,并给出每个类目的置信度(0~1 之间的小数);\n" "2. 将 query 改写为更适合检索的形式;\n" - "3. 提取 query 的关键词。\n\n" + "3. 提取 query 的关键词;\n" + "4. 提取 query 中的实体(人名、产品名、技术名词、组织机构等);\n" + "5. 判断 query 的查询意图,从 [查询, 对比, 操作, 定义, 排查] 中选最接近的一个;\n" + "6. 提取 query 中的时间范围(如 '2023年'、'上个月'、'Q1'),无则留空字符串。\n\n" f"可选类目:\n{category_block}\n\n" "要求:\n" "- categories 中的 name 只能从上面的类目名中选择,不要输出其他名称;\n" "- 若没有明显命中的类目,categories 返回空列表;\n" + "- entities 没有则返回空数组;\n" "- 只输出 JSON,不要输出任何其他内容。\n\n" '输出格式:{"categories": [{"name": "类目名", "confidence": 0.0}], ' - '"rewrite": "改写后的 query", "keywords": ["关键词"]}\n\n' + '"rewrite": "改写后的 query", "keywords": ["关键词"], ' + '"entities": ["实体"], "intent": "查询", "time_range": ""}\n\n' f"用户 query:{query}" ) diff --git a/app/core/result_summarizer.py b/app/core/result_summarizer.py new file mode 100644 index 0000000..eedef9b --- /dev/null +++ b/app/core/result_summarizer.py @@ -0,0 +1,66 @@ +"""检索结果 AI 总结 + +对检索返回的 chunk 命中结果,调用 Ollama 本地模型生成一段针对用户 query 的总结回答。 +仅基于检索结果内容,不编造未提及的信息。 +""" + +import structlog + +from app.config import settings +from app.models.search import SearchHit +from app.services.ollama import OllamaClient + +logger = structlog.get_logger() + + +class ResultSummarizer: + """检索结果总结器""" + + def __init__(self, ollama: OllamaClient | None = None) -> None: + self.ollama = ollama or OllamaClient() + + async def summarize(self, query: str, hits: list[SearchHit]) -> str: + """对检索结果生成针对 query 的总结 + + 取前 settings.result_summary_max_hits 条命中拼接为上下文, + 无命中时返回空字符串(不调 LLM)。 + """ + if not hits: + return "" + max_hits = settings.result_summary_max_hits + selected = hits[:max_hits] + context = self._build_context(selected) + prompt = self._build_prompt(query, context) + try: + summary = await self.ollama.generate(prompt) + except Exception: + logger.warning("检索结果总结生成失败,返回空字符串", exc_info=True) + return "" + logger.info("检索结果总结完成", query=query, hits_count=len(selected), summary_len=len(summary)) + return summary.strip() + + @staticmethod + def _build_context(hits: list[SearchHit]) -> str: + """拼接命中结果为带编号的上下文""" + blocks: list[str] = [] + for i, hit in enumerate(hits, start=1): + header_parts = [f"[{i}]"] + if hit.title: + header_parts.append(hit.title) + if hit.section_path: + header_parts.append(hit.section_path) + header = " / ".join(header_parts) + blocks.append(f"{header}\n{hit.text}") + return "\n\n---\n\n".join(blocks) + + @staticmethod + def _build_prompt(query: str, context: str) -> str: + return ( + "请根据以下检索结果,针对用户问题生成一段简洁的总结回答。\n" + "要求:\n" + "- 综合多条结果信息,不要简单逐条罗列;\n" + "- 只基于检索结果内容,不编造未提及的信息;\n" + "- 用中文回答,简洁明了。\n\n" + f"用户问题:{query}\n\n" + f"检索结果:\n{context}" + ) diff --git a/app/core/retriever.py b/app/core/retriever.py index 58d7e00..9240f68 100644 --- a/app/core/retriever.py +++ b/app/core/retriever.py @@ -15,11 +15,12 @@ from qdrant_client import models from app.config import settings from app.core.embeddings import EmbeddingService, create_embedding_service -from app.core.query_parser import QueryParser +from app.core.query_parser import QueryParser, RouteDecision from app.core.ranker import finalize, rrf_fuse +from app.core.result_summarizer import ResultSummarizer from app.core.sparse import SparseEncoder from app.models.knowledge import load_taxonomy -from app.models.search import SearchHit, SearchRequest, SearchResponse +from app.models.search import ExtractedInfo, SearchHit, SearchRequest, SearchResponse from app.services.ollama import OllamaClient from app.services.qdrant import ( COLLECTION_CHUNKS, @@ -47,6 +48,7 @@ class Retriever: query_parser: QueryParser | None = None, embedding: EmbeddingService | None = None, sparse_encoder: SparseEncoder | None = None, + result_summarizer: ResultSummarizer | None = None, ) -> None: self.qdrant = qdrant or QdrantService() self.query_parser = query_parser or QueryParser( @@ -55,6 +57,7 @@ class Retriever: ) self.embedding = embedding or create_embedding_service() self.sparse_encoder = sparse_encoder or SparseEncoder() + self.result_summarizer = result_summarizer or ResultSummarizer() async def search(self, request: SearchRequest) -> SearchResponse: """分层检索主流程""" @@ -74,12 +77,8 @@ class Retriever: # L1 无候选文档 → 全库 chunk 兜底 chunk_hits = await self._search_collection(COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, None) logger.info("L1 无命中,全库 chunk 兜底", hits=len(chunk_hits)) - return SearchResponse( - query=request.query, - hits=self._to_hits(finalize(chunk_hits, self._final_k(request))), - routed_categories=route.filter_categories or [], - fallback=True, - ) + hits = self._to_hits(finalize(chunk_hits, self._final_k(request))) + return await self._build_response(request, route, hits, fallback=True) doc_ids = _unique((p.payload or {}).get("doc_id") for p in l1_hits) @@ -125,12 +124,8 @@ class Retriever: logger.info("chunk 检索完成", hits=len(chunk_hits)) final_points = finalize(rrf_fuse([chunk_hits]), self._final_k(request)) - return SearchResponse( - query=request.query, - hits=self._to_hits(final_points), - routed_categories=route.filter_categories or [], - fallback=route.fallback, - ) + hits = self._to_hits(final_points) + return await self._build_response(request, route, hits, fallback=route.fallback) async def _search_collection( self, @@ -167,3 +162,33 @@ class Retriever: ) ) return hits + + async def _build_response( + self, + request: SearchRequest, + route: RouteDecision, + hits: list[SearchHit], + *, + fallback: bool, + ) -> SearchResponse: + """构造最终响应:组装 AI 提取信息,按需生成结果总结""" + parsed = route.parsed + extracted = ExtractedInfo( + rewrite=parsed.rewrite, + keywords=parsed.keywords, + entities=parsed.entities, + intent=parsed.intent, + time_range=parsed.time_range, + categories=[c.name for c in parsed.categories], + ) + summary: str | None = None + if request.summarize: + summary = await self.result_summarizer.summarize(request.query, hits) + return SearchResponse( + query=request.query, + hits=hits, + routed_categories=route.filter_categories or [], + fallback=fallback, + extracted_info=extracted, + summary=summary, + ) diff --git a/app/main.py b/app/main.py index 7ba8dc0..28e045d 100644 --- a/app/main.py +++ b/app/main.py @@ -8,10 +8,12 @@ from fastapi.exceptions import RequestValidationError from fastapi.responses import FileResponse, JSONResponse from app.api.response import ApiError, error +from app.api.v1.auth import router as auth_router from app.api.v1.document import router as document_router from app.api.v1.knowledge import router as knowledge_router from app.api.v1.search import router as search_router from app.config import settings +from app.core.auth import ensure_default_admin from app.services.qdrant import QdrantService logger = structlog.get_logger() @@ -19,15 +21,24 @@ logger = structlog.get_logger() @asynccontextmanager async def lifespan(_: FastAPI) -> AsyncIterator[None]: - """应用生命周期:启动时初始化 Qdrant 集合 + """应用生命周期:启动时初始化 Qdrant 集合与默认管理员 - 初始化失败仅记录日志、不阻止启动(本地开发可能无 Qdrant)。 + 初始化失败仅记录日志、不阻止启动(本地开发可能无 Qdrant/Redis)。 """ try: await QdrantService().ensure_collections() logger.info("Qdrant 集合初始化完成") except Exception: logger.error("Qdrant 集合初始化失败,跳过初始化继续启动") + try: + await ensure_default_admin() + except Exception: + logger.warning("默认管理员初始化失败,跳过", exc_info=True) + try: + Path(settings.upload_dir).mkdir(parents=True, exist_ok=True) + logger.info("上传目录已就绪", upload_dir=settings.upload_dir) + except Exception: + logger.warning("上传目录初始化失败,文件落盘将按需创建", exc_info=True) yield @@ -38,6 +49,7 @@ app = FastAPI( lifespan=lifespan, ) +app.include_router(auth_router) app.include_router(search_router) app.include_router(document_router) app.include_router(knowledge_router) diff --git a/app/models/auth.py b/app/models/auth.py new file mode 100644 index 0000000..20ab85c --- /dev/null +++ b/app/models/auth.py @@ -0,0 +1,42 @@ +"""认证相关数据模型""" + +from datetime import datetime + +from pydantic import BaseModel, Field + + +class LoginRequest(BaseModel): + """登录请求""" + + username: str = Field(description="用户名") + password: str = Field(description="明文密码") + + +class RegisterRequest(BaseModel): + """注册请求""" + + username: str = Field(min_length=3, max_length=32, description="用户名(3-32 字符)") + password: str = Field(min_length=6, max_length=128, description="密码(6-128 字符)") + + +class AuthUser(BaseModel): + """对外暴露的用户信息(不含密码)""" + + username: str = Field(description="用户名") + role: str = Field(default="user", description="角色:admin | user") + created_at: datetime = Field(description="创建时间") + + +class StoredUser(AuthUser): + """存储层用户:含密码哈希,仅内部使用,不对外暴露""" + + hashed_password: str = Field(description="bcrypt 密码哈希") + + +class TokenResponse(BaseModel): + """登录成功返回的 token 信息""" + + access_token: str = Field(description="JWT access token") + token_type: str = Field(default="bearer", description="token 类型") + expires_in: int = Field(description="token 有效期(秒)") + user: AuthUser = Field(description="登录用户信息") diff --git a/app/models/document.py b/app/models/document.py index 4d18cb6..f117fd0 100644 --- a/app/models/document.py +++ b/app/models/document.py @@ -49,3 +49,4 @@ class IngestionResult(BaseModel): chunks_count: int = Field(default=0, description="写入的 chunk 数量") tags: list[str] = Field(default_factory=list, description="附加分类标签") category_confidence: float = Field(default=0.0, description="主类目分类置信度") + deduplicated: bool = Field(default=False, description="是否命中去重复用旧文档(True 时 document_id 为既有文档 ID)") diff --git a/app/models/search.py b/app/models/search.py index f132de6..d209be2 100644 --- a/app/models/search.py +++ b/app/models/search.py @@ -8,6 +8,7 @@ class SearchRequest(BaseModel): query: str = Field(description="查询文本") top_k: int | None = Field(default=None, description="返回结果数,为空时使用 settings.retrieval_final_k") + summarize: bool = Field(default=False, description="是否对检索结果生成 AI 总结") class SearchHit(BaseModel): @@ -21,6 +22,17 @@ class SearchHit(BaseModel): doc_summary: str = Field(default="", description="L1 文档总结,仅用于上下文标注") +class ExtractedInfo(BaseModel): + """AI 从 query 中提取的关键信息""" + + rewrite: str = Field(default="", description="改写后的 query") + keywords: list[str] = Field(default_factory=list, description="关键词") + entities: list[str] = Field(default_factory=list, description="实体(人名/产品/技术等)") + intent: str = Field(default="", description="查询意图分类") + time_range: str = Field(default="", description="时间范围,空表示无") + categories: list[str] = Field(default_factory=list, description="命中类目名(含未达阈值的候选)") + + class SearchResponse(BaseModel): """检索响应""" @@ -28,3 +40,5 @@ class SearchResponse(BaseModel): hits: list[SearchHit] = Field(default_factory=list, description="命中结果列表") routed_categories: list[str] = Field(default_factory=list, description="query 路由命中的类目") fallback: bool = Field(default=False, description="是否走了全库兜底路径") + extracted_info: ExtractedInfo | None = Field(default=None, description="AI 提取的 query 关键信息") + summary: str | None = Field(default=None, description="检索结果 AI 总结,仅 summarize=true 时返回") diff --git a/app/static/admin.html b/app/static/admin.html index d41d2d2..295e6b1 100644 --- a/app/static/admin.html +++ b/app/static/admin.html @@ -115,11 +115,56 @@ .status-running { background: #eff6ff; color: #1d4ed8; border-color: #93c5fd; } .status-done { background: #f0fdf4; color: #15803d; border-color: #86efac; } .status-failed { background: #fef2f2; color: #b91c1c; border-color: #fca5a5; } + header { position: relative; } + .user-area { position: absolute; top: 14px; right: 24px; display: flex; align-items: center; gap: 8px; } + .user-area .action { padding: 4px 10px; } + .login-overlay { + position: fixed; inset: 0; background: rgba(0,0,0,0.45); + display: flex; align-items: center; justify-content: center; z-index: 100; + } + .login-box { + background: #fff; border-radius: 8px; padding: 24px 28px; width: 320px; + box-shadow: 0 6px 20px rgba(0,0,0,0.25); + } + .login-box h2 { font-size: 16px; margin: 0 0 16px; } + .extracted-info { + background: #f0f9ff; border: 1px solid #bae6fd; border-radius: 6px; + padding: 10px 12px; margin-bottom: 12px; font-size: 13px; + } + .extracted-info .ei-title { font-weight: 600; margin-bottom: 6px; color: #0369a1; } + .extracted-info .ei-row { margin-bottom: 3px; } + .extracted-info .ei-k { color: #6b7280; margin-right: 4px; } + .summary-box { + background: #f0fdf4; border: 1px solid #bbf7d0; border-radius: 6px; + padding: 12px 14px; margin-bottom: 12px; font-size: 13px; white-space: pre-wrap; + } + .summary-box .sum-title { font-weight: 600; margin-bottom: 6px; color: #15803d; }
+