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

View File

@ -5,9 +5,13 @@ COPY --from=ghcr.io/astral-sh/uv:0.7.2 /uv /uvx /bin/
# 设置工作目录
WORKDIR /app
# 设置时区为 UTC+8
ENV TZ=Asia/Shanghai
ENV UV_LINK_MODE=copy
# 环境变量设置
ARG http_proxy
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

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
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如果存在
COPY ./web/package*.json ./

View File

@ -1,3 +1,7 @@
### 如何优雅的拉取镜像?
使用 `bash docker/pull_image.sh python:3.12` 就可以。
### 如何配置本地大语言模型?
支持添加以 OpenAI 兼容模式运行的本地模型,可在 Web 设置中直接添加(适用于 vllm 和 Ollama 等)。
@ -44,7 +48,7 @@ services:
```bash
docker compose down
docker compose up --build -d
docker compose up --build -d
```
注:添加本地向量模型由于在 docker 内外的路径差异很大,因此建议参考前面的路径映射之后,也在这里添加。
@ -69,3 +73,41 @@ docker compose up --build -d
name: nomic-embed-text
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",
"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
Base = declarative_base()
# 导入所有模型文件确保它们被注册到Base.metadata
from . import user_model
from . import kb_models
from . import thread_model
Base = declarative_base()

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.sql import func
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")
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)
return result

View File

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

View File

@ -1,14 +1,14 @@
import json
import re
import os
from typing import Any, Callable, Optional, Type, Union, Annotated
from collections.abc import Callable
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_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
def regist_tool(
@ -16,9 +16,9 @@ def regist_tool(
title: str = "",
description: str = "",
return_direct: bool = False,
args_schema: Optional[Type[BaseModel]] = None,
args_schema: type[BaseModel] | None = None,
infer_schema: bool = True,
) -> Union[Callable, BaseTool]:
) -> Callable | BaseTool:
"""
wrapper of langchain tool decorator
add tool to registry automatically
@ -66,7 +66,12 @@ def regist_tool(
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
if not content:
if stream_flag == True:
if stream_flag:
print()
stream_flag = False
continue
if stream_flag == False and content:
if not stream_flag and content:
print(f"AI: {content}", end="", flush=True)
stream_flag = True
continue

View File

@ -127,7 +127,7 @@ class Config(SimpleConfig):
if 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:
logger.warning(f"提醒MODEL_DIR {self.model_dir} 不存在,如果未配置,请忽略,如果配置了,请检查是否配置正确,比如 docker-compose 文件中的映射")

View File

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

View File

@ -4,12 +4,14 @@ import time
import traceback
import shutil
from sqlalchemy.orm import joinedload
from pathlib import Path
import asyncio
from pymilvus import MilvusClient, MilvusException
from src import config
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.models.kb_models import KnowledgeDatabase, KnowledgeFile, KnowledgeNode
from src.utils.db_migration import migrate_knowledge_db
@ -367,7 +369,7 @@ class KnowledgeBase:
for line in lines:
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}
def get_kb_by_id(self, db_id):
@ -387,25 +389,35 @@ class KnowledgeBase:
"""
file_infos = {}
for file in files:
file_id = "file_" + hashstr(file + str(time.time()))
file_type = file.split(".")[-1].lower()
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_path = Path(file)
file_id = "file_" + hashstr(str(file_path) + str(time.time()))
file_type = file_path.suffix.lower().replace(".", "")
file_info_dict = {
"file_id": file_id,
"filename": os.path.basename(file),
"path": file,
"filename": file_path.name,
"path": str(file_path),
"type": file_type,
"status": "waiting",
"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
@ -426,56 +438,32 @@ class KnowledgeBase:
file_infos = {}
# 使用UnstructuredURLLoader加载URL内容
# loader = UnstructuredURLLoader(urls=urls, continue_on_failure=True)
for url_idx, url in enumerate(urls):
for url in urls:
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:
# 加载单个URL内容
single_loader = UnstructuredURLLoader(urls=[url], continue_on_failure=False)
documents = await single_loader.aload()
# 将文档内容合并
# 合并文档内容
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:
logger.error(f"处理URL {url} 时出错: {e}")
file_infos[file_id] = {
"file_id": file_id,
"filename": url[:100] + "..." if len(url) > 100 else url,
"path": url,
"type": "url",
"status": "failed",
"created_at": time.time(),
"error": str(e),
"nodes": []
}
file_info_dict["status"] = "failed"
file_info_dict["error"] = str(e)
file_infos[file_id] = file_info_dict
return file_infos
@ -765,3 +753,30 @@ class KnowledgeBase:
updated_db = self.update_database_record(db_id, name, description)
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:
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}`"
f"如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`")
logger.debug("如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
f"如果还是不行,建议手动下载到某个文件夹,比如 {os.getenv('MODEL_DIR', '/models')}/BAAI/bge-m3 目录下;")
self.model = HuggingFaceEmbeddings(

View File

@ -123,17 +123,23 @@ class OCRPlugin:
raise FileNotFoundError(f"PDF file not found: {pdf_path}")
try:
# 检查是否为文本PDF
if is_text_pdf(pdf_path):
logger.info(f"PDF file is text, use llama_index.readers.file to read")
return pdfreader(pdf_path)
# 检查是否为文本PDF可能会出现错误比如每一页都有可读取的水印文字但是内容本身是扫描件需要使用OCR处理
# if is_text_pdf(pdf_path):
# from src.core.indexing import pdfreader
# logger.info("PDF file is text, use llama_index.readers.file to read")
# return pdfreader(pdf_path)
# 将PDF转换为图像
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 = []
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 = []
@ -147,46 +153,58 @@ class OCRPlugin:
logger.error(f"PDF processing error: {str(e)}")
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):
imgs = []
img_dir = os.path.join(output_dir, 'imgs')
if not os.path.exists(img_dir):
os.makedirs(img_dir)
pdfDoc = fitz.open(pdf_path)
totalPage = pdfDoc.page_count
for pg in tqdm(range(totalPage), desc='to imgs', ncols=100):
page = pdfDoc[pg]
rotate = int(0)
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]
# 读取PDF文件
with open(pdf_path, 'rb') as f:
files = {'file': f}
data = {
'parse_method': 'ocr', # 使用OCR模式
'is_json_md_dump': False, # 不需要保存中间文件
'return_layout': False, # 不需要返回布局信息
'return_info': False, # 不需要返回额外信息
'return_content_list': False, # 不需要返回内容列表
'return_images': False, # 不需要返回图片
}
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):
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):
"""读取普通文本文件并返回text文本"""
assert os.path.exists(file_path), "File not found"

View File

@ -9,14 +9,14 @@ def is_text_pdf(pdf_path):
total_pages = len(doc)
if total_pages == 0:
return False
text_pages = 0
for page_num in range(total_pages):
page = doc.load_page(page_num)
text = page.get_text()
if text.strip(): # 检查是否有文本内容
text_pages += 1
# 计算有文本内容的页面比例
text_ratio = text_pages / total_pages
# 如果超过50%的页面有文本内容则认为是文本PDF

View File

@ -1,5 +1,4 @@
import os
from typing import List, Dict
from tavily import TavilyClient
from src.utils.logging_config import logger
@ -11,7 +10,7 @@ class WebSearcher:
self.client = TavilyClient(api_key)
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 搜索相关内容
@ -45,7 +44,7 @@ class WebSearcher:
logger.error(f"Error during web search: {str(e)}")
return []
def format_search_results(self, results: List[Dict]) -> str:
def format_search_results(self, results: list[dict]) -> str:
"""
将搜索结果格式化为文本
@ -64,4 +63,4 @@ class WebSearcher:
formatted_text += f" {result['content']}\n"
formatted_text += f" 来源: {result['url']}\n\n"
return formatted_text
return formatted_text

View File

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

View File

@ -93,7 +93,7 @@
v-if="configStore.config.enable_web_search"
@click="meta.use_web=!meta.use_web"
>
<CompassOutlined style="margin-right: 3px;"/>
<Compass style="margin-right: 3px;" size="14"/>
联网搜索
</div>
<div
@ -101,7 +101,7 @@
v-if="configStore.config.enable_knowledge_graph"
@click="meta.use_graph=!meta.use_graph"
>
<DeploymentUnitOutlined style="margin-right: 3px;"/>
<Waypoints style="margin-right: 3px;" size="14"/>
知识图谱
</div>
<a-dropdown
@ -109,7 +109,7 @@
:class="{'opt-item': true, 'active': meta.selectedKB !== null}"
>
<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>
</a>
<template #overlay>
@ -147,7 +147,7 @@ import {
PlusCircleOutlined,
DeploymentUnitOutlined,
} 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 { useConfigStore } from '@/stores/config'
import { useUserStore } from '@/stores/user'
@ -972,6 +972,13 @@ const findLastIndex = (array, predicate) => {
max-width: 1200px;
}
.opt-item {
display: flex;
justify-content: center;
align-items: center;
gap: 4px;
}
.note {
width: 100%;
font-size: small;

View File

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