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:
parent
1f3e24cc37
commit
b0db3c6b0d
@ -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
|
||||
|
||||
@ -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
38
docker/pull_image.sh
Normal 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!"
|
||||
@ -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 ./
|
||||
|
||||
|
||||
@ -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
|
||||
```
|
||||
@ -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"] # 忽略的规则
|
||||
63
scripts/mineru-api/Dockerfile
Normal file
63
scripts/mineru-api/Dockerfile
Normal 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
303
scripts/mineru-api/app.py
Normal 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)
|
||||
44
scripts/mineru-api/magic-pdf.json
Normal file
44
scripts/mineru-api/magic-pdf.json
Normal 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"
|
||||
}
|
||||
@ -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()
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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("无效的令牌")
|
||||
|
||||
@ -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=(
|
||||
"查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,"
|
||||
"不要直接使用用户的原始输入去查询。"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 文件中的映射")
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
@ -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(
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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;
|
||||
|
||||
@ -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;
|
||||
|
||||
@ -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([]);
|
||||
|
||||
Loading…
Reference in New Issue
Block a user