1359 lines
55 KiB
Python
1359 lines
55 KiB
Python
from io import BytesIO
|
||
import logging
|
||
import os
|
||
import re
|
||
import tempfile
|
||
from fastapi import FastAPI, File, Path, UploadFile, HTTPException, Form, Request, Header
|
||
from fastapi.responses import JSONResponse, StreamingResponse
|
||
from fastapi.middleware.cors import CORSMiddleware
|
||
from pydantic import BaseModel
|
||
from dotenv import load_dotenv
|
||
from pathlib import Path as PATH
|
||
from starlette.datastructures import Headers
|
||
import httpx
|
||
import pdfplumber
|
||
from markdown_process.markdown_preprocess import process_pdf_logic, convert_doc_to_pdf, convert_docx_to_pdf
|
||
from markdown_process.split_markdown import hierarchical_save, split_text, save_md, split_page_chunk
|
||
from markdown_process.chunk_postprocess import context_chunk_process
|
||
from database.milvus_api import Milvus_Database, sanitize_collection_name
|
||
from modelsAPI.model_api import OpenaiAPI, query_rewrite, prepare_rag_prompt, rerank_text, rerank_3d,resetQuery
|
||
from typing import List, Optional
|
||
from collections import defaultdict
|
||
from fuzzywuzzy import fuzz, process
|
||
from collections import defaultdict,Counter
|
||
import uuid
|
||
import hanlp
|
||
from contentprocess.contentprocess import extract_unique_nouns_from_file
|
||
from itertools import groupby
|
||
import requests
|
||
app = FastAPI(max_request_size=1024 * 1024 * 10)
|
||
load_dotenv()
|
||
|
||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||
logger = logging.getLogger(__name__)
|
||
|
||
app.add_middleware(
|
||
CORSMiddleware,
|
||
allow_origins=["*"], # Allow all origins
|
||
allow_credentials=True,
|
||
allow_methods=["*"], # Allow all HTTP methods
|
||
allow_headers=["*"], # Allow all headers
|
||
)
|
||
print("正在加载 HanLP 模型...")
|
||
HanLPmodel = hanlp.load(hanlp.pretrained.mtl.CLOSE_TOK_POS_NER_SRL_DEP_SDP_CON_ELECTRA_SMALL_ZH)
|
||
|
||
@app.get("/debug/routes")
|
||
async def list_routes():
|
||
return [{"path": route.path, "methods": route.methods}
|
||
for route in app.routes]
|
||
|
||
|
||
# class MarkdownRequest(BaseModel):
|
||
# pdf_file: UploadFile = File(...) # 接收Markdown格式的文本
|
||
# database_name: Optional[str] = "JXTest" # 接收类型列表,类型应该是 ['01', '02', '03']
|
||
# collection_name: Optional[str] = "MarkdownTest")#
|
||
|
||
class SearchRequest(BaseModel):
|
||
query: Optional[str] = None
|
||
top_k: Optional[int] = 5
|
||
database_name: Optional[str] = "JXTest"
|
||
collection_name: Optional[str] = "RAGTest"
|
||
query_vector: Optional[List[float]] = None
|
||
|
||
|
||
@app.post("/search")
|
||
async def search_markdown(request: SearchRequest):
|
||
logger.info(f"Received search query: {request.query}")
|
||
# 在这里实现搜索逻辑
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
# query 改写
|
||
query = query_rewrite(query=request.query)
|
||
query_vector = OpenaiAPI.get_embeddings(list(query))
|
||
exist_database = milvus_db.list_databases()
|
||
if request.database_name not in exist_database:
|
||
logger.error(f"database {request.database_name} not exist, please create it first")
|
||
raise HTTPException(400, f" {request.database_name} not exist, please create it first")
|
||
milvus_db.using_database(database_name=request.database_name)
|
||
|
||
search_results = milvus_db.search(collection_name=request.collection_name, query_vector=query_vector,
|
||
top_k=request.top_k)
|
||
reranker_result_list, relevance_scores = OpenaiAPI.rerank_query(query=request.query, documents=search_results,
|
||
top_n=request.top_k)
|
||
rag_prompt = prepare_rag_prompt(query=request.query, documents=reranker_result_list)
|
||
return StreamingResponse(OpenaiAPI.open_api_chat_stream(query=rag_prompt), media_type="text/plain")
|
||
|
||
|
||
class CreateCollectionRequest(BaseModel):
|
||
collection_name: str
|
||
database_name: str
|
||
chunk_type: str = "01"
|
||
|
||
|
||
@app.post("/create_collection")
|
||
async def create_collection(request: CreateCollectionRequest):
|
||
"""
|
||
创建数据库和集合
|
||
chunk_type: 01 表示智能存储方式,02表示按照章节存储,03表示按字数存储
|
||
"""
|
||
collection_name = request.collection_name
|
||
collection_name = sanitize_collection_name(collection_name)
|
||
database_name = request.database_name
|
||
chunk_type = request.chunk_type
|
||
|
||
logger.info(f"Received create collection request: {database_name},{collection_name}")
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
try:
|
||
# 创建数据库与collections
|
||
exist_database = milvus_db.list_databases()
|
||
if database_name not in exist_database:
|
||
res_milvus_create = milvus_db.create_database(database_name=database_name)
|
||
logger.info(f"create database result: {res_milvus_create}")
|
||
else:
|
||
logger.info(f"database {database_name} already exist")
|
||
|
||
milvus_db.using_database(database_name=database_name)
|
||
exist_collection = milvus_db.has_collection(collection_name=collection_name, database_name=database_name)
|
||
if not exist_collection:
|
||
if chunk_type == "03":
|
||
save_title = False
|
||
else:
|
||
save_title = True
|
||
res_milvus_create = milvus_db.create_hybrid_collection(collection_name=collection_name,
|
||
database_name=database_name, save_title=save_title)
|
||
logger.info(f"create collection result: {res_milvus_create}")
|
||
|
||
return JSONResponse(content={"message": "success", "code": "200"}, status_code=200)
|
||
else:
|
||
logger.info(f"collection {collection_name} already exist")
|
||
raise HTTPException(400, f"collection {collection_name} already exist")
|
||
except Exception as e:
|
||
logger.error(f"Error creating collection: {str(e)}")
|
||
raise HTTPException(500, "Error creating collection")
|
||
finally:
|
||
|
||
milvus_db.close()
|
||
|
||
|
||
@app.get("/has_collection")
|
||
async def collection(database_name, collection_name):
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
try:
|
||
collection = milvus_db.has_collection(collection_name=collection_name, database_name=database_name)
|
||
return {"message": "success", "code": "200", "data": {"has_collection": collection}}
|
||
except Exception as e:
|
||
logger.error(f"Error getting collections: {str(e)}")
|
||
raise HTTPException(500, "Error getting collections")
|
||
finally:
|
||
milvus_db.close()
|
||
|
||
|
||
@app.delete("/delete_collection/{collection_name}")
|
||
async def delete_collection(collection_name: str, database_name: Optional[str]):
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
try:
|
||
collection_name = sanitize_collection_name(collection_name)
|
||
milvus_db.drop_collection(collection_name=collection_name, database_name=database_name)
|
||
return JSONResponse(content={"message": "success", "code": "200"}, status_code=200)
|
||
except Exception as e:
|
||
logger.error(f"Error deleting collection: {str(e)}")
|
||
raise HTTPException(500, "Error deleting collection")
|
||
finally:
|
||
milvus_db.close()
|
||
|
||
|
||
@app.get("/get_collections")
|
||
async def get_collections(database_name):
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
try:
|
||
collections = milvus_db.list_collections(database_name=database_name)
|
||
if collections: # 确保列表非空
|
||
collections.pop(0)
|
||
return {"message": "success", "code": "200", "data": {"collections": collections}}
|
||
except Exception as e:
|
||
logger.error(f"Error getting collections: {str(e)}")
|
||
raise HTTPException(500, "Error getting collections")
|
||
finally:
|
||
milvus_db.close()
|
||
|
||
|
||
@app.post("/upload_pdf")
|
||
async def process_pdf(pdf_file: UploadFile = File(...), database_name: Optional[str] = Form("JXTest"),
|
||
collection_name: Optional[str] = Form("MarkdownTest"), chunk_type: str = "01",
|
||
chunk_size: int = 256, chunk_overlap: int = 50):
|
||
"""
|
||
接收 PDF 文件并转发到解析服务
|
||
chunk_type: list[str] 01: 智能分割 ; 02: 按照章节分割; 03: 按照字数分割
|
||
chunk_size: int 按照字数分割
|
||
chunk_overlap: int 分块重叠
|
||
"""
|
||
logger.info(
|
||
f"Received PDF file: {pdf_file.filename},database_name : {database_name},collection_name : {collection_name}")
|
||
# 验证文件类型
|
||
if not pdf_file.content_type.startswith("application/pdf"):
|
||
logger.error(f"Invalid file type: {pdf_file.content_type}")
|
||
raise HTTPException(400, "仅支持 PDF 文件")
|
||
filename = pdf_file.filename
|
||
save_title = True
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
try:
|
||
# 创建数据库与collections
|
||
|
||
exist_database = milvus_db.list_databases()
|
||
if database_name not in exist_database:
|
||
res_milvus_create = milvus_db.create_database(database_name=database_name)
|
||
logger.info(f"create database result: {res_milvus_create}")
|
||
else:
|
||
logger.info(f"database {database_name} already exist")
|
||
|
||
milvus_db.using_database(database_name=database_name)
|
||
exist_collection = milvus_db.has_collection(collection_name=collection_name)
|
||
if not exist_collection:
|
||
# 创建向量检索集合
|
||
# res_milvus_create = milvus_db.create_collections(collection_name,drop_=False)
|
||
# 创建混检集合
|
||
if chunk_type == "03":
|
||
save_title = False
|
||
res_milvus_create = milvus_db.create_hybrid_collection(collection_name=collection_name,
|
||
database_name=database_name, save_title=save_title)
|
||
logger.info(f"create collection result: {res_milvus_create}")
|
||
else:
|
||
logger.info(f"collection {collection_name} already exist")
|
||
|
||
md_content, content_chunk_list = await process_pdf_logic(pdf_file)
|
||
logger.info("PDF file processed successfully")
|
||
# 以下部分需调整
|
||
md_header_splits = split_text(resource=filename, mdDocs=md_content, logger=logger, database_name=database_name,
|
||
collection_name=collection_name, chunk_type=chunk_type, chunk_length=chunk_size,
|
||
chunk_overlap=chunk_overlap)
|
||
# md_header_splits = split_by_md_header(mdDocs=md_content,logger=logger,database_name=database_name,collection_name=collection_name)
|
||
logger.info("Markdown content split successfully")
|
||
|
||
return {"message": "PDF save to Milvus successfully", "code": 200,
|
||
"data": {"md_head": md_header_splits, "md_content": md_content}}
|
||
|
||
except httpx.HTTPStatusError as e:
|
||
logger.error(f"HTTP status error from parsing service: {e.response.text}")
|
||
raise HTTPException(e.response.status_code, f"解析服务返回错误: {e.response.text}")
|
||
except httpx.RequestError as e:
|
||
logger.error(f"Failed to connect to parsing service: {str(e)}")
|
||
raise HTTPException(503, f"无法连接到解析服务: {str(e)}")
|
||
except Exception as e:
|
||
logger.error(f"Internal server error: {str(e)}")
|
||
raise HTTPException(500, f"服务器内部错误: {str(e)}")
|
||
finally:
|
||
milvus_db.close()
|
||
|
||
|
||
@app.delete("/delete_resource/{resource}")
|
||
async def delete_resource(resource: str, collection_name, database_name: str, ):
|
||
"""
|
||
删除指定数据库和集合
|
||
"""
|
||
logger.info(
|
||
f"delete resource : {resource}, delete database_name : {database_name},collection_name : {collection_name}")
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
try:
|
||
milvus_db.using_database(database_name=database_name)
|
||
collection_name = sanitize_collection_name(collection_name)
|
||
res = milvus_db.delete_resource(collection_name=collection_name, resource=resource, database_name=database_name)
|
||
|
||
return {"code": 200, "message": "success", "data": {"result": res}}
|
||
except Exception as e:
|
||
logger.error(f"Delete failed: {str(e)}", exc_info=True)
|
||
raise HTTPException(
|
||
status_code=500,
|
||
detail=f"删除失败: {str(e)}"
|
||
)
|
||
finally:
|
||
milvus_db.close()
|
||
|
||
|
||
@app.post("/get_markdown")
|
||
async def get_markdown(file: UploadFile = File(...)):
|
||
filename = file.filename
|
||
file_ext = os.path.splitext(filename)[1].lower()
|
||
headers = Headers({"content-type": "application/pdf"})
|
||
if file.content_type.startswith("application/pdf"):
|
||
logger.info(f"Received file: {filename}")
|
||
md_content, content_chunk_list = await process_pdf_logic(file)
|
||
logger.info("PDF file processed successfully")
|
||
if not md_content:
|
||
content = await file.read()
|
||
await file.seek(0)
|
||
md_content = ""
|
||
try:
|
||
with pdfplumber.open(BytesIO(content)) as pdf:
|
||
for page in pdf.pages:
|
||
md_content += page.extract_text() or ""
|
||
except Exception as e:
|
||
logger.error(f"Failed to parse PDF: {e}")
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="PDF file read failed"
|
||
)
|
||
# elif file.content_type == EXT_TO_CONTENT_TYPE[".docx"]:
|
||
elif file_ext == ".docx":
|
||
# 处理DOCX文件
|
||
converted_path = await convert_docx_to_pdf(os.getenv("CONVERT_DOC_URL", "http://192.168.0.46:59070"), file)
|
||
|
||
# 读取转换后的DOCX文件
|
||
with open(converted_path, 'rb') as f:
|
||
converted_file = UploadFile(filename=f"{PATH(filename).stem}.pdf", file=f, headers=headers)
|
||
md_content, content_chunk_list = await process_pdf_logic(converted_file)
|
||
os.unlink(converted_path)
|
||
if not md_content:
|
||
content = await file.read()
|
||
await file.seek(0)
|
||
md_content = ""
|
||
try:
|
||
with pdfplumber.open(BytesIO(content)) as pdf:
|
||
for page in pdf.pages:
|
||
md_content += page.extract_text() or ""
|
||
except Exception as e:
|
||
logger.error(f"Failed to parse PDF: {e}")
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="PDF file read failed")
|
||
# elif file.content_type == EXT_TO_CONTENT_TYPE[".doc"]:
|
||
elif file_ext == ".doc":
|
||
# 处理DOC文件,先转换为DOCX
|
||
try:
|
||
logger.info(f"Converting DOC file: {filename}")
|
||
converted_path = await convert_doc_to_pdf(os.getenv("CONVERT_DOC_URL", "http://192.168.0.46:59070"), file)
|
||
|
||
# 读取转换后的DOCX文件
|
||
with open(converted_path, 'rb') as f:
|
||
converted_file = UploadFile(filename=f"{PATH(filename).stem}.pdf", file=f, headers=headers)
|
||
md_content, content_chunk_list = await process_pdf_logic(converted_file)
|
||
|
||
# 获取图片信息
|
||
img_information = [chunk for chunk in content_chunk_list if chunk.get("type") == "image"]
|
||
|
||
# 清理临时文件
|
||
os.unlink(converted_path)
|
||
if not md_content:
|
||
content = await file.read()
|
||
await file.seek(0)
|
||
md_content = ""
|
||
try:
|
||
with pdfplumber.open(BytesIO(content)) as pdf:
|
||
for page in pdf.pages:
|
||
md_content += page.extract_text() or ""
|
||
except Exception as e:
|
||
logger.error(f"Failed to parse PDF: {e}")
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="PDF file read failed")
|
||
|
||
except Exception as e:
|
||
logger.error(f"Failed to convert DOC: {e}")
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"DOC conversion failed: {str(e)}"
|
||
)
|
||
else:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="Unsupported file type. Only PDF, DOC and DOCX are supported."
|
||
)
|
||
img_information = [chunk for chunk in content_chunk_list if chunk.get("type") == "image"]
|
||
|
||
return {
|
||
"message": "success",
|
||
"code": 200,
|
||
"data": {
|
||
"markdown_content": md_content,
|
||
"content_chunk_list": content_chunk_list,
|
||
"img_information": img_information
|
||
}
|
||
}
|
||
|
||
|
||
class SplitContextRequest(BaseModel):
|
||
file_name: str = "测试文件"
|
||
context: Optional[str] = "测试文本"
|
||
chunk_type: str = "01"
|
||
chunk_size: int = 1024
|
||
chunk_overlap: int = 100
|
||
max_level: Optional[int] = 3
|
||
chunk_list: List[dict] = []
|
||
|
||
|
||
@app.post("/split_result")
|
||
async def split_result(file: UploadFile = File(...),
|
||
chunk_type: str = Form("01"),
|
||
chunk_size: int = Form(1024),
|
||
chunk_overlap: int = Form(100),
|
||
max_level: Optional[int] = Form(3),
|
||
chunk_list: Optional[List[dict]] = Form([])
|
||
):
|
||
"""
|
||
拆分结果
|
||
"""
|
||
filename = file.filename
|
||
file_ext = os.path.splitext(filename)[1]
|
||
headers = Headers({"content-type": "application/pdf"})
|
||
if file.content_type.startswith("application/pdf"):
|
||
logger.info(f"Received file: {filename}")
|
||
md_content, content_chunk_list = await process_pdf_logic(file)
|
||
logger.info("PDF file processed successfully")
|
||
elif file_ext == ".docx":
|
||
# 处理DOCX文件
|
||
converted_path = await convert_docx_to_pdf(os.getenv("CONVERT_DOC_URL", "http://192.168.0.46:59070"), file)
|
||
|
||
# 读取转换后的DOCX文件
|
||
with open(converted_path, 'rb') as f:
|
||
converted_file = UploadFile(filename=f"{PATH(filename).stem}.pdf", file=f, headers=headers)
|
||
md_content, content_chunk_list = await process_pdf_logic(converted_file)
|
||
os.unlink(converted_path)
|
||
elif file_ext == ".docx":
|
||
# 处理DOC文件,先转换为DOCX
|
||
try:
|
||
logger.info(f"Converting DOC file: {filename}")
|
||
converted_path = await convert_doc_to_pdf(os.getenv("CONVERT_DOC_URL", "http://192.168.0.46:59070"), file)
|
||
|
||
# 读取转换后的DOCX文件
|
||
with open(converted_path, 'rb') as f:
|
||
converted_file = UploadFile(filename=f"{PATH(filename).stem}.pdf", file=f, headers=headers)
|
||
md_content, content_chunk_list = await process_pdf_logic(converted_file)
|
||
|
||
# 清理临时文件
|
||
os.unlink(converted_path)
|
||
|
||
except Exception as e:
|
||
logger.error(f"Failed to convert DOC: {e}")
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"DOC conversion failed: {str(e)}"
|
||
)
|
||
else:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="Unsupported file type. Only PDF, DOC and DOCX are supported."
|
||
)
|
||
if not md_content:
|
||
content = await file.read()
|
||
await file.seek(0)
|
||
md_content = ""
|
||
try:
|
||
with pdfplumber.open(BytesIO(content)) as pdf:
|
||
for page in pdf.pages:
|
||
md_content += page.extract_text() or ""
|
||
except Exception as e:
|
||
logger.error(f"Failed to parse PDF: {e}")
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="PDF file read failed")
|
||
if chunk_list:
|
||
md_header_splits = split_page_chunk(
|
||
resource=filename,
|
||
content_chunk_list=content_chunk_list,
|
||
chunk_size=chunk_size,
|
||
chunk_overlap=chunk_overlap,
|
||
logger=logger,
|
||
max_level=max_level,
|
||
chunk_type=chunk_type
|
||
)
|
||
else:
|
||
md_header_splits = split_text(
|
||
resource=filename,
|
||
mdDocs=md_content,
|
||
logger=logger,
|
||
chunk_type=chunk_type,
|
||
chunk_size=chunk_size,
|
||
chunk_overlap=chunk_overlap,
|
||
max_level=max_level
|
||
)
|
||
return {
|
||
"message": "success",
|
||
"code": 200,
|
||
"data": {
|
||
"split_result": md_header_splits,
|
||
}
|
||
}
|
||
|
||
|
||
@app.post("/split_context")
|
||
async def split_context(request: SplitContextRequest):
|
||
"""
|
||
拆分结果
|
||
"""
|
||
filename = request.file_name
|
||
context = request.context
|
||
chunk_type = request.chunk_type
|
||
chunk_size = request.chunk_size
|
||
chunk_overlap = request.chunk_overlap
|
||
max_level = request.max_level
|
||
chunk_list = request.chunk_list
|
||
if chunk_list:
|
||
md_header_splits = split_page_chunk(
|
||
resource=filename,
|
||
content_chunk_list=chunk_list,
|
||
chunk_type=chunk_type,
|
||
chunk_size=chunk_size,
|
||
chunk_overlap=chunk_overlap,
|
||
logger=logger,
|
||
max_level=max_level
|
||
)
|
||
else:
|
||
if context:
|
||
md_header_splits = split_text(resource=filename, mdDocs=context,
|
||
logger=logger, chunk_type=chunk_type,
|
||
chunk_size=chunk_size, chunk_overlap=chunk_overlap,
|
||
max_level=max_level)
|
||
return {"message": "sucess", "code": 200, "data": {"split_result": md_header_splits}}
|
||
|
||
|
||
@app.post("/split_context_chunk")
|
||
async def split_context_chunk(request: SplitContextRequest):
|
||
"""
|
||
拆分结果
|
||
"""
|
||
filename = request.file_name
|
||
context = request.context
|
||
chunk_type = request.chunk_type
|
||
chunk_size = request.chunk_size
|
||
chunk_overlap = request.chunk_overlap
|
||
max_level = request.max_level
|
||
chunk_list = request.chunk_list
|
||
md_header_splits = split_page_chunk(
|
||
resource=filename,
|
||
content_chunk_list=chunk_list,
|
||
chunk_type=chunk_type,
|
||
chunk_size=chunk_size,
|
||
chunk_overlap=chunk_overlap,
|
||
logger=logger,
|
||
max_level=max_level
|
||
)
|
||
context_chunk_splits = context_chunk_process(markdown_context=context, chunks=md_header_splits)
|
||
# context_chunk_splits = OpenaiAPI().context_chunk_process(markdown_context=context,chunks=md_header_splits)
|
||
return {"message": "sucess", "code": 200, "data": {"split_result": context_chunk_splits}}
|
||
|
||
|
||
class InsertRequest(BaseModel):
|
||
split_list: list
|
||
collection_name: str
|
||
database_name: str # 默认值
|
||
|
||
|
||
@app.post("/insert_md")
|
||
async def insert_md(request: InsertRequest):
|
||
"""
|
||
插入md
|
||
"""
|
||
collection_name = request.collection_name
|
||
logger.info(f"get {collection_name}")
|
||
database_name = request.database_name
|
||
split_list = request.split_list
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
for split_result in split_list:
|
||
reource = split_result.get("resource")
|
||
text = milvus_db.get_text_by_resource(collection_name=collection_name, database_name=database_name,
|
||
resource=reource, logger=logger)
|
||
if text:
|
||
logger.info("已经存在相同的内容")
|
||
return {"code": 200, "message": "success", "data": {"result": split_list}}
|
||
else:
|
||
logger.info("内容不在库内,跳过后续查询")
|
||
break
|
||
try:
|
||
# if SPLIT_LIST:
|
||
# logger.info(f"Inserting {SPLIT_LIST}")
|
||
res = save_md(results=request.split_list, collection_name=collection_name, database_name=request.database_name)
|
||
# res = save_md(results=SPLIT_LIST,collection_name=collection_name,database_name=request.database_name)
|
||
return {"code": 200, "message": "success", "data": {"result": res}}
|
||
# else:
|
||
# return {"code":200,"message":"success","data":{"result":split_list}}
|
||
except Exception as e:
|
||
logger.error(f"Insert failed: {str(e)}", exc_info=True)
|
||
raise HTTPException(400, f"插入失败: {str(e)}")
|
||
|
||
|
||
@app.post("/hierarchical_insert")
|
||
async def hierarchical_insert(request: InsertRequest):
|
||
collection_name = request.collection_name
|
||
logger.info(f"get {collection_name}")
|
||
database_name = request.database_name
|
||
split_list = request.split_list
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
for split_result in split_list:
|
||
reource = split_result.get("resource")
|
||
text = milvus_db.get_text_by_resource(collection_name=collection_name, database_name=database_name,
|
||
resource=reource, logger=logger)
|
||
if text:
|
||
logger.info("已经存在相同的内容")
|
||
return {"code": 200, "message": "success", "data": {"result": split_list}}
|
||
else:
|
||
logger.info("内容不在库内,跳过后续查询")
|
||
break
|
||
try:
|
||
res = hierarchical_save(results=request.split_list, collection_name=collection_name,
|
||
database_name=request.database_name)
|
||
return {"code": 200, "message": "success", "data": {"result": res}}
|
||
# else:
|
||
# return {"code":200,"message":"success","data":{"result":split_list}}
|
||
except Exception as e:
|
||
logger.error(f"Insert failed: {str(e)}", exc_info=True)
|
||
raise HTTPException(400, f"插入失败: {str(e)}")
|
||
|
||
|
||
@app.post("/search_source")
|
||
async def search_source(request: SearchRequest):
|
||
"""
|
||
搜索与给定查询最相关的 Markdown 文档
|
||
"""
|
||
logger.info(f"Received search query: {request.query},collection_name: {request.collection_name}")
|
||
|
||
# 初始化 Milvus 数据库
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
|
||
# 获取查询的向量表示
|
||
# query = query_rewrite(request.query)
|
||
query = request.query
|
||
query = resetQuery(query)
|
||
logger.info(f"改写query:{query}")
|
||
query_vector = OpenaiAPI.get_embeddings(list(query))
|
||
|
||
collection_name = request.collection_name
|
||
final_results = []
|
||
# 搜索相关文档
|
||
search_results, search_resources, recall_results = milvus_db.hybrid_search(
|
||
collection_name=collection_name,
|
||
database_name=request.database_name,
|
||
query_vector=query_vector,
|
||
top_k=request.top_k,
|
||
query=query,
|
||
embedding_rate=0.4
|
||
)
|
||
logger.info(f"文本混合检索结果: {recall_results}")
|
||
if not recall_results:
|
||
search_results, search_resources, recall_results = milvus_db.embedding_search(
|
||
collection_name=collection_name,
|
||
query_vector=query_vector,
|
||
database_name=request.database_name,
|
||
top_k=request.top_k
|
||
)
|
||
logger.info(f"文本纯向量检索结果: {search_results}")
|
||
|
||
for recall_result in recall_results:
|
||
text_1 = recall_result["text"]
|
||
text_2 = query
|
||
score = OpenaiAPI.embedding_cosine_similarity(text_1, text_2)
|
||
recall_result["score"] = score
|
||
# recall_results = sorted(recall_results, key=lambda x: x["score"], reverse=True)
|
||
# for idx, item in enumerate(recall_results, start=1):
|
||
# item["index"] = idx
|
||
logger.info("------------------------------------")
|
||
logger.info(f"第一轮精排结果: {recall_results}")
|
||
recall_results = rerank_text(query, recall_results)
|
||
recall_results = sorted(recall_results, key=lambda x: x["score"], reverse=True)
|
||
for idx, item in enumerate(recall_results, start=1):
|
||
item["index"] = idx
|
||
threshold = 0.2
|
||
recall_results = [item for item in recall_results if item.get('score', 0) > threshold]
|
||
logger.info("------------------------------------")
|
||
logger.info(f"第二轮精排结果: {recall_results}")
|
||
|
||
return {"code": 200, "message": "success", "data": recall_results, "query_vector": query_vector}
|
||
|
||
# 图谱检索服务地址(可通过环境变量覆盖)
|
||
GRAPH_SEARCH_URL = os.getenv("GRAPH_SEARCH_URL", "http://192.168.0.46:59085/search_graph")
|
||
|
||
|
||
async def call_graph_search(query: str, top_k: int = 10, graph_timeout: float = 30.0):
|
||
"""
|
||
异步调用图谱检索接口,并解析返回的 data.results
|
||
"""
|
||
if not query:
|
||
return []
|
||
|
||
try:
|
||
# 使用较长超时时间,避免复杂图谱查询被过早中断
|
||
async with httpx.AsyncClient(timeout=graph_timeout) as client:
|
||
resp = await client.post(
|
||
GRAPH_SEARCH_URL,
|
||
json={
|
||
"query": query,
|
||
"top_k": top_k,
|
||
"graph_timeout": graph_timeout,
|
||
# 如需入口节点可在这里扩展: "entry_nodes": ...
|
||
},
|
||
)
|
||
resp.raise_for_status()
|
||
result = resp.json()
|
||
except httpx.RequestError as e:
|
||
# 打印更详细的错误信息,便于排查(连接失败 / 读超时等)
|
||
logger.error(
|
||
f"[call_graph_search] 请求图谱检索服务失败: {e.__class__.__name__} - {e}",
|
||
exc_info=True
|
||
)
|
||
return []
|
||
except httpx.HTTPStatusError as e:
|
||
logger.error(f"[call_graph_search] 图谱检索服务返回异常状态码: {e.response.status_code}, body={e.response.text}")
|
||
return []
|
||
except Exception as e:
|
||
logger.error(f"[call_graph_search] 解析图谱检索结果失败: {e}", exc_info=True)
|
||
return []
|
||
|
||
# 按照约定结构提取 data.results
|
||
try:
|
||
data = result.get("data", {})
|
||
if isinstance(data, dict):
|
||
graph_results = data.get("results", [])
|
||
else:
|
||
graph_results = []
|
||
if not isinstance(graph_results, list):
|
||
graph_results = []
|
||
return graph_results
|
||
except Exception as e:
|
||
logger.error(f"[call_graph_search] 提取 data.results 失败: {e}", exc_info=True)
|
||
return []
|
||
|
||
|
||
|
||
"""
|
||
reranker_result_list = OpenaiAPI.rerank_query(
|
||
query=query,
|
||
documents=search_results,
|
||
top_n=request.top_k
|
||
)
|
||
# print("###########")
|
||
# print("reranker_result_list:",reranker_result_list)
|
||
|
||
|
||
for res in reranker_result_list:
|
||
index = res["index"]
|
||
resource,img_path = search_resources[index]["resource"],search_resources[index]["img_path"]
|
||
res["resource"] = resource
|
||
res["img_path"] = img_path
|
||
|
||
#内容去重
|
||
# 基于text字段去重
|
||
seen_texts = set()
|
||
unique_results = []
|
||
for res in reranker_result_list:
|
||
if res["text"] not in seen_texts:
|
||
seen_texts.add(res["text"])
|
||
unique_results.append(res)
|
||
milvus_db.close()
|
||
return {"code": 200, "message": "success", "data": unique_results}
|
||
"""
|
||
class SearchBatchRequest(BaseModel):
|
||
queries: Optional[List[str]] = None # 新增:支持批量文本
|
||
query_vectors: Optional[List[List[float]]] = None # 新增:支持批量向量
|
||
top_k: Optional[int] = 5
|
||
database_name: Optional[str] = "XIAN"
|
||
collection_name: Optional[str] = "RAGTest"
|
||
|
||
@app.post("/search_batch")
|
||
async def search_source(request: SearchBatchRequest):
|
||
"""
|
||
搜索与给定查询最相关的 Markdown 文档
|
||
"""
|
||
# ✅ 修复:使用 request.queries 而不是 request.query
|
||
logger.info(f"Received search queries: {request.queries}, collection_name: {request.collection_name}")
|
||
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
|
||
# 处理文本或向量输入
|
||
if request.queries:
|
||
query_texts = request.queries
|
||
try:
|
||
result = OpenaiAPI.batch_embeddings(query_texts)
|
||
query_vector = [item['embedding'] for item in result]
|
||
except Exception as e:
|
||
logger.error(f"Embedding failed: {e}")
|
||
raise HTTPException(status_code=500, detail="Failed to generate embeddings")
|
||
elif request.query_vectors:
|
||
query_texts = request.queries # 可能为空
|
||
query_vector = request.query_vectors
|
||
else:
|
||
raise HTTPException(status_code=400, detail="You must provide either 'queries' or 'query_vectors'")
|
||
|
||
# 执行混合搜索
|
||
search_results, search_resources, recall_results = milvus_db.hybrid_search_batch(
|
||
collection_name=request.collection_name,
|
||
database_name=request.database_name,
|
||
query_vectors=query_vector,
|
||
top_k=request.top_k,
|
||
queries=query_texts, # 可为 None
|
||
embedding_rate=0.1
|
||
)
|
||
logger.info(f"3D 模型混合检索结果: {recall_results}")
|
||
|
||
return {"code": 200, "message": "success", "data": recall_results}
|
||
|
||
class Search3DtextRequest(BaseModel):
|
||
content : str
|
||
search3d_result : list
|
||
top_k: Optional[int] = 5
|
||
database_name: Optional[str] = "XIAN"
|
||
collection_name: Optional[str] = "b9d47d5a6d66cbe2cd44a8ea4ad021fc"
|
||
|
||
|
||
@app.post("/search_3Dtext")
|
||
async def search_source(request: Search3DtextRequest):
|
||
content = request.content
|
||
search3d_result = request.search3d_result
|
||
model = HanLPmodel
|
||
allowed_resources = list(set([item['resource'] for item in search3d_result]))
|
||
SCORE_THRESHOLD = 0.3
|
||
logger.info(f"Received search content: {content}, collection_name: {request.collection_name}")
|
||
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
# 处理文本或向量输入
|
||
if content:
|
||
query_texts = extract_unique_nouns_from_file(content,model=model)
|
||
try:
|
||
result = OpenaiAPI.batch_embeddings(query_texts)
|
||
query_vector = [item['embedding'] for item in result]
|
||
except Exception as e:
|
||
logger.error(f"Embedding failed: {e}")
|
||
raise HTTPException(status_code=500, detail="Failed to generate embeddings")
|
||
else:
|
||
raise HTTPException(status_code=400, detail="You must provide either 'content' ")
|
||
|
||
# 执行混合搜索
|
||
search_results, search_resources, recall_results = milvus_db.hybrid_search_batch(
|
||
collection_name=request.collection_name,
|
||
database_name=request.database_name,
|
||
query_vectors=query_vector,
|
||
top_k=request.top_k,
|
||
queries=query_texts, # 可为 None
|
||
embedding_rate=0.1
|
||
)
|
||
search_result = []
|
||
for query_results in recall_results:
|
||
filtered_query_results = [
|
||
item for item in query_results
|
||
if item["score"] >= SCORE_THRESHOLD and item["resource"] in allowed_resources
|
||
]
|
||
search_result.append(filtered_query_results)
|
||
search_result = [sublist for sublist in search_result if sublist]
|
||
# for item in search_result[0]:
|
||
# if item["score"] >= 0.9:
|
||
# continue
|
||
# else:
|
||
# item["score"] += 0.1
|
||
if search_result:
|
||
for item in search_result[0]:
|
||
if item["score"] < 0.9:
|
||
item["score"] += 0.1
|
||
|
||
logger.info(f"大模型生成内容零部件3D检索召回:{search_result}")
|
||
if search_result:
|
||
return {"code": 200, "message": "success", "data": search_result}
|
||
else:
|
||
return {"code": 200, "message": "success", "data": []}
|
||
|
||
@app.post("/search_3D")
|
||
async def search_3D(request: SearchRequest):
|
||
"""
|
||
搜索与给定查询最相关的 Markdown 文档
|
||
"""
|
||
logger.info(f"Received search query: {request.query},collection_name: {request.collection_name}")
|
||
# 初始化 Milvus 数据库
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
|
||
# 获取查询的向量表示
|
||
# query = query_rewrite(request.query)
|
||
if request.query:
|
||
query = request.query
|
||
# query_vector = OpenaiAPI.get_embeddings(list(query))
|
||
# else:
|
||
# query_vector = request.query_vector
|
||
if request.query_vector:
|
||
query_vector = request.query_vector
|
||
else:
|
||
query_vector = OpenaiAPI.get_embeddings(list(query))
|
||
|
||
# 检查数据库是否存在
|
||
# exist_database = milvus_db.list_databases()
|
||
# if request.database_name not in exist_database:
|
||
# logger.error(f"database {request.database_name} not exist, please create it first")
|
||
# raise HTTPException(400, f" {request.database_name} not exist, please create it first")
|
||
|
||
# 使用指定数据库
|
||
# milvus_db.using_database(database_name=request.database_name)
|
||
# collection_name = sanitize_collection_name(request.collection_name)
|
||
collection_name = request.collection_name
|
||
search_results, search_resources, recall_results = milvus_db.hybrid_search(
|
||
collection_name=collection_name,
|
||
database_name=request.database_name,
|
||
query_vector=query_vector,
|
||
top_k=request.top_k,
|
||
query=request.query,
|
||
embedding_rate=0.1
|
||
)
|
||
logger.info(f"3D 模型混合检索结果: {recall_results}")
|
||
|
||
threshold = 0.2
|
||
logger.info(f"阈值:{threshold}")
|
||
|
||
search_result = [item for item in recall_results if item.get('score', 0) > threshold]
|
||
logger.info(f"精排召回第一轮:{search_result}")
|
||
search_result = rerank_3d(query, search_result)
|
||
search_result = [item for item in search_result if item.get('score', 0) > threshold]
|
||
logger.info(f"精排召回第二轮:{search_result}")
|
||
if search_result:
|
||
return {"code": 200, "message": "success", "data": search_result}
|
||
else:
|
||
return {"code": 200, "message": "success", "data": []}
|
||
|
||
|
||
class SearchChunkRequest(BaseModel):
|
||
resource: str
|
||
database_name: str
|
||
collection_name: str
|
||
|
||
|
||
@app.post("/search_chunk")
|
||
async def search_chunk(request: SearchChunkRequest):
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
results = milvus_db.get_chunk_by_resource(collection_name=request.collection_name,
|
||
database_name=request.database_name, resource=request.resource)
|
||
return {"code": 200, "message": "success", "data": results}
|
||
|
||
|
||
@app.post("/search")
|
||
async def search_markdown(request: Request, query: str, top_k: Optional[int] = 5,
|
||
database_name: Optional[str] = "JXTest", collection_name: Optional[str] = "RAGTest"):
|
||
"""
|
||
搜索与给定查询最相关的 Markdown 文档
|
||
"""
|
||
logger.info(f"Received search query: {query}")
|
||
# 在这里实现搜索逻辑,例如使用 Milvus 进行向量搜索
|
||
milvus_db = Milvus_Database(user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"), uri=os.getenv("Milvus_URI"))
|
||
query_vector = OpenaiAPI.get_embeddings(list(query))
|
||
exist_database = milvus_db.list_databases()
|
||
if database_name not in exist_database:
|
||
logger.error(f"database {database_name} not exist, please create it first")
|
||
raise HTTPException(400, f" {database_name} not exist, please create it first")
|
||
# 混搜部分
|
||
search_results = milvus_db.hybrid_search(query_vector=query_vector, query=query, top_k=top_k,
|
||
collection_name=collection_name, database_name=database_name)
|
||
reranker_result_list, relevance_scores = OpenaiAPI.rerank_query(query=query, documents=search_results, top_n=top_k)
|
||
rag_prompt = prepare_rag_prompt(query=query, documents=reranker_result_list)
|
||
return StreamingResponse(OpenaiAPI.open_api_chat_stream(query=rag_prompt), media_type="text/plain")
|
||
|
||
|
||
@app.post("/search_context_resource")
|
||
def search_context_resource(request: SearchRequest):
|
||
logger.info(f"Received search query: {request.query},collection_name: {request.collection_name}")
|
||
|
||
# 初始化 Milvus 数据库
|
||
milvus_db = Milvus_Database(
|
||
user=os.getenv("Milvus_USER", "user"),
|
||
password=os.getenv("Milvus_PASSWORD", "password"),
|
||
uri=os.getenv("Milvus_URI")
|
||
)
|
||
|
||
# 获取查询的向量表示
|
||
# query = query_rewrite(request.query)
|
||
query = request.query
|
||
query_vector = OpenaiAPI.get_embeddings(list(query))
|
||
|
||
# 检查数据库是否存在
|
||
# exist_database = milvus_db.list_databases()
|
||
# if request.database_name not in exist_database:
|
||
# logger.error(f"database {request.database_name} not exist, please create it first")
|
||
# raise HTTPException(400, f" {request.database_name} not exist, please create it first")
|
||
|
||
# 使用指定数据库
|
||
# milvus_db.using_database(database_name=request.database_name)
|
||
# collection_name = sanitize_collection_name(request.collection_name)
|
||
collection_name = request.collection_name
|
||
final_results = []
|
||
# 搜索相关文档
|
||
search_results, search_resources, recall_results = milvus_db.hierarchical_search(
|
||
collection_name=collection_name,
|
||
database_name=request.database_name,
|
||
query_vector=query_vector,
|
||
top_k=request.top_k,
|
||
query=query
|
||
)
|
||
logger.info(f"文本混合检索结果: {recall_results}")
|
||
|
||
if not recall_results:
|
||
search_results, search_resources, recall_results = milvus_db.embedding_search(
|
||
collection_name=collection_name,
|
||
query_vector=query_vector,
|
||
database_name=request.database_name,
|
||
top_k=request.top_k
|
||
)
|
||
logger.info(f"文本纯向量检索结果: {search_results}")
|
||
|
||
logger.info(f"recall_results: {recall_results}")
|
||
return {"code": 200, "message": "success", "data": recall_results, "query_vector": query_vector}
|
||
|
||
|
||
class HightLightRequest(BaseModel):
|
||
markdown_content: str
|
||
chunk: str
|
||
|
||
|
||
@app.post("/get_highlight_markdown")
|
||
async def get_hightlight(request: HightLightRequest):
|
||
logger.info(f"Received chunk :{request.chunk}")
|
||
|
||
def find_best_match(text, chunk, threshold=10):
|
||
"""返回得分最高的单个匹配位置及得分"""
|
||
normalized_text = text
|
||
normalized_chunk = chunk
|
||
|
||
# 使用滑动窗口确保找到全局最佳匹配
|
||
best_match = None
|
||
best_score = 0
|
||
|
||
# 窗口大小设为待匹配文本长度的2倍
|
||
window_size = min(len(normalized_chunk), len(normalized_text))
|
||
step = len(normalized_chunk) # 滑动步长
|
||
|
||
for i in range(1, len(normalized_text) - len(normalized_chunk), step):
|
||
window = normalized_text[i:i + window_size]
|
||
|
||
# 使用多种匹配策略组合
|
||
current_score = 0
|
||
current_score += fuzz.ratio(normalized_chunk, window) * 0.3
|
||
current_score += fuzz.partial_ratio(normalized_chunk, window) * 0.4
|
||
current_score += fuzz.token_sort_ratio(normalized_chunk, window) * 0.3
|
||
|
||
if current_score > best_score:
|
||
best_score = current_score
|
||
# 记录在原始文本中的位置(而非预处理后的位置)
|
||
raw_start = text.find(window[i:i + len(normalized_chunk)])
|
||
raw_end = raw_start + len(normalized_chunk)
|
||
best_match = (raw_start, raw_end, best_score)
|
||
|
||
# 仅返回超过阈值的匹配
|
||
return best_match if best_score >= threshold else None
|
||
|
||
def make_flexible_pattern(chunk):
|
||
"""生成允许换行符和空格数量不一致的正则模式"""
|
||
# 转义特殊字符(保留 \n 和空格)
|
||
escaped = re.escape(chunk)
|
||
|
||
# 将 \n 替换为 \s*(匹配任意空白,包括 \n 和空格)
|
||
flexible_pattern = escaped.replace(r"\n", r"\s+") # \s+ 确保至少一个空白
|
||
|
||
# 进一步优化:如果原文本有连续空格,允许匹配变体
|
||
flexible_pattern = re.sub(r"\\ ", r"\\s+", flexible_pattern)
|
||
|
||
# 完整模式:必须完整匹配整个chunk(避免部分匹配)
|
||
return re.compile(rf"\b{flexible_pattern}\b", flags=re.DOTALL)
|
||
|
||
def highlight_markdown(text, chunk):
|
||
"""在文本中找到最相似的部分并高亮"""
|
||
logger.info(f"chunk length: {len(chunk)}")
|
||
# 把换行符变成可选匹配(\s* 匹配任意空白)
|
||
# pattern = re.compile(re.escape(clean_chunk))
|
||
pattern = make_flexible_pattern(chunk)
|
||
match = pattern.search(text)
|
||
|
||
if match:
|
||
matched_text = match.group()
|
||
return text.replace(matched_text, f'<span style="background-color: #fff3cd;">{matched_text}</span>')
|
||
else:
|
||
logger.info("无法匹配,尝试调整换行符...")
|
||
chunk_0 = chunk.split(";")[0]
|
||
chunk_1 = chunk.split("\n")[0]
|
||
if chunk_0 in text:
|
||
start = text.find(chunk_0)
|
||
return text[
|
||
:start] + f'<span style="background-color: #fff3cd;">{text[start:start + len(chunk)]}</span>' + text[
|
||
start + len(
|
||
chunk):]
|
||
elif chunk_1 in text:
|
||
start = text.find(chunk_1)
|
||
return text[
|
||
:start] + f'<span style="background-color: #fff3cd;">{text[start:start + len(chunk)]}</span>' + text[
|
||
start + len(
|
||
chunk):]
|
||
|
||
first_line = chunk.split('\n')[0].split(";")[0] # 获取首行内容
|
||
if first_line and first_line in text:
|
||
# 找到首行位置
|
||
start = text.find(first_line)
|
||
# 找到chunk在原文中的结束位置(可能跨越多行)
|
||
end = start
|
||
for line in chunk.split('\n'):
|
||
line_end = text.find(line, end)
|
||
if line_end == -1:
|
||
break
|
||
end = line_end + len(line)
|
||
|
||
# 确保至少匹配了首行
|
||
if end > start:
|
||
return text[:start] + f'<span style="background-color: #fff3cd;">{text[start:end]}</span>' + text[
|
||
end:]
|
||
else:
|
||
logger.warning("仍然无法匹配!,采用模糊匹配")
|
||
result = find_best_match(text, chunk)
|
||
if not result:
|
||
return text
|
||
|
||
start, end, score = result
|
||
logger.info(f"最佳匹配得分: {score}")
|
||
|
||
# 确保不重复高亮(处理嵌套匹配)
|
||
if '<span style=' not in text[start:end]:
|
||
return text[:start] + \
|
||
f'<span style="background-color: #fff3cd;">{text[start:end]}</span>' + \
|
||
text[end:]
|
||
return text
|
||
|
||
highlighted = highlight_markdown(
|
||
# markdown_content,
|
||
request.markdown_content,
|
||
# all_content,
|
||
request.chunk
|
||
)
|
||
# logger.info(f"hightlighted text :{hightlighted_text}")
|
||
return {"code": 200, "message": "success", "data": [{"highlighted_markdown": highlighted}]}
|
||
|
||
|
||
import shutil
|
||
from fastapi.staticfiles import StaticFiles
|
||
app.mount("/converted_pdfs", StaticFiles(directory="converted_pdfs"), name="converted_pdfs")
|
||
@app.post("/word_to_pdf")
|
||
async def get_markdown(file: UploadFile = File(...)):
|
||
filename = file.filename
|
||
file_ext = os.path.splitext(filename)[1].lower()
|
||
file_name = os.path.splitext(filename)[0].lower()
|
||
pdf_url = None
|
||
|
||
if file_ext == ".docx" or file_ext == ".doc":
|
||
converted_path = await convert_docx_to_pdf(
|
||
os.getenv("CONVERT_DOC_URL", "http://192.168.0.46:59070"), file
|
||
)
|
||
|
||
try:
|
||
output_dir = "converted_pdfs"
|
||
os.makedirs(output_dir, exist_ok=True)
|
||
|
||
pdf_filename = f"{file_name}_{uuid.uuid4().hex[:8]}.pdf"
|
||
persistent_pdf_path = os.path.join(output_dir, pdf_filename)
|
||
shutil.copy2(converted_path, persistent_pdf_path)
|
||
|
||
# 构造可访问的 HTTP URL
|
||
base_url = os.getenv("BASE_URL", "http://192.168.0.46:59079")
|
||
pdf_url = f"{base_url}/converted_pdfs/{pdf_filename}"
|
||
logger.info(f"DOCX converted to PDF, accessible at: {pdf_url}")
|
||
|
||
except Exception as e:
|
||
logger.error(f"Failed to save PDF file: {e}")
|
||
raise HTTPException(status_code=500, detail="Failed to process and save PDF.")
|
||
|
||
else:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="Unsupported file type. Only DOC and DOCX are supported."
|
||
)
|
||
|
||
return {
|
||
"message": "success",
|
||
"code": 200,
|
||
"data": {
|
||
"pdf_url": pdf_url, # 返回完整可下载链接
|
||
}
|
||
}
|
||
|
||
|
||
class hightlightboxRequest(BaseModel):
|
||
split_chunk: List[dict] = []
|
||
content_middle: List[dict] = []
|
||
|
||
|
||
@app.post("/get_highlight_bbox")
|
||
async def get_hightlight(request: hightlightboxRequest):
|
||
split_chunk = request.split_chunk
|
||
content_middle = request.content_middle
|
||
|
||
|
||
def sort_by_page_idx_frequency(data):
|
||
|
||
page_idx_counts = Counter(item['page_idx'] for item in data)
|
||
|
||
# 获取出现频率最高的前三个 page_idx(如果有并列,取前三个)
|
||
top_three_page_idxs = {page_idx for page_idx, count in page_idx_counts.most_common(3)}
|
||
|
||
# 筛选出所有属于前三高频 page_idx 的原始记录
|
||
filtered_data = [item for item in data if item['page_idx'] in top_three_page_idxs]
|
||
|
||
# 在这些高频记录中找出最大的 page_idx
|
||
max_page_idx = max(top_three_page_idxs) # 因为 top_three_page_idxs 是前三高频的键
|
||
|
||
return filtered_data, max_page_idx
|
||
|
||
def find_and_filter_records(record_list, text):
|
||
# 尝试找到第一个“content”作为text前缀的记录的索引
|
||
start_index = next((i for i, record in enumerate(record_list) if
|
||
record.get('content', '').strip() and text.startswith(record['content'])), None)
|
||
# 如果找到了这样的记录,返回从这条记录开始到最后的所有记录
|
||
if start_index is not None:
|
||
record_list = record_list[start_index:]
|
||
|
||
# 如果没有找到以“content”开头的记录,则尝试找到第一个“content”作为text后缀的记录的索引
|
||
end_index = next((i for i, record in enumerate(record_list) if
|
||
record.get('content', '').strip() and text.endswith(record['content'])), None)
|
||
#
|
||
# # 如果找到了这样的记录,返回从这条记录开始到最后的所有记录
|
||
if end_index is not None:
|
||
record_list = record_list[:end_index+1]
|
||
|
||
# 如果既没有找到匹配的前缀也没有找到匹配的后缀,则返回原list
|
||
return record_list
|
||
|
||
|
||
|
||
def getbbox(result_devide, content_middle):
|
||
for idx, item in enumerate(content_middle):
|
||
item['id'] = idx + 1 # 从1开始编号
|
||
|
||
# for item in content_middle:
|
||
# print(item)
|
||
resource = result_devide[0]['resource']
|
||
for re in result_devide:
|
||
if 'chunklist' not in re:
|
||
re['chunklist'] = []
|
||
temp_chunk_list = []
|
||
|
||
for record in content_middle:
|
||
if record['type'] == 'text':
|
||
if record['content'] in re['text']:
|
||
temp_chunk_list.append(record)
|
||
elif record['type'] == 'image':
|
||
if record['image_path'] in re['text']:
|
||
temp_chunk_list.append(record)
|
||
elif record['type'] == 'table':
|
||
if record['html'] in re['text']:
|
||
temp_chunk_list.append(record)
|
||
re['chunklist'] = temp_chunk_list
|
||
result_devide = [item for item in result_devide if item['chunklist']]
|
||
|
||
length_result_devide = len(result_devide)
|
||
result = {}
|
||
|
||
for i in range(length_result_devide):
|
||
temp_uuid = str(uuid.uuid4())
|
||
result[temp_uuid] = []
|
||
chunk = result_devide[i]['text']
|
||
result_devide[i]["id"] = temp_uuid
|
||
|
||
# 假设输入数据存储在名为data的列表中
|
||
data = result_devide[i]['chunklist'] # 您的数据列表
|
||
|
||
sorted_data = sorted(data, key=lambda x: x['id'])
|
||
|
||
# 步骤2:使用 groupby 找出连续 id 的组
|
||
groups = []
|
||
current_group = []
|
||
|
||
for item in sorted_data:
|
||
if not current_group:
|
||
current_group.append(item)
|
||
else:
|
||
# 检查当前 id 是否与前一个连续
|
||
if item['id'] == current_group[-1]['id'] + 1:
|
||
current_group.append(item)
|
||
else:
|
||
# 不连续,保存当前组并开始新组
|
||
if len(current_group) >= 1: # 可设最小长度过滤
|
||
groups.append(current_group)
|
||
current_group = [item]
|
||
|
||
# 别忘了最后一个组
|
||
if current_group:
|
||
groups.append(current_group)
|
||
|
||
# 按组的长度从大到小排序
|
||
sorted_groups = sorted(groups, key=len, reverse=True)
|
||
|
||
sorted_data, most_common_page_idx1 = sort_by_page_idx_frequency(sorted_groups[0])
|
||
sorted_filtered_items = sorted(
|
||
(item for sublist in sorted_groups for item in sublist if
|
||
most_common_page_idx1 - 1 <= item['page_idx'] <= most_common_page_idx1 + 1),
|
||
key=lambda x: x['id']
|
||
)
|
||
sorted_filtered_items = find_and_filter_records(sorted_filtered_items, chunk)
|
||
if not sorted_filtered_items :
|
||
most_common_page_idx = 0
|
||
else:
|
||
most_common_page_idx = sorted_filtered_items[0]['page_idx']
|
||
|
||
for chunk in sorted_filtered_items:
|
||
chunk['resource'] = resource
|
||
result[temp_uuid].append(chunk)
|
||
result_devide[i]['page_idx'] = most_common_page_idx
|
||
# result[temp_uuid] =sorted_data
|
||
#
|
||
data = []
|
||
for ins in result_devide:
|
||
data.append(
|
||
{
|
||
'id': ins['id'],
|
||
'Header_1': ins['Header_1'],
|
||
'Header_2': ins['Header_2'],
|
||
'Header_3': ins['Header_3'],
|
||
'text': ins['text'],
|
||
'resource': resource,
|
||
'img_path': ins['img_path'],
|
||
'origin_text': ins['origin_text'],
|
||
'page': None,
|
||
'page_idx': ins['page_idx']
|
||
}
|
||
|
||
)
|
||
return data, result
|
||
|
||
data, result = getbbox(split_chunk, content_middle)
|
||
if not data or not result:
|
||
return {"code": 200, "message": "success", "data": []}
|
||
else:
|
||
return {"code": 200, "message": "success", "data": [{"data": data, "result": result}]}
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import uvicorn
|
||
|
||
uvicorn.run(app, host="0.0.0.0", port=9079)
|