diff --git a/docker-compose.yml b/docker-compose.yml index 2b804876..0023fbd5 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -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 diff --git a/docker/api.Dockerfile b/docker/api.Dockerfile index fcc3ca69..01b30d1d 100644 --- a/docker/api.Dockerfile +++ b/docker/api.Dockerfile @@ -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 diff --git a/docker/pull_image.sh b/docker/pull_image.sh new file mode 100644 index 00000000..69471a9d --- /dev/null +++ b/docker/pull_image.sh @@ -0,0 +1,38 @@ +#!/bin/bash + +if [ $# -ne 1 ]; then + echo "Usage: $0 " + 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!" \ No newline at end of file diff --git a/docker/web.Dockerfile b/docker/web.Dockerfile index 472bffab..766b4bf2 100644 --- a/docker/web.Dockerfile +++ b/docker/web.Dockerfile @@ -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 ./ diff --git a/docs/how-to.md b/docs/how-to.md index f4a45c3b..094a7995 100644 --- a/docs/how-to.md +++ b/docs/how-to.md @@ -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 +``` \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 7274896e..87910759 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] # 忽略的规则 \ No newline at end of file diff --git a/scripts/mineru-api/Dockerfile b/scripts/mineru-api/Dockerfile new file mode 100644 index 00000000..1a32eaf2 --- /dev/null +++ b/scripts/mineru-api/Dockerfile @@ -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 \"$@\"", "--"] diff --git a/scripts/mineru-api/app.py b/scripts/mineru-api/app.py new file mode 100644 index 00000000..072a238a --- /dev/null +++ b/scripts/mineru-api/app.py @@ -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) diff --git a/scripts/mineru-api/magic-pdf.json b/scripts/mineru-api/magic-pdf.json new file mode 100644 index 00000000..a2dc7de0 --- /dev/null +++ b/scripts/mineru-api/magic-pdf.json @@ -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" +} diff --git a/server/models/__init__.py b/server/models/__init__.py index 66ad4806..7c2377ae 100644 --- a/server/models/__init__.py +++ b/server/models/__init__.py @@ -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 \ No newline at end of file +Base = declarative_base() \ No newline at end of file diff --git a/server/models/kb_models.py b/server/models/kb_models.py index a18b6542..eaa1c874 100644 --- a/server/models/kb_models.py +++ b/server/models/kb_models.py @@ -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 diff --git a/server/routers/data_router.py b/server/routers/data_router.py index 0112519b..29584608 100644 --- a/server/routers/data_router.py +++ b/server/routers/data_router.py @@ -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 diff --git a/server/utils/auth_utils.py b/server/utils/auth_utils.py index 18ce903a..4ea6fe98 100644 --- a/server/utils/auth_utils.py +++ b/server/utils/auth_utils.py @@ -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("无效的令牌") \ No newline at end of file + raise ValueError("无效的令牌") diff --git a/src/agents/tools_factory.py b/src/agents/tools_factory.py index 739f0b75..5bdc3ef7 100644 --- a/src/agents/tools_factory.py +++ b/src/agents/tools_factory.py @@ -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=( + "查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询," + "不要直接使用用户的原始输入去查询。" + ) + ) diff --git a/src/agents/utils.py b/src/agents/utils.py index 4785805f..f49168df 100644 --- a/src/agents/utils.py +++ b/src/agents/utils.py @@ -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 diff --git a/src/config/__init__.py b/src/config/__init__.py index a8aa7c54..5c47911a 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -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 文件中的映射") diff --git a/src/core/indexing.py b/src/core/indexing.py index aafd61eb..fd783134 100644 --- a/src/core/indexing.py +++ b/src/core/indexing.py @@ -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) diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 22aaab7b..59fdfa1e 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -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 \ No newline at end of file diff --git a/src/models/embedding.py b/src/models/embedding.py index 462a27dd..4fff41fb 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -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( diff --git a/src/plugins/_ocr.py b/src/plugins/_ocr.py index 6de4febc..f18b0dd5 100644 --- a/src/plugins/_ocr.py +++ b/src/plugins/_ocr.py @@ -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" diff --git a/src/utils/__init__.py b/src/utils/__init__.py index 95083746..c9dc555d 100644 --- a/src/utils/__init__.py +++ b/src/utils/__init__.py @@ -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 diff --git a/src/utils/web_search.py b/src/utils/web_search.py index cb343ddd..d8b0e3c4 100644 --- a/src/utils/web_search.py +++ b/src/utils/web_search.py @@ -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 \ No newline at end of file + return formatted_text diff --git a/web/src/assets/base.css b/web/src/assets/base.css index 95bada1f..ae96fdf7 100644 --- a/web/src/assets/base.css +++ b/web/src/assets/base.css @@ -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; diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 2496c916..123d2d5c 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -93,7 +93,7 @@ v-if="configStore.config.enable_web_search" @click="meta.use_web=!meta.use_web" > - + 联网搜索
- + 知识图谱
- + {{ meta.selectedKB === null ? '不使用知识库' : opts.databases[meta.selectedKB]?.name }}