feat: 更新 Docker 配置和添加 MinerU OCR 服务

- 修改了 docker-compose.yml,更新了 API 和 Web 服务的镜像名称,并添加了 MinerU OCR 服务。
- 在 pyproject.toml 中添加了 Ruff 代码检查工具的配置。
- 更新了 API 和 Web Dockerfile,添加了代理环境变量。
- 新增了用于拉取 Docker 镜像的脚本。
- 在文档中添加了关于如何使用 MinerU OCR 的说明。
- 优化了代码结构和日志记录,提升了可读性和维护性。
This commit is contained in:
Wenjie Zhang 2025-05-23 15:30:14 +08:00
parent 1f3e24cc37
commit b0db3c6b0d
25 changed files with 824 additions and 206 deletions

View File

@ -3,7 +3,7 @@ services:
build: build:
context: . context: .
dockerfile: docker/api.Dockerfile dockerfile: docker/api.Dockerfile
image: api:0.1.0 image: yuxi-api:0.1.0
container_name: api-dev container_name: api-dev
working_dir: /app working_dir: /app
volumes: volumes:
@ -16,7 +16,7 @@ services:
reservations: reservations:
devices: devices:
- driver: nvidia - driver: nvidia
count: 1 device_ids: ['2']
capabilities: [gpu] capabilities: [gpu]
ports: ports:
- "5050:5050" - "5050:5050"
@ -29,6 +29,7 @@ services:
- NEO4J_USERNAME=${NEO4J_USERNAME:-neo4j} - NEO4J_USERNAME=${NEO4J_USERNAME:-neo4j}
- NEO4J_PASSWORD=${NEO4J_PASSWORD:-0123456789} - NEO4J_PASSWORD=${NEO4J_PASSWORD:-0123456789}
- MILVUS_URI=http://milvus:19530 - MILVUS_URI=http://milvus:19530
- MINERU_OCR_URI=http://mineru-api:5051
- MODEL_DIR=/models - MODEL_DIR=/models
- RUNNING_IN_DOCKER=true - RUNNING_IN_DOCKER=true
command: uv run uvicorn server.main:app --host 0.0.0.0 --port 5050 --reload command: uv run uvicorn server.main:app --host 0.0.0.0 --port 5050 --reload
@ -39,7 +40,7 @@ services:
context: . context: .
dockerfile: docker/web.Dockerfile dockerfile: docker/web.Dockerfile
target: development target: development
image: web:0.1.0 image: yuxi-web:0.1.0
container_name: web-dev container_name: web-dev
volumes: volumes:
- ./web:/app - ./web:/app
@ -58,7 +59,7 @@ services:
graph: graph:
image: neo4j:5.26 image: neo4j:5.26
container_name: graph-dev container_name: graph
ports: ports:
- "7474:7474" - "7474:7474"
- "7687:7687" - "7687:7687"
@ -96,7 +97,7 @@ services:
restart: unless-stopped restart: unless-stopped
minio: minio:
container_name: milvus-minio-dev container_name: milvus-minio
image: minio/minio:RELEASE.2023-03-20T20-16-18Z image: minio/minio:RELEASE.2023-03-20T20-16-18Z
environment: environment:
MINIO_ACCESS_KEY: ${MINIO_ACCESS_KEY:-minioadmin} MINIO_ACCESS_KEY: ${MINIO_ACCESS_KEY:-minioadmin}
@ -116,7 +117,7 @@ services:
milvus: milvus:
image: milvusdb/milvus:v2.5.6 image: milvusdb/milvus:v2.5.6
container_name: milvus-standalone-dev container_name: milvus
command: ["milvus", "run", "standalone"] command: ["milvus", "run", "standalone"]
security_opt: security_opt:
- seccomp:unconfined - seccomp:unconfined
@ -143,6 +144,26 @@ services:
- app-network - app-network
restart: unless-stopped restart: unless-stopped
mineru-api:
build:
context: scripts/mineru-api
dockerfile: Dockerfile
image: mineru-api:latest
container_name: mineru-api
deploy:
resources:
reservations:
devices:
- driver: nvidia
device_ids: ['2']
capabilities: [gpu]
ports:
- "5051:5051"
networks:
- app-network
restart: unless-stopped
command: python -m uvicorn app:app --host 0.0.0.0 --port 5051 --app-dir /app
networks: networks:
app-network: app-network:
driver: bridge driver: bridge

View File

@ -5,9 +5,13 @@ COPY --from=ghcr.io/astral-sh/uv:0.7.2 /uv /uvx /bin/
# 设置工作目录 # 设置工作目录
WORKDIR /app WORKDIR /app
# 设置时区为 UTC+8 # 环境变量设置
ENV TZ=Asia/Shanghai ARG http_proxy
ENV UV_LINK_MODE=copy ARG https_proxy
ENV http_proxy=$http_proxy \
https_proxy=$https_proxy \
TZ=Asia/Shanghai \
UV_LINK_MODE=copy
RUN ln -snf /usr/share/zoneinfo/$TZ /etc/localtime && echo $TZ > /etc/timezone RUN ln -snf /usr/share/zoneinfo/$TZ /etc/localtime && echo $TZ > /etc/timezone

38
docker/pull_image.sh Normal file
View File

@ -0,0 +1,38 @@
#!/bin/bash
if [ $# -ne 1 ]; then
echo "Usage: $0 <image:tag>"
exit 1
fi
set -e # 当命令失败时,立即退出脚本
IMAGE_TAG=$1
# Count the number of slashes to determine the image format
SLASH_COUNT=$(echo $IMAGE_TAG | tr -cd '/' | wc -c)
# Set mirror URL based on image format
if [ $SLASH_COUNT -eq 0 ]; then
# No prefix (e.g., python:3.12-slim)
MIRROR_URL="m.daocloud.io/docker.io/library"
elif [ $SLASH_COUNT -eq 1 ]; then
# One prefix (e.g., milvusdb/milvus:latest)
MIRROR_URL="m.daocloud.io/docker.io"
else
# Two or more prefixes (e.g., quay.io/coreos/etcd:v3.5.5)
MIRROR_URL="m.daocloud.io"
fi
# Pull image from mirror
docker pull $MIRROR_URL/$IMAGE_TAG
# Tag image with original name
docker tag $MIRROR_URL/$IMAGE_TAG $IMAGE_TAG
# Remove mirror image
docker rmi $MIRROR_URL/$IMAGE_TAG
docker images
echo "Process completed successfully!"

View File

@ -2,6 +2,13 @@
FROM node:latest AS development FROM node:latest AS development
WORKDIR /app WORKDIR /app
ARG http_proxy
ARG https_proxy
ENV http_proxy=$http_proxy \
https_proxy=$https_proxy \
TZ=Asia/Shanghai
# 复制 package.json 和 package-lock.json如果存在 # 复制 package.json 和 package-lock.json如果存在
COPY ./web/package*.json ./ COPY ./web/package*.json ./

View File

@ -1,3 +1,7 @@
### 如何优雅的拉取镜像?
使用 `bash docker/pull_image.sh python:3.12` 就可以。
### 如何配置本地大语言模型? ### 如何配置本地大语言模型?
支持添加以 OpenAI 兼容模式运行的本地模型,可在 Web 设置中直接添加(适用于 vllm 和 Ollama 等)。 支持添加以 OpenAI 兼容模式运行的本地模型,可在 Web 设置中直接添加(适用于 vllm 和 Ollama 等)。
@ -69,3 +73,41 @@ docker compose up --build -d
name: nomic-embed-text name: nomic-embed-text
dimension: 768 dimension: 768
``` ```
### 如何配置 MinerU 抽取数据
在 PDF 数据处理中,可以选择配置 [MinerU](https://github.com/opendatalab/MinerU) 来实现更快速、更准确的 PDF 识别效果。
```yml
mineru-api:
build:
context: scripts/mineru-api
dockerfile: Dockerfile
image: mineru-api:latest
container_name: mineru-api
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: 1
capabilities: [gpu]
ports:
- "5051:5051"
networks:
- app-network
restart: unless-stopped
command: python -m uvicorn app:app --host 0.0.0.0 --port 5051 --app-dir /app
```
如果要添加代理的话,可以添加 build args
```yml
mineru-api:
build:
context: scripts/mineru-api
dockerfile: Dockerfile
args:
http_proxy: http://宿主机IP:7890
https_proxy: http://宿主机IP:7890
```

View File

@ -39,3 +39,12 @@ dependencies = [
"uvicorn[standard]>=0.34.2", "uvicorn[standard]>=0.34.2",
"zhipuai>=2.1.5.20250421", "zhipuai>=2.1.5.20250421",
] ]
[tool.ruff]
line-length = 140 # 代码最大行宽
select = [ # 选择的规则
"F",
"E",
"W",
"UP",
]
ignore = ["F401"] # 忽略的规则

View File

@ -0,0 +1,63 @@
# Use the official Ubuntu base image
FROM ubuntu:22.04
# Set environment variables to non-interactive to avoid prompts during installation
# 环境变量设置
ARG http_proxy
ARG https_proxy
ENV http_proxy=$http_proxy \
https_proxy=$https_proxy \
DEBIAN_FRONTEND=noninteractive
# Update the package list and install necessary packages
RUN apt-get update && \
apt-get install -y \
software-properties-common && \
add-apt-repository ppa:deadsnakes/ppa && \
apt-get update && \
apt-get install -y \
python3.10 \
python3.10-venv \
python3.10-distutils \
python3-pip \
wget \
git \
libgl1 \
libreoffice \
fonts-noto-cjk \
fonts-wqy-zenhei \
fonts-wqy-microhei \
ttf-mscorefonts-installer \
fontconfig \
libglib2.0-0 \
libxrender1 \
libsm6 \
libxext6 \
poppler-utils \
&& rm -rf /var/lib/apt/lists/*
# Set Python 3.10 as the default python3
RUN update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.10 1
# Create a virtual environment for MinerU
RUN python3 -m venv /opt/mineru_venv
# Copy the configuration file template and install magic-pdf latest
RUN /bin/bash -c "source /opt/mineru_venv/bin/activate && \
pip3 install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple && \
pip3 install -U magic-pdf[full] -i https://mirrors.aliyun.com/pypi/simple"
# Download models and update the configuration file
COPY magic-pdf.json /root/magic-pdf.json
RUN /bin/bash -c "pip3 install modelscope -i https://mirrors.aliyun.com/pypi/simple && \
wget https://gcore.jsdelivr.net/gh/opendatalab/MinerU@master/scripts/download_models.py -O download_models.py && \
python3 download_models.py && \
sed -i 's|cpu|cuda|g' /root/magic-pdf.json"
COPY app.py /app/app.py
RUN /bin/bash -c "source /opt/mineru_venv/bin/activate && \
pip3 install fastapi uvicorn python-multipart loguru -i https://mirrors.aliyun.com/pypi/simple"
# Set the entry point to activate the virtual environment and run the command line tool
ENTRYPOINT ["/bin/bash", "-c", "source /opt/mineru_venv/bin/activate && exec \"$@\"", "--"]

303
scripts/mineru-api/app.py Normal file
View File

@ -0,0 +1,303 @@
import json
import os
import tempfile
from base64 import b64encode
from glob import glob
from io import StringIO
from typing import Tuple, Union
import magic_pdf.model as model_config
import uvicorn
from fastapi import FastAPI, Form, UploadFile
from fastapi.responses import JSONResponse
from loguru import logger
from magic_pdf.config.enums import SupportedPdfParseMethod
from magic_pdf.data.data_reader_writer import DataWriter, FileBasedDataWriter
from magic_pdf.data.data_reader_writer.s3 import S3DataReader, S3DataWriter
from magic_pdf.data.dataset import ImageDataset, PymuDocDataset
from magic_pdf.data.read_api import read_local_images, read_local_office
from magic_pdf.libs.config_reader import get_bucket_name, get_s3_config
from magic_pdf.model.doc_analyze_by_custom_model import doc_analyze
from magic_pdf.operators.models import InferenceResult
from magic_pdf.operators.pipes import PipeResult
model_config.__use_inside_model__ = True
app = FastAPI()
pdf_extensions = [".pdf"]
office_extensions = [".ppt", ".pptx", ".doc", ".docx"]
image_extensions = [".png", ".jpg", ".jpeg"]
class MemoryDataWriter(DataWriter):
def __init__(self):
self.buffer = StringIO()
def write(self, path: str, data: bytes) -> None:
if isinstance(data, str):
self.buffer.write(data)
else:
self.buffer.write(data.decode("utf-8"))
def write_string(self, path: str, data: str) -> None:
self.buffer.write(data)
def get_value(self) -> str:
return self.buffer.getvalue()
def close(self):
self.buffer.close()
def init_writers(
file_path: str = None,
file: UploadFile = None,
output_path: str = None,
output_image_path: str = None,
) -> Tuple[
Union[S3DataWriter, FileBasedDataWriter],
Union[S3DataWriter, FileBasedDataWriter],
bytes,
]:
"""
Initialize writers based on path type
Args:
file_path: file path (local path or S3 path)
file: Uploaded file object
output_path: Output directory path
output_image_path: Image output directory path
Returns:
Tuple[writer, image_writer, file_bytes]: Returns initialized writer tuple and file content
"""
file_extension:str = None
if file_path:
is_s3_path = file_path.startswith("s3://")
if is_s3_path:
bucket = get_bucket_name(file_path)
ak, sk, endpoint = get_s3_config(bucket)
writer = S3DataWriter(
output_path, bucket=bucket, ak=ak, sk=sk, endpoint_url=endpoint
)
image_writer = S3DataWriter(
output_image_path, bucket=bucket, ak=ak, sk=sk, endpoint_url=endpoint
)
# 临时创建reader读取文件内容
temp_reader = S3DataReader(
"", bucket=bucket, ak=ak, sk=sk, endpoint_url=endpoint
)
file_bytes = temp_reader.read(file_path)
file_extension = os.path.splitext(file_path)[1]
else:
writer = FileBasedDataWriter(output_path)
image_writer = FileBasedDataWriter(output_image_path)
os.makedirs(output_image_path, exist_ok=True)
with open(file_path, "rb") as f:
file_bytes = f.read()
file_extension = os.path.splitext(file_path)[1]
else:
# 处理上传的文件
file_bytes = file.file.read()
file_extension = os.path.splitext(file.filename)[1]
writer = FileBasedDataWriter(output_path)
image_writer = FileBasedDataWriter(output_image_path)
os.makedirs(output_image_path, exist_ok=True)
return writer, image_writer, file_bytes, file_extension
def process_file(
file_bytes: bytes,
file_extension: str,
parse_method: str,
image_writer: Union[S3DataWriter, FileBasedDataWriter],
) -> Tuple[InferenceResult, PipeResult]:
"""
Process PDF file content
Args:
file_bytes: Binary content of file
file_extension: file extension
parse_method: Parse method ('ocr', 'txt', 'auto')
image_writer: Image writer
Returns:
Tuple[InferenceResult, PipeResult]: Returns inference result and pipeline result
"""
ds: Union[PymuDocDataset, ImageDataset] = None
if file_extension in pdf_extensions:
ds = PymuDocDataset(file_bytes)
elif file_extension in office_extensions:
# 需要使用office解析
temp_dir = tempfile.mkdtemp()
with open(os.path.join(temp_dir, f"temp_file.{file_extension}"), "wb") as f:
f.write(file_bytes)
ds = read_local_office(temp_dir)[0]
elif file_extension in image_extensions:
# 需要使用ocr解析
temp_dir = tempfile.mkdtemp()
with open(os.path.join(temp_dir, f"temp_file.{file_extension}"), "wb") as f:
f.write(file_bytes)
ds = read_local_images(temp_dir)[0]
infer_result: InferenceResult = None
pipe_result: PipeResult = None
if parse_method == "ocr":
infer_result = ds.apply(doc_analyze, ocr=True)
pipe_result = infer_result.pipe_ocr_mode(image_writer)
elif parse_method == "txt":
infer_result = ds.apply(doc_analyze, ocr=False)
pipe_result = infer_result.pipe_txt_mode(image_writer)
else: # auto
if ds.classify() == SupportedPdfParseMethod.OCR:
infer_result = ds.apply(doc_analyze, ocr=True)
pipe_result = infer_result.pipe_ocr_mode(image_writer)
else:
infer_result = ds.apply(doc_analyze, ocr=False)
pipe_result = infer_result.pipe_txt_mode(image_writer)
return infer_result, pipe_result
def encode_image(image_path: str) -> str:
"""Encode image using base64"""
with open(image_path, "rb") as f:
return b64encode(f.read()).decode()
@app.post(
"/file_parse",
tags=["projects"],
summary="Parse files (supports local files and S3)",
)
async def file_parse(
file: UploadFile = None,
file_path: str = Form(None),
parse_method: str = Form("auto"),
is_json_md_dump: bool = Form(False),
output_dir: str = Form("output"),
return_layout: bool = Form(False),
return_info: bool = Form(False),
return_content_list: bool = Form(False),
return_images: bool = Form(False),
):
"""
Execute the process of converting PDF to JSON and MD, outputting MD and JSON files
to the specified directory.
Args:
file: The PDF file to be parsed. Must not be specified together with
`file_path`
file_path: The path to the PDF file to be parsed. Must not be specified together
with `file`
parse_method: Parsing method, can be auto, ocr, or txt. Default is auto. If
results are not satisfactory, try ocr
is_json_md_dump: Whether to write parsed data to .json and .md files. Default
to False. Different stages of data will be written to different .json files
(3 in total), md content will be saved to .md file
output_dir: Output directory for results. A folder named after the PDF file
will be created to store all results
return_layout: Whether to return parsed PDF layout. Default to False
return_info: Whether to return parsed PDF info. Default to False
return_content_list: Whether to return parsed PDF content list. Default to False
"""
try:
if (file is None and file_path is None) or (
file is not None and file_path is not None
):
return JSONResponse(
content={"error": "Must provide either file or file_path"},
status_code=400,
)
# Get PDF filename
file_name = os.path.basename(file_path if file_path else file.filename).split(
"."
)[0]
output_path = f"{output_dir}/{file_name}"
output_image_path = f"{output_path}/images"
# Initialize readers/writers and get PDF content
writer, image_writer, file_bytes, file_extension = init_writers(
file_path=file_path,
file=file,
output_path=output_path,
output_image_path=output_image_path,
)
# Process PDF
infer_result, pipe_result = process_file(file_bytes, file_extension, parse_method, image_writer)
# Use MemoryDataWriter to get results
content_list_writer = MemoryDataWriter()
md_content_writer = MemoryDataWriter()
middle_json_writer = MemoryDataWriter()
# Use PipeResult's dump method to get data
pipe_result.dump_content_list(content_list_writer, "", "images")
pipe_result.dump_md(md_content_writer, "", "images")
pipe_result.dump_middle_json(middle_json_writer, "")
# Get content
content_list = json.loads(content_list_writer.get_value())
md_content = md_content_writer.get_value()
middle_json = json.loads(middle_json_writer.get_value())
model_json = infer_result.get_infer_res()
# If results need to be saved
if is_json_md_dump:
writer.write_string(
f"{file_name}_content_list.json", content_list_writer.get_value()
)
writer.write_string(f"{file_name}.md", md_content)
writer.write_string(
f"{file_name}_middle.json", middle_json_writer.get_value()
)
writer.write_string(
f"{file_name}_model.json",
json.dumps(model_json, indent=4, ensure_ascii=False),
)
# Save visualization results
pipe_result.draw_layout(os.path.join(output_path, f"{file_name}_layout.pdf"))
pipe_result.draw_span(os.path.join(output_path, f"{file_name}_spans.pdf"))
pipe_result.draw_line_sort(
os.path.join(output_path, f"{file_name}_line_sort.pdf")
)
infer_result.draw_model(os.path.join(output_path, f"{file_name}_model.pdf"))
# Build return data
data = {}
if return_layout:
data["layout"] = model_json
if return_info:
data["info"] = middle_json
if return_content_list:
data["content_list"] = content_list
if return_images:
image_paths = glob(f"{output_image_path}/*.jpg")
data["images"] = {
os.path.basename(
image_path
): f"data:image/jpeg;base64,{encode_image(image_path)}"
for image_path in image_paths
}
data["md_content"] = md_content # md_content is always returned
# Clean up memory writers
content_list_writer.close()
md_content_writer.close()
middle_json_writer.close()
return JSONResponse(data, status_code=200)
except Exception as e:
logger.exception(e)
return JSONResponse(content={"error": str(e)}, status_code=500)
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=8888)

View File

@ -0,0 +1,44 @@
{
"bucket_info":{
"bucket-name-1":["ak", "sk", "endpoint"],
"bucket-name-2":["ak", "sk", "endpoint"]
},
"models-dir":"/opt/models",
"layoutreader-model-dir":"/opt/layoutreader",
"device-mode":"cuda",
"layout-config": {
"model": "doclayout_yolo"
},
"formula-config": {
"mfd_model": "yolo_v8_mfd",
"mfr_model": "unimernet_small",
"enable": true
},
"table-config": {
"model": "rapid_table",
"sub_model": "slanet_plus",
"enable": true,
"max_time": 400
},
"llm-aided-config": {
"formula_aided": {
"api_key": "your_api_key",
"base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"model": "qwen2.5-7b-instruct",
"enable": false
},
"text_aided": {
"api_key": "your_api_key",
"base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"model": "qwen2.5-7b-instruct",
"enable": false
},
"title_aided": {
"api_key": "your_api_key",
"base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"model": "qwen2.5-32b-instruct",
"enable": false
}
},
"config_version": "1.2.0"
}

View File

@ -1,8 +1,3 @@
from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.ext.declarative import declarative_base
Base = declarative_base() Base = declarative_base()
# 导入所有模型文件确保它们被注册到Base.metadata
from . import user_model
from . import kb_models
from . import thread_model

View File

@ -1,4 +1,4 @@
from sqlalchemy import Column, Integer, String, DateTime, JSON, Float, ForeignKey, Text from sqlalchemy import Column, Integer, String, DateTime, JSON, ForeignKey, Text
from sqlalchemy.orm import relationship from sqlalchemy.orm import relationship
from sqlalchemy.sql import func from sqlalchemy.sql import func
import time import time

View File

@ -54,7 +54,7 @@ async def query_test(query: str = Body(...), meta: dict = Body(...), current_use
@data.post("/file-to-chunk") @data.post("/file-to-chunk")
async def file_to_chunk(files: List[str] = Body(...), params: dict = Body(...), current_user: User = Depends(get_admin_user)): async def file_to_chunk(files: List[str] = Body(...), params: dict = Body(...), current_user: User = Depends(get_admin_user)):
logger.debug(f"File to chunk: {files}") logger.debug(f"File to chunk: {files} {params=}")
result = await knowledge_base.file_to_chunk(files, params=params) result = await knowledge_base.file_to_chunk(files, params=params)
return result return result

View File

@ -2,7 +2,7 @@ import hashlib
import os import os
import jwt import jwt
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Optional, Dict, Any from typing import Any
# JWT配置 # JWT配置
JWT_SECRET_KEY = os.environ.get("JWT_SECRET_KEY", "yuxi_know_secure_key") JWT_SECRET_KEY = os.environ.get("JWT_SECRET_KEY", "yuxi_know_secure_key")
@ -38,7 +38,7 @@ class AuthUtils:
return hashed == check_hash return hashed == check_hash
@staticmethod @staticmethod
def create_access_token(data: Dict[str, Any], expires_delta: Optional[timedelta] = None) -> str: def create_access_token(data: dict[str, Any], expires_delta: timedelta | None = None) -> str:
"""创建JWT访问令牌""" """创建JWT访问令牌"""
to_encode = data.copy() to_encode = data.copy()
@ -55,7 +55,7 @@ class AuthUtils:
return encoded_jwt return encoded_jwt
@staticmethod @staticmethod
def decode_token(token: str) -> Optional[Dict[str, Any]]: def decode_token(token: str) -> dict[str, Any] | None:
"""解码验证JWT令牌""" """解码验证JWT令牌"""
try: try:
payload = jwt.decode(token, JWT_SECRET_KEY, algorithms=[JWT_ALGORITHM]) payload = jwt.decode(token, JWT_SECRET_KEY, algorithms=[JWT_ALGORITHM])
@ -64,7 +64,7 @@ class AuthUtils:
return None return None
@staticmethod @staticmethod
def verify_access_token(token: str) -> Dict[str, Any]: def verify_access_token(token: str) -> dict[str, Any]:
"""验证访问令牌,如果无效则抛出异常""" """验证访问令牌,如果无效则抛出异常"""
try: try:
payload = jwt.decode(token, JWT_SECRET_KEY, algorithms=[JWT_ALGORITHM]) payload = jwt.decode(token, JWT_SECRET_KEY, algorithms=[JWT_ALGORITHM])

View File

@ -1,14 +1,14 @@
import json import json
import re import re
import os from collections.abc import Callable
from typing import Any, Callable, Optional, Type, Union, Annotated from typing import Annotated, Any
from pydantic import BaseModel, Field
from langchain_core.tools import tool, BaseTool, StructuredTool
from langchain_community.tools.tavily_search import TavilySearchResults from langchain_community.tools.tavily_search import TavilySearchResults
from langchain_core.tools import BaseTool, StructuredTool, tool
from pydantic import BaseModel, Field
from src import config, graph_base, knowledge_base
from src import graph_base, knowledge_base, config
# refs https://github.com/chatchat-space/LangGraph-Chatchat chatchat-server/chatchat/server/agent/tools_factory/tools_registry.py # refs https://github.com/chatchat-space/LangGraph-Chatchat chatchat-server/chatchat/server/agent/tools_factory/tools_registry.py
def regist_tool( def regist_tool(
@ -16,9 +16,9 @@ def regist_tool(
title: str = "", title: str = "",
description: str = "", description: str = "",
return_direct: bool = False, return_direct: bool = False,
args_schema: Optional[Type[BaseModel]] = None, args_schema: type[BaseModel] | None = None,
infer_schema: bool = True, infer_schema: bool = True,
) -> Union[Callable, BaseTool]: ) -> Callable | BaseTool:
""" """
wrapper of langchain tool decorator wrapper of langchain tool decorator
add tool to registry automatically add tool to registry automatically
@ -66,7 +66,12 @@ def regist_tool(
class KnowledgeRetrieverModel(BaseModel): class KnowledgeRetrieverModel(BaseModel):
query: str = Field(description="查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。") query: str = Field(
description=(
"查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,"
"不要直接使用用户的原始输入去查询。"
)
)

View File

@ -40,12 +40,12 @@ async def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
content = msg.content or msg.tool_calls content = msg.content or msg.tool_calls
if not content: if not content:
if stream_flag == True: if stream_flag:
print() print()
stream_flag = False stream_flag = False
continue continue
if stream_flag == False and content: if not stream_flag and content:
print(f"AI: {content}", end="", flush=True) print(f"AI: {content}", end="", flush=True)
stream_flag = True stream_flag = True
continue continue

View File

@ -127,7 +127,7 @@ class Config(SimpleConfig):
if self.model_dir: if self.model_dir:
if os.path.exists(self.model_dir): if os.path.exists(self.model_dir):
logger.info(f"MODEL_DIR {self.model_dir} 下面的文件夹: {os.listdir(self.model_dir)}") logger.debug(f"MODEL_DIR {self.model_dir} 下面的文件夹: {os.listdir(self.model_dir)}")
else: else:
logger.warning(f"提醒MODEL_DIR {self.model_dir} 不存在,如果未配置,请忽略,如果配置了,请检查是否配置正确,比如 docker-compose 文件中的映射") logger.warning(f"提醒MODEL_DIR {self.model_dir} 不存在,如果未配置,请忽略,如果配置了,请检查是否配置正确,比如 docker-compose 文件中的映射")

View File

@ -1,99 +1,140 @@
import os import os
import asyncio import asyncio
from pathlib import Path from pathlib import Path
from llama_index.core import Document from langchain.schema.document import Document
from llama_index.core.node_parser import SimpleFileNodeParser from langchain.text_splitter import RecursiveCharacterTextSplitter
from llama_index.core.node_parser import SentenceSplitter from langchain_community.document_loaders import (
from llama_index.readers.file import FlatReader, DocxReader TextLoader,
PyPDFLoader,
Docx2txtLoader,
UnstructuredMarkdownLoader,
UnstructuredHTMLLoader,
CSVLoader,
JSONLoader
)
from src.utils import hashstr, logger from src.utils import hashstr, logger
def chunk(text_or_path, params=None): def chunk_with_parser(file_path, params=None):
""" """
将文本或文件切分成固定大小的块 使用文件解析器将文件切分成固定大小的块
Args: Args:
text_or_path: 文本或文件路径 file_path: 文件路径
params: 参数 params: 参数
chunk_size: 块大小
chunk_overlap: 块重叠大小
use_parser: 是否使用文件解析器
Returns:
nodes: 节点列表
""" """
params = params or {} params = params or {}
chunk_size = int(params.get("chunk_size", 500)) chunk_size = int(params.get("chunk_size", 500))
chunk_overlap = int(params.get("chunk_overlap", 100)) chunk_overlap = int(params.get("chunk_overlap", 100))
splitter = SentenceSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
)
# 如果文件存在,并且是当前目录下的文件,则使用文件解析器 file_type = Path(file_path).suffix.lower()
if os.path.isfile(text_or_path) and os.path.exists(text_or_path) and os.path.abspath(text_or_path).startswith(os.getcwd()):
parser = SimpleFileNodeParser()
file_type = Path(text_or_path).suffix.lower()
if file_type in [".txt", ".json", ".md"]:
docs = FlatReader().load_data(Path(text_or_path))
elif file_type in [".docx"]:
docs = DocxReader().load_data(Path(text_or_path))
else:
raise ValueError(f"Unsupported file type `{file_type}`")
if params.get("use_parser"): # 选择合适的加载器
nodes = parser.get_nodes_from_documents(docs) if file_type in ['.txt']:
else: loader = TextLoader(file_path)
nodes = splitter.get_nodes_from_documents(docs)
elif file_type in ['.md']:
loader = UnstructuredMarkdownLoader(file_path)
elif file_type in ['.docx', '.doc']:
loader = Docx2txtLoader(file_path)
elif file_type in ['.html', '.htm']:
loader = UnstructuredHTMLLoader(file_path)
elif file_type in ['.json']:
loader = JSONLoader(file_path, jq_schema=".")
elif file_type in ['.csv']:
loader = CSVLoader(file_path)
else: else:
docs = [Document(id_=hashstr(text_or_path), text=text_or_path)] raise ValueError(f"不支持的文件类型: {file_type}")
nodes = splitter.get_nodes_from_documents(docs)
# 加载文档
docs = loader.load()
# 创建文本分割器
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
separators=["\n\n", "\n", ".", " ", ""],
)
# 分割文档
nodes = text_splitter.split_documents(docs)
# 添加序号信息到metadata
for i, node in enumerate(nodes):
if node.metadata is None:
node.metadata = {}
node.metadata["chunk_idx"] = i
return nodes return nodes
def chunk_text(text, params=None):
"""
将文本切分成固定大小的块
"""
params = params or {}
chunk_size = int(params.get("chunk_size", 500))
chunk_overlap = int(params.get("chunk_overlap", 100))
# 创建文本分割器
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
separators=["\n\n", "\n", ".", " ", ""]
)
def pdfreader(file_path): # 分割文档
nodes = text_splitter.split_text(text)
# 添加序号信息到metadata
nodes = [{"text": node, "metadata": {"chunk_idx": i}} for i, node in enumerate(nodes)]
return nodes
def chunk(text_or_path, params=None):
raise NotImplementedError("chunk is deprecated, use chunk_with_parser or chunk_text instead")
def pdfreader(file_path, params=None):
"""读取PDF文件并返回text文本""" """读取PDF文件并返回text文本"""
assert os.path.exists(file_path), "File not found" assert file_path.exists(), "File not found"
assert file_path.endswith(".pdf"), "File format not supported" assert file_path.suffix.lower() == ".pdf", "File format not supported"
from llama_index.readers.file import PDFReader # 使用LangChain的PDF加载器
doc = PDFReader().load_data(file=Path(file_path)) loader = PyPDFLoader(str(file_path))
docs = loader.load()
# 简单的拼接起来之后返回纯文本 # 简单的拼接起来之后返回纯文本
text = "\n\n".join([d.get_content() for d in doc]) text = "\n\n".join([d.page_content for d in docs])
return text return text
def plainreader(file_path): def plainreader(file_path):
"""读取普通文本文件并返回text文本""" """读取普通文本文件并返回text文本"""
assert os.path.exists(file_path), "File not found" assert os.path.exists(file_path), "File not found"
with open(file_path, "r") as f: # 使用LangChain的文本加载器
text = f.read() loader = TextLoader(str(file_path))
docs = loader.load()
text = "\n\n".join([d.page_content for d in docs])
return text return text
def read_text(file, params=None): def parse_pdf(file, params=None):
support_format = [".pdf", ".txt", ".md"] params = params or {}
assert os.path.exists(file), "File not found" opt_ocr = params.get("enable_ocr", "disable")
logger.info(f"Try to read file {file}")
if not os.path.isfile(file): if opt_ocr == "onnx_rapid_ocr":
logger.error(f"Directory not supported now!")
raise NotImplementedError("Directory not supported now!")
if file.endswith(".pdf"):
from src.plugins import ocr from src.plugins import ocr
return ocr.process_pdf(file) return ocr.process_pdf(file)
elif file.endswith(".txt") or file.endswith(".md"): elif opt_ocr == "mineru_ocr":
return plainreader(file) from src.plugins import ocr
return ocr.process_pdf_mineru(file)
else: else:
logger.error(f"File format not supported, only support {support_format}") return pdfreader(file, params=params)
raise Exception(f"File format not supported, only support {support_format}")
async def parse_pdf_async(file, params=None):
async def read_text_async(file): return await asyncio.to_thread(parse_pdf, file, params=params)
return await asyncio.to_thread(read_text, file)

View File

@ -4,12 +4,14 @@ import time
import traceback import traceback
import shutil import shutil
from sqlalchemy.orm import joinedload from sqlalchemy.orm import joinedload
from pathlib import Path
import asyncio
from pymilvus import MilvusClient, MilvusException from pymilvus import MilvusClient, MilvusException
from src import config from src import config
from src.utils import logger, hashstr from src.utils import logger, hashstr
from src.core.indexing import chunk, read_text_async from src.core.indexing import chunk_with_parser, chunk_text, parse_pdf_async
from server.db_manager import db_manager from server.db_manager import db_manager
from server.models.kb_models import KnowledgeDatabase, KnowledgeFile, KnowledgeNode from server.models.kb_models import KnowledgeDatabase, KnowledgeFile, KnowledgeNode
from src.utils.db_migration import migrate_knowledge_db from src.utils.db_migration import migrate_knowledge_db
@ -367,7 +369,7 @@ class KnowledgeBase:
for line in lines: for line in lines:
line.pop("vector") line.pop("vector")
lines.sort(key=lambda x: x.get("start_char_idx") or 0) lines.sort(key=lambda x: x.get("start_char_idx") or x.get("metadata", {}).get("chunk_idx", 0))
return {"lines": lines} return {"lines": lines}
def get_kb_by_id(self, db_id): def get_kb_by_id(self, db_id):
@ -387,25 +389,35 @@ class KnowledgeBase:
""" """
file_infos = {} file_infos = {}
for file in files: for file in files:
file_id = "file_" + hashstr(file + str(time.time())) file_path = Path(file)
file_id = "file_" + hashstr(str(file_path) + str(time.time()))
file_type = file.split(".")[-1].lower() file_type = file_path.suffix.lower().replace(".", "")
file_info_dict = {
if file_type == "pdf":
texts = await read_text_async(file)
nodes = chunk(texts, params=params)
else:
nodes = chunk(file, params=params)
file_infos[file_id] = {
"file_id": file_id, "file_id": file_id,
"filename": os.path.basename(file), "filename": file_path.name,
"path": file, "path": str(file_path),
"type": file_type, "type": file_type,
"status": "waiting", "status": "waiting",
"created_at": time.time(), "created_at": time.time(),
"nodes": [node.dict() for node in nodes] "nodes": []
} }
logger.debug(f"{file_info_dict=}")
try:
if file_type == "pdf":
texts = await parse_pdf_async(file_path, params=params)
nodes = chunk_text(texts, params=params)
else:
nodes = chunk_with_parser(file_path, params=params)
processed_nodes = [parse_node_data(node) for node in nodes]
file_info_dict["nodes"] = processed_nodes
except Exception as e:
logger.error(f"处理文件 {file_path} 时出错: {e}")
file_info_dict["status"] = "failed"
file_info_dict["error"] = str(e)
file_infos[file_id] = file_info_dict
return file_infos return file_infos
@ -426,56 +438,32 @@ class KnowledgeBase:
file_infos = {} file_infos = {}
# 使用UnstructuredURLLoader加载URL内容 for url in urls:
# loader = UnstructuredURLLoader(urls=urls, continue_on_failure=True)
for url_idx, url in enumerate(urls):
file_id = "url_" + hashstr(url + str(time.time())) file_id = "url_" + hashstr(url + str(time.time()))
file_info_dict = {
"file_id": file_id,
"filename": gen_filename_from_url(url),
"path": url,
"type": "url",
"status": "waiting",
"created_at": time.time(),
"nodes": []
}
try: try:
# 加载单个URL内容 # 加载单个URL内容
single_loader = UnstructuredURLLoader(urls=[url], continue_on_failure=False) single_loader = UnstructuredURLLoader(urls=[url], continue_on_failure=False)
documents = await single_loader.aload() documents = await single_loader.aload()
# 将文档内容合并 # 合并文档内容
text_content = "\n\n".join([doc.page_content for doc in documents]) text_content = "\n\n".join([doc.page_content for doc in documents])
file_info_dict["nodes"] = chunk_text(text_content, params=params)
# 对内容进行分块
nodes = chunk(text_content, params=params)
# 从URL中提取域名作为文件名
from urllib.parse import urlparse, unquote
parsed_url = urlparse(unquote(url))
domain = parsed_url.netloc
path = parsed_url.path
filename = f"{domain}{path}"
if filename.endswith('/'):
filename = filename[:-1]
if len(filename) > 100:
filename = filename[:97] + "..."
filename = filename.replace('/', '_')
file_infos[file_id] = {
"file_id": file_id,
"filename": filename,
"path": url,
"type": "url",
"status": "waiting",
"created_at": time.time(),
"nodes": [node.dict() for node in nodes]
}
except Exception as e: except Exception as e:
logger.error(f"处理URL {url} 时出错: {e}") logger.error(f"处理URL {url} 时出错: {e}")
file_infos[file_id] = { file_info_dict["status"] = "failed"
"file_id": file_id, file_info_dict["error"] = str(e)
"filename": url[:100] + "..." if len(url) > 100 else url,
"path": url, file_infos[file_id] = file_info_dict
"type": "url",
"status": "failed",
"created_at": time.time(),
"error": str(e),
"nodes": []
}
return file_infos return file_infos
@ -765,3 +753,30 @@ class KnowledgeBase:
updated_db = self.update_database_record(db_id, name, description) updated_db = self.update_database_record(db_id, name, description)
return updated_db return updated_db
def parse_node_data(node):
node_data = node.model_dump() if hasattr(node, "model_dump") else node
node_text = node_data.get("page_content", node_data.get("text", ""))
node_dict = {
"text": node_text,
"hash": hashstr(node_text, with_salt=True),
"start_char_idx": node_data.get("start_char_idx", None),
"end_char_idx": node_data.get("end_char_idx", None),
"metadata": node_data.get("metadata", {}),
}
return node_dict
def gen_filename_from_url(url):
# 从URL中提取域名作为文件名
from urllib.parse import urlparse, unquote
parsed_url = urlparse(unquote(url))
domain = parsed_url.netloc
path = parsed_url.path
filename = f"{domain}{path}"
if filename.endswith('/'):
filename = filename[:-1]
if len(filename) > 100:
filename = filename[:97] + "..."
filename = filename.replace('/', '_')
return filename

View File

@ -82,8 +82,8 @@ class LocalEmbeddingModel(BaseEmbeddingModel):
else: else:
logger.warning(f"Local model `{info['name']}` not found in `{self.model}`, using `{info['name']}`") logger.warning(f"Local model `{info['name']}` not found in `{self.model}`, using `{info['name']}`")
logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`" logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`")
f"如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;" logger.debug("如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
f"如果还是不行,建议手动下载到某个文件夹,比如 {os.getenv('MODEL_DIR', '/models')}/BAAI/bge-m3 目录下;") f"如果还是不行,建议手动下载到某个文件夹,比如 {os.getenv('MODEL_DIR', '/models')}/BAAI/bge-m3 目录下;")
self.model = HuggingFaceEmbeddings( self.model = HuggingFaceEmbeddings(

View File

@ -123,17 +123,23 @@ class OCRPlugin:
raise FileNotFoundError(f"PDF file not found: {pdf_path}") raise FileNotFoundError(f"PDF file not found: {pdf_path}")
try: try:
# 检查是否为文本PDF # 检查是否为文本PDF可能会出现错误比如每一页都有可读取的水印文字但是内容本身是扫描件需要使用OCR处理
if is_text_pdf(pdf_path): # if is_text_pdf(pdf_path):
logger.info(f"PDF file is text, use llama_index.readers.file to read") # from src.core.indexing import pdfreader
return pdfreader(pdf_path) # logger.info("PDF file is text, use llama_index.readers.file to read")
# return pdfreader(pdf_path)
# 将PDF转换为图像 images = []
filename = os.path.basename(pdf_path).split('.')[0]
output_dir = os.path.join('saves', 'data', 'pdf2txt', filename)
os.makedirs(output_dir, exist_ok=True)
images = self.convert_imgs(pdf_path, output_dir) pdfDoc = fitz.open(pdf_path)
totalPage = pdfDoc.page_count
for pg in tqdm(range(totalPage), desc='to images', ncols=100):
page = pdfDoc[pg]
rotate, zoom_x, zoom_y = 0, 2, 2
mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate)
pix = page.get_pixmap(matrix=mat, alpha=False)
img_pil = Image.frombytes("RGB", [pix.width, pix.height], pix.samples)
images.append(img_pil)
# 处理每个图像并合并文本 # 处理每个图像并合并文本
all_text = [] all_text = []
@ -147,46 +153,58 @@ class OCRPlugin:
logger.error(f"PDF processing error: {str(e)}") logger.error(f"PDF processing error: {str(e)}")
return "" return ""
def process_pdf_mineru(self, pdf_path):
"""
使用Mineru OCR处理PDF文件
:param pdf_path: PDF文件路径
:return: 提取的文本
"""
mineru_ocr_uri = os.getenv("MINERU_OCR_URI", "http://localhost:5051")
import requests
import json
def convert_imgs(self, pdf_path, output_dir): # 读取PDF文件
imgs = [] with open(pdf_path, 'rb') as f:
img_dir = os.path.join(output_dir, 'imgs') files = {'file': f}
if not os.path.exists(img_dir): data = {
os.makedirs(img_dir) 'parse_method': 'ocr', # 使用OCR模式
pdfDoc = fitz.open(pdf_path) 'is_json_md_dump': False, # 不需要保存中间文件
totalPage = pdfDoc.page_count 'return_layout': False, # 不需要返回布局信息
for pg in tqdm(range(totalPage), desc='to imgs', ncols=100): 'return_info': False, # 不需要返回额外信息
page = pdfDoc[pg] 'return_content_list': False, # 不需要返回内容列表
rotate = int(0) 'return_images': False, # 不需要返回图片
zoom_x = 2 }
zoom_y = 2
mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate)
pix = page.get_pixmap(matrix=mat, alpha=False)
img_filename = os.path.join(img_dir, f'images_{pg+1}.png')
pix.save(img_filename) # os.sep
imgs.append(img_filename)
else:
img_names = sorted(os.listdir(img_dir))
imgs = [os.path.join(img_dir, img_name) for img_name in img_names]
return imgs try:
# 发送POST请求到Mineru OCR服务
response = requests.post(
f"{mineru_ocr_uri}/file_parse",
files=files,
data=data
)
response.raise_for_status() # 检查响应状态
# 解析响应
result = response.json()
if 'md_content' in result:
return result['md_content']
else:
logger.error("Mineru OCR response does not contain md_content")
return ""
except requests.exceptions.RequestException as e:
logger.error(f"Mineru OCR request failed: {str(e)}")
return ""
except json.JSONDecodeError as e:
logger.error(f"Failed to parse Mineru OCR response: {str(e)}")
return ""
except Exception as e:
logger.error(f"Unexpected error in Mineru OCR processing: {str(e)}")
return ""
def get_state(task_id): def get_state(task_id):
return GOLBAL_STATE.get(task_id, {}) return GOLBAL_STATE.get(task_id, {})
def pdfreader(file_path):
"""读取PDF文件并返回text文本"""
assert os.path.exists(file_path), "File not found"
assert file_path.endswith(".pdf"), "File format not supported"
from llama_index.readers.file import PDFReader
doc = PDFReader().load_data(file=Path(file_path))
# 简单的拼接起来之后返回纯文本
text = "\n\n".join([d.get_content() for d in doc])
return text
def plainreader(file_path): def plainreader(file_path):
"""读取普通文本文件并返回text文本""" """读取普通文本文件并返回text文本"""
assert os.path.exists(file_path), "File not found" assert os.path.exists(file_path), "File not found"

View File

@ -1,5 +1,4 @@
import os import os
from typing import List, Dict
from tavily import TavilyClient from tavily import TavilyClient
from src.utils.logging_config import logger from src.utils.logging_config import logger
@ -11,7 +10,7 @@ class WebSearcher:
self.client = TavilyClient(api_key) self.client = TavilyClient(api_key)
logger.info("WebSearcher initialized with Tavily client") logger.info("WebSearcher initialized with Tavily client")
def search(self, query: str, max_results: int = 1) -> List[Dict]: def search(self, query: str, max_results: int = 1) -> list[dict]:
""" """
使用 Tavily 搜索相关内容 使用 Tavily 搜索相关内容
@ -45,7 +44,7 @@ class WebSearcher:
logger.error(f"Error during web search: {str(e)}") logger.error(f"Error during web search: {str(e)}")
return [] return []
def format_search_results(self, results: List[Dict]) -> str: def format_search_results(self, results: list[dict]) -> str:
""" """
将搜索结果格式化为文本 将搜索结果格式化为文本

View File

@ -34,6 +34,7 @@
--main-color: #1c6586; --main-color: #1c6586;
--main-color-dark: #004d5c; --main-color-dark: #004d5c;
--main-color-text: #104461;
--main-light-1: #0076AB; --main-light-1: #0076AB;
--main-light-2: #DAEAED; --main-light-2: #DAEAED;
--main-light-3: #EDF0F1; --main-light-3: #EDF0F1;

View File

@ -93,7 +93,7 @@
v-if="configStore.config.enable_web_search" v-if="configStore.config.enable_web_search"
@click="meta.use_web=!meta.use_web" @click="meta.use_web=!meta.use_web"
> >
<CompassOutlined style="margin-right: 3px;"/> <Compass style="margin-right: 3px;" size="14"/>
联网搜索 联网搜索
</div> </div>
<div <div
@ -101,7 +101,7 @@
v-if="configStore.config.enable_knowledge_graph" v-if="configStore.config.enable_knowledge_graph"
@click="meta.use_graph=!meta.use_graph" @click="meta.use_graph=!meta.use_graph"
> >
<DeploymentUnitOutlined style="margin-right: 3px;"/> <Waypoints style="margin-right: 3px;" size="14"/>
知识图谱 知识图谱
</div> </div>
<a-dropdown <a-dropdown
@ -109,7 +109,7 @@
:class="{'opt-item': true, 'active': meta.selectedKB !== null}" :class="{'opt-item': true, 'active': meta.selectedKB !== null}"
> >
<a class="ant-dropdown-link" @click.prevent> <a class="ant-dropdown-link" @click.prevent>
<BookOutlined style="margin-right: 3px;"/> <BookCheck style="margin-right: 3px;" size="14"/>
<span class="text">{{ meta.selectedKB === null ? '不使用知识库' : opts.databases[meta.selectedKB]?.name }}</span> <span class="text">{{ meta.selectedKB === null ? '不使用知识库' : opts.databases[meta.selectedKB]?.name }}</span>
</a> </a>
<template #overlay> <template #overlay>
@ -147,7 +147,7 @@ import {
PlusCircleOutlined, PlusCircleOutlined,
DeploymentUnitOutlined, DeploymentUnitOutlined,
} from '@ant-design/icons-vue' } from '@ant-design/icons-vue'
import { Ellipsis, PanelLeftOpen, MessageSquarePlus } from 'lucide-vue-next' import { Ellipsis, PanelLeftOpen, MessageSquarePlus, Compass, Waypoints, BookCheck } from 'lucide-vue-next'
import { onClickOutside } from '@vueuse/core' import { onClickOutside } from '@vueuse/core'
import { useConfigStore } from '@/stores/config' import { useConfigStore } from '@/stores/config'
import { useUserStore } from '@/stores/user' import { useUserStore } from '@/stores/user'
@ -972,6 +972,13 @@ const findLastIndex = (array, predicate) => {
max-width: 1200px; max-width: 1200px;
} }
.opt-item {
display: flex;
justify-content: center;
align-items: center;
gap: 4px;
}
.note { .note {
width: 100%; width: 100%;
font-size: small; font-size: small;

View File

@ -115,9 +115,9 @@
<a-input-number v-model:value="chunkParams.chunk_overlap" :min="0" :max="1000" /> <a-input-number v-model:value="chunkParams.chunk_overlap" :min="0" :max="1000" />
<p class="param-description">相邻文本片段间的重叠字符数</p> <p class="param-description">相邻文本片段间的重叠字符数</p>
</a-form-item> </a-form-item>
<a-form-item label="使用文件节点解析器" name="use_parser"> <a-form-item label="使用OCR" name="enable_ocr">
<a-switch v-model:checked="chunkParams.use_parser" /> <a-select v-model:value="chunkParams.enable_ocr" :options="enable_ocr_options" />
<p class="param-description">启用特定文件格式的智能分析</p> <p class="param-description">启用OCR功能支持PDF文件的文本提取</p>
</a-form-item> </a-form-item>
</a-form> </a-form>
</div> </div>
@ -389,6 +389,12 @@ const meta = reactive({
sortBy: 'rerank_score', sortBy: 'rerank_score',
}); });
const enable_ocr_options = ref([
{ value: 'disable', payload: { title: '不启用' } },
{ value: 'onnx_rapid_ocr', payload: { title: 'ONNX with RapidOCR' } },
{ value: 'mineru_ocr', payload: { title: 'MinerU OCR' } },
])
const use_rewrite_queryOptions = ref([ const use_rewrite_queryOptions = ref([
{ value: 'off', payload: { title: 'off', subTitle: '不启用' } }, { value: 'off', payload: { title: 'off', subTitle: '不启用' } },
{ value: 'on', payload: { title: 'on', subTitle: '启用重写' } }, { value: 'on', payload: { title: 'on', subTitle: '启用重写' } },
@ -635,7 +641,7 @@ const deleteFile = (fileId) => {
const chunkParams = ref({ const chunkParams = ref({
chunk_size: 1000, chunk_size: 1000,
chunk_overlap: 200, chunk_overlap: 200,
use_parser: false, enable_ocr: 'disable',
}) })
const chunkResults = ref([]); const chunkResults = ref([]);