增加图片文本化接口,优化切片接口图片文本抽取,kgrag环境安装rapidocr-onnxruntime

This commit is contained in:
Defeng 2026-07-17 15:02:19 +08:00
parent 1b3cd3bfc3
commit 626139226d
5 changed files with 472 additions and 20 deletions

79
app.py
View File

@ -57,7 +57,7 @@ from concurrent.futures import ThreadPoolExecutor
import base64
from extract_excel_node_relation import get_excel_node_relation
from doc2pdf import Doc2PDF
from chunk_text import get_chunk_bbox, split_text_preserve_sentences, data_replace, chunk_check,merge_short_slices,find_and_read_content_list,process_pdf_file,process_other_file
from chunk_text import get_chunk_bbox, split_text_preserve_sentences, data_replace, chunk_check,merge_short_slices,find_and_read_content_list,process_pdf_file,process_other_file,safe_original_filename,image_base64_to_data_url,extract_image_text,prepare_split_image
from filename_proceess_and_kgquery import get_entity
from kg_build.extract_filtertext import get_filtertext_node_relation
import os
@ -92,7 +92,7 @@ from default_ontology_config import (
get_relationships,
get_relationship_types,
)
from config import NEO4J_CONFIG,URL,API_URLS,API_OTHER_URLS
from config import NEO4J_CONFIG,URL,API_URLS,API_OTHER_URLS,VLM_CONCURRENCY
converter = Doc2PDF()
DATA_DIR = "/app/files"
# ================== 全局资源容器 ==================
@ -167,6 +167,7 @@ async def lifespan(app):
app = FastAPI(max_request_size=1024 * 1024 * 10, lifespan=lifespan)
VLM_SEMAPHORE = asyncio.Semaphore(max(1, VLM_CONCURRENCY))
setup_logging()
logger = logging.getLogger(__name__)
@ -2028,6 +2029,77 @@ class RequestWrapper(BaseModel):
pdf_contents: List[PdfContentItem]
class Base64ImageRequest(BaseModel):
content: Optional[str] = Field(
None,
description="Optional text around the image, for example: 图3.1 船舶主发动机组成图images/test.jpg",
)
filename: str = Field(..., description="Original image filename, for example test.png")
image_base64: str = Field(..., description="Base64 image content, with or without data:image/... prefix")
async def analyze_image_with_vlm(image_bytes: bytes, suffix: str) -> str:
image_url = image_base64_to_data_url(image_bytes, suffix)
async with VLM_SEMAPHORE:
return await OpenaiAPI.open_api_vl_without_thinking(image_url)
@app.post("/split_image")
async def image_base64(request_data: Base64ImageRequest):
filename = safe_original_filename(request_data.filename)
suffix = PathLib(filename).suffix.lower()
if suffix not in {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".gif", ".tif", ".tiff"}:
raise HTTPException(status_code=400, detail="Unsupported image type")
image_bytes, processed_suffix = prepare_split_image(request_data.image_base64, suffix)
if not image_bytes:
raise HTTPException(status_code=400, detail="decoded image is empty")
temp_path = None
try:
with tempfile.NamedTemporaryFile(delete=False, suffix=processed_suffix) as tmp:
tmp.write(image_bytes)
temp_path = PathLib(tmp.name)
ocr_data = await extract_image_text(temp_path)
try:
temp_path.unlink()
except Exception as e:
logger.warning(f"Failed to delete temp image {temp_path}: {e}")
finally:
temp_path = None
content = (request_data.content or "").strip()
ocr_full_text = (ocr_data.get("full_text") or "").strip()
if not content and not ocr_full_text:
vlm_text = await analyze_image_with_vlm(image_bytes, processed_suffix)
image_content = "图片文本描述为:" + (vlm_text or "").strip()
elif content and ocr_full_text:
image_content = f"图片上下文内容为:{content}, {ocr_full_text}"
elif content:
image_content = "图片上下文内容为:" + content
else:
image_content = ocr_full_text
image_content = "图片文件名为:" + filename + ",图片相关内容为:" + image_content
return JSONResponse(
{
"code": 200,
"message": "ok",
"data": {
"image_content": image_content,
"filename": filename,
},
}
)
finally:
if temp_path and temp_path.exists():
try:
temp_path.unlink()
except Exception as e:
logger.warning(f"Failed to delete temp image {temp_path}: {e}")
@app.post("/split_content_list")
@ -2920,6 +2992,3 @@ if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=9085)

View File

@ -8,21 +8,293 @@ import re
from typing import List, Dict, Optional, Union, Any, Tuple
# 指定 JSON 文件路径
# file_path = '/storage01/home/hdf/project/wxkgrag/kgextract/output/122-06A0014-B01001_发动机-维修手册_content_list.json'
from filename_proceess_and_kgquery import get_entity
from pathlib import PurePath, Path as PathLib
try:
from filename_proceess_and_kgquery import get_entity
except ImportError:
get_entity = None
from pathlib import PurePath, PurePosixPath, Path as PathLib
from fastapi import FastAPI, File, Path, UploadFile, HTTPException, Form, Request, Header, Body
from fileparse_util import process_document
from io import BytesIO
from urllib.parse import unquote, urlparse
import asyncio
import base64
import binascii
import threading
import aiofiles
import httpx
import logging
from config import SHIP_MODEL_NAME,API_URLS
from rapidocr_onnxruntime import RapidOCR
from PIL import Image, ImageOps, UnidentifiedImageError
from config import (
API_URLS,
IMAGE_DIR,
MAX_SPLIT_IMAGE_BASE64_CHARS,
MAX_SPLIT_IMAGE_BYTES,
MAX_SPLIT_IMAGE_INPUT_BYTES,
OCR_CONCURRENCY,
SPLIT_IMAGE_COMPRESS_THRESHOLD_BYTES,
SPLIT_IMAGE_JPEG_QUALITY,
SPLIT_IMAGE_MAX_DIMENSION,
)
try:
from config import SHIP_MODEL_NAME
except ImportError:
SHIP_MODEL_NAME = ""
logger = logging.getLogger(__name__)
ocr_engine = RapidOCR()
OCR_SEMAPHORE = asyncio.Semaphore(max(1, OCR_CONCURRENCY))
_ocr_thread_lock = threading.Lock()
_image_ocr_cache: Dict[str, str] = {}
_image_ocr_cache_lock = threading.Lock()
_api_url_inflight = [0] * len(API_URLS)
_api_url_lock = threading.Lock()
def _image_mime_type(suffix: str) -> str:
return {
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".png": "image/png",
".bmp": "image/bmp",
".webp": "image/webp",
".gif": "image/gif",
".tif": "image/tiff",
".tiff": "image/tiff",
}.get((suffix or "").lower(), "image/png")
def _image_save_format(suffix: str) -> tuple[str, dict]:
suffix = (suffix or "").lower()
if suffix in {".jpg", ".jpeg"}:
return "JPEG", {"quality": SPLIT_IMAGE_JPEG_QUALITY, "optimize": True}
if suffix == ".webp":
return "WEBP", {"quality": SPLIT_IMAGE_JPEG_QUALITY, "method": 6}
if suffix == ".png":
return "PNG", {"optimize": True, "compress_level": 9}
if suffix in {".tif", ".tiff"}:
return "TIFF", {"compression": "tiff_lzw"}
if suffix == ".gif":
return "GIF", {"optimize": True}
if suffix == ".bmp":
return "BMP", {}
return "PNG", {"optimize": True, "compress_level": 9}
def _compression_suffix(suffix: str) -> str:
suffix = (suffix or "").lower()
if suffix in {".jpg", ".jpeg", ".png", ".webp"}:
return suffix
if suffix in {".bmp", ".gif", ".tif", ".tiff"}:
return ".png"
return ".png"
def _save_image_to_bytes(img: Image.Image, suffix: str) -> bytes:
image_format, save_options = _image_save_format(suffix)
if image_format in {"JPEG", "WEBP", "BMP"} and img.mode not in {"RGB", "L"}:
img = img.convert("RGB")
elif image_format in {"PNG", "TIFF", "GIF"} and img.mode not in {"RGB", "RGBA", "L", "P"}:
img = img.convert("RGBA")
output = BytesIO()
img.save(output, format=image_format, **save_options)
return output.getvalue()
def _save_image_as_jpeg(img: Image.Image) -> bytes:
if img.mode not in {"RGB", "L"}:
img = img.convert("RGB")
last_bytes = b""
for quality in (SPLIT_IMAGE_JPEG_QUALITY, 82, 74):
output = BytesIO()
img.save(output, format="JPEG", quality=quality, optimize=True)
last_bytes = output.getvalue()
if len(last_bytes) <= MAX_SPLIT_IMAGE_BYTES:
break
return last_bytes
def compress_image_bytes(image_bytes: bytes, suffix: str = "") -> tuple[bytes, str]:
if not image_bytes:
raise HTTPException(status_code=400, detail="decoded image is empty")
if len(image_bytes) > MAX_SPLIT_IMAGE_INPUT_BYTES:
raise HTTPException(status_code=413, detail="image input is too large")
suffix = (suffix or "").lower()
try:
with Image.open(BytesIO(image_bytes)) as img:
img = ImageOps.exif_transpose(img)
max_side = max(img.size)
should_resize = SPLIT_IMAGE_MAX_DIMENSION > 0 and max_side > SPLIT_IMAGE_MAX_DIMENSION
should_reencode = len(image_bytes) > SPLIT_IMAGE_COMPRESS_THRESHOLD_BYTES
if not should_resize and not should_reencode:
if len(image_bytes) > MAX_SPLIT_IMAGE_BYTES:
raise HTTPException(status_code=413, detail="image is too large")
return image_bytes, suffix
if should_resize:
scale = SPLIT_IMAGE_MAX_DIMENSION / max_side
new_size = (
max(1, int(img.width * scale)),
max(1, int(img.height * scale)),
)
resample = getattr(getattr(Image, "Resampling", Image), "LANCZOS")
img = img.resize(new_size, resample)
compressed_suffix = _compression_suffix(suffix)
compressed_bytes = _save_image_to_bytes(img, compressed_suffix)
if len(compressed_bytes) > MAX_SPLIT_IMAGE_BYTES:
compressed_bytes = _save_image_as_jpeg(img)
compressed_suffix = ".jpg"
except UnidentifiedImageError as exc:
raise HTTPException(status_code=400, detail="image content is not a supported image") from exc
except HTTPException:
raise
except Exception as exc:
logger.warning(f"Failed to compress image, using original bytes: {exc}")
compressed_bytes = image_bytes
compressed_suffix = suffix
if len(compressed_bytes) >= len(image_bytes) and len(image_bytes) <= MAX_SPLIT_IMAGE_BYTES:
return image_bytes, suffix
if len(compressed_bytes) > MAX_SPLIT_IMAGE_BYTES:
raise HTTPException(status_code=413, detail="image is too large after compression")
return compressed_bytes, compressed_suffix
def prepare_split_image(image_base64: str, suffix: str = "") -> tuple[bytes, str]:
raw = (image_base64 or "").strip()
if not raw:
raise HTTPException(status_code=400, detail="image_base64 is empty")
if "," in raw and raw.split(",", 1)[0].lower().startswith("data:image"):
raw = raw.split(",", 1)[1]
raw = "".join(raw.split())
if len(raw) > MAX_SPLIT_IMAGE_BASE64_CHARS:
raise HTTPException(status_code=413, detail="image_base64 is too large")
try:
image_bytes = base64.b64decode(raw, validate=True)
except (binascii.Error, ValueError) as exc:
raise HTTPException(status_code=400, detail="image_base64 is not valid base64") from exc
return compress_image_bytes(image_bytes, suffix)
def decode_base64_image(image_base64: str, suffix: str = "") -> bytes:
image_bytes, _ = prepare_split_image(image_base64, suffix)
return image_bytes
def safe_original_filename(filename: str) -> str:
name = os.path.basename((filename or "").strip())
if not name:
raise HTTPException(status_code=400, detail="filename is empty")
safe_name = re.sub(r"[^\w.\-\u4e00-\u9fff]+", "_", name).strip("._")
if not safe_name:
raise HTTPException(status_code=400, detail="filename is invalid")
return safe_name
def _strip_markdown_image_path(img_path: str) -> str:
value = (img_path or "").strip()
match = re.match(r"!\[[^\]]*\]\((.*?)\)", value)
return match.group(1).strip() if match else value
def resolve_local_image_path(img_path: str) -> Optional[PathLib]:
value = _strip_markdown_image_path(str(img_path or ""))
if not value:
return None
parsed_path = unquote(urlparse(value).path or value)
marker = "/images/"
if marker in parsed_path:
relative_path = parsed_path.split(marker, 1)[1]
else:
normalized = parsed_path.replace("\\", "/").lstrip("/")
if normalized.startswith("images/"):
relative_path = normalized.split("/", 1)[1]
else:
relative_path = normalized
relative_path = relative_path.replace("\\", "/").lstrip("/")
if not relative_path:
return None
image_root = PathLib(IMAGE_DIR).resolve()
local_path = (image_root / PurePosixPath(relative_path)).resolve()
try:
local_path.relative_to(image_root)
except ValueError:
logger.warning(f"Image path is outside IMAGE_DIR and will be ignored: {img_path}")
return None
return local_path if local_path.is_file() else None
def extract_image_text_sync(image_path: Union[str, PathLib]) -> dict:
resolved_path = str(PathLib(image_path).resolve())
with _image_ocr_cache_lock:
cached_text = _image_ocr_cache.get(resolved_path)
if cached_text is not None:
return {"texts": [], "full_text": cached_text, "elapse": None}
with _ocr_thread_lock:
ocr_result, elapse = ocr_engine(resolved_path)
texts = [item[1] for item in (ocr_result or [])]
full_text = "图片包含的文字内容为:" + ",".join(texts) if texts else ""
with _image_ocr_cache_lock:
_image_ocr_cache[resolved_path] = full_text
return {
"texts": texts,
"full_text": full_text,
"elapse": elapse,
}
def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optional[PathLib]) -> str:
if "图片包含的文字内容为:" in markdown_image:
return markdown_image
if not local_image_path:
return markdown_image
try:
ocr_data = extract_image_text_sync(local_image_path)
except Exception as exc:
logger.warning(f"Failed to OCR local image {local_image_path}: {exc}")
return markdown_image
ocr_full_text = (ocr_data.get("full_text") or "").strip()
if not ocr_full_text:
return markdown_image
return f"{markdown_image}\n{ocr_full_text}"
def image_base64_to_data_url(image_base64: Union[str, bytes], suffix: str) -> str:
if isinstance(image_base64, bytes):
raw = base64.b64encode(image_base64).decode("ascii")
else:
raw = (image_base64 or "").strip()
if "," in raw and raw.split(",", 1)[0].lower().startswith("data:image"):
raw = raw.split(",", 1)[1]
raw = "".join(raw.split())
mime_type = _image_mime_type(suffix)
return f"data:{mime_type};base64,{raw}"
async def extract_image_text(image_path: PathLib) -> dict:
async with OCR_SEMAPHORE:
return await asyncio.to_thread(extract_image_text_sync, image_path)
def format_entity_text(data):
# 1. 定义实体类型与属性的映射关系
# 格式: '实体类型': [ ('前缀1', '属性1'), ('前缀2', '属性2') ]
@ -710,18 +982,35 @@ def data_replace(data, prefix):
遍历 data 列表若元素包含 'img_path' 字段则移除前缀 'images/' 并拼接新前缀
使用 pathlib 安全处理路径
"""
from pathlib import PurePosixPath
for ins in data:
if isinstance(ins, dict) and "img_path" in ins:
img_path = ins["img_path"]
if isinstance(img_path, list):
ins["img_path"] = "\n".join(
data_replace([{"img_path": str(item)}], prefix)[0]["img_path"]
for item in img_path
if str(item).strip()
)
continue
img_path = str(img_path)
if img_path.strip().startswith("!["):
local_image_path = resolve_local_image_path(img_path)
ins["img_path"] = build_image_markdown_with_ocr(img_path, local_image_path)
continue
p = PurePosixPath(img_path)
# 如果以 images/ 开头,去掉第一级目录
if p.parts and p.parts[0] == "images":
relative_path = str(PurePosixPath(*p.parts[1:]))
ins["img_path"] = f"![图片]({prefix}{relative_path})"
markdown_image = f"![图片]({prefix}{relative_path})"
local_image_path = resolve_local_image_path(relative_path)
ins["img_path"] = build_image_markdown_with_ocr(markdown_image, local_image_path)
else:
# 否则保留原路径(或按需处理)
ins["img_path"] = f"![图片]({prefix}{img_path})"
markdown_image = f"![图片]({prefix}{img_path})"
local_image_path = resolve_local_image_path(img_path)
ins["img_path"] = build_image_markdown_with_ocr(markdown_image, local_image_path)
return data
def filter_content_bbox(data, max_length=3000):
@ -1057,7 +1346,8 @@ async def process_pdf_file(pdf_path: Path, image_prefix: str,filename:str) -> Op
slices = get_chunk_bbox(replaced_content_list)
slices_check = chunk_check(slices, 8000)
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
print(11111111111111111111111111111111111111)
print(len(images))
return {"slices": slices_check, "images": images}
except json.JSONDecodeError:
@ -1144,8 +1434,3 @@ if __name__ == "__main__":
for ins in slices_2:
print(ins)

View File

@ -74,4 +74,18 @@ API_OTHER_URLS = [
# "http://192.168.0.111:9975/analyze-otherfile",
]
SHIP_MODEL_NAME = "E:\ZKYNLP\Hjunproject\project0506\kgrag\ship_xinghao.json"
# ==================== 切片图片处理配置 ====================
MAX_SPLIT_IMAGE_INPUT_BYTES = 50 * 1024 * 1024
MAX_SPLIT_IMAGE_BYTES = 10 * 1024 * 1024
MAX_SPLIT_IMAGE_BASE64_CHARS = (MAX_SPLIT_IMAGE_INPUT_BYTES * 4 // 3) + 4096
SPLIT_IMAGE_COMPRESS_THRESHOLD_BYTES = MAX_SPLIT_IMAGE_BYTES
SPLIT_IMAGE_MAX_DIMENSION = 4096
SPLIT_IMAGE_JPEG_QUALITY = 88
OCR_CONCURRENCY = 1
VLM_CONCURRENCY = 2
SHIP_MODEL_NAME = "/app/ship_xinghao.json"
TREE_JSON_PATH = "/app/tree_data.json"
# ==================== mineru和kgrag目录映射地址 ====================
SEARCH_DIR ="/app/mineru_output"
IMAGE_DIR = "/app/mineru_output/images"

View File

@ -5,7 +5,7 @@ import asyncio
import logging
from typing import Optional, Tuple, Dict
import aiofiles # pip install aiofiles
from config import SEARCH_DIR,IMAGE_DIR
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
@ -103,7 +103,6 @@ async def encode_images_to_base64(
async def process_document(
input_file_name: str,
search_dir: Optional[str] = "/app/mineru_output",
local_image_dir: Optional[str] = None,
concurrency: int = 32,
) -> Optional[Tuple[list, Dict[str, str]]]:
@ -120,6 +119,7 @@ async def process_document(
成功: (content_list, images_dict)
失败: None
"""
search_dir = SEARCH_DIR
content, content_path = await find_and_read_content_list(search_dir, input_file_name)
if content is None:
logger.error("❌ 未找到指定的 _content_list.json 文件。")

View File

@ -15,7 +15,19 @@ from neo4j_graphrag.embeddings.base import Embedder
from config import LLM_CONFIG, EMBEDDING_CONFIG
# 模块级单例客户端,避免每次调用都创建新连接
_async_client: AsyncOpenAI = None
def _get_async_client() -> AsyncOpenAI:
"""获取或创建 AsyncOpenAI 单例客户端"""
global _async_client
if _async_client is None:
_async_client = AsyncOpenAI(
api_key=LLM_CONFIG["api_key"],
base_url=LLM_CONFIG["base_url"],
)
return _async_client
class OpenaiAPI:
@staticmethod
@ -148,7 +160,79 @@ class OpenaiAPI:
extra_body={"chat_template_kwargs": {"enable_thinking": False}}
)
return response.choices[0].message.content
@staticmethod
async def open_api_vl_without_thinking(
image_url: str
) -> str:
"""
专门用于从图片中提取设备名称故障现象
自动过滤思考过程仅返回核心结果
参数:
image_url: Base64 格式的图片数据 (data:image/...;base64,...)
model: 模型名称默认为配置中的模型
custom_instruction: 额外的特定指令 (可选)
"""
# 1. 确定模型名称
model = LLM_CONFIG.get("model")
# 2. 构建强约束的 System Prompt
# 核心目标:禁止思考标签,禁止废话,只给结果
system_prompt = (
"你是一个视觉分析专家。你的任务是从图片中识别内容并进行描述。\n"
"【严格约束】\n"
"1. 绝对禁止输出 <think>, <think>, <reasoning> 等任何思考过程标签。\n"
"2. 绝对禁止输出'好的''根据图片''分析如下'等开场白或结束语。\n"
"3. 直接输出最终结论,格式必须严格遵守下面的模板。"
)
# 3. 构建针对性的 User Prompt
default_task = (
"请分析这张图片,描述图片内容,重点关注以下信息\n"
"1. 图片整体阐述了什么主题,描述了哪些内容\n"
"请用专业、客观、简短的语言描述。"
)
final_user_prompt = default_task
# 4. 初始化客户端
client = _get_async_client()
# 5. 构建请求参数
kwargs = {
"model": model,
"messages": [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "text", "text": final_user_prompt},
{"type": "image_url", "image_url": {"url": image_url}}
]
}
],
"temperature": 0.1, # 低温度以保证事实准确性
"stream": False,
"max_tokens": LLM_CONFIG.get("max_tokens"),
}
# 尝试通过参数关闭思考 (取决于后端支持情况)
kwargs["extra_body"] = {"chat_template_kwargs": {"enable_thinking": False}}
try:
print(f"🚀 正在调用 {model} 进行图片内容分析...")
response = await client.chat.completions.create(**kwargs)
raw_content = response.choices[0].message.content or ""
print(f"✅ 分析完成:\n{raw_content}")
return raw_content
except Exception as e:
import traceback
error_msg = f"❌ 图片内容分析失败:{str(e)}\n{traceback.format_exc()}"
print(error_msg)
return f"Error: {str(e)}"
@staticmethod
async def open_api_chat_async_json(query: str, model: str = None, temperature: float = None) -> str:
if model is None: