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'{matched_text}') 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'{text[start:start + len(chunk)]}' + text[ start + len( chunk):] elif chunk_1 in text: start = text.find(chunk_1) return text[ :start] + f'{text[start:start + len(chunk)]}' + 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'{text[start:end]}' + 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 '{text[start:end]}' + \ 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)