kgrag/app_pkg/app.py
2026-07-29 18:10:19 +08:00

1359 lines
55 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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)