import httpx
from openai import OpenAI
import os
import requests
import json
import asyncio
import aiohttp
from typing import List, Dict, Any
from tqdm import tqdm # 进度条支持(可选)
import re
class OpenaiAPI:
# def __init__(self):
# self.openai_embedding_api_base = os.getenv("OPENAI_API_EMBEDDING_BASE")
# self.openai_embedding_model = os.getenv("OPENAI_EMBEDDING_MODEL")
# self.openai_chat_api_base = os.getenv("OPENAI_API_BASE")
# self.openai_chat_model = os.getenv("OPENAI_MODEL")
# self.openai_api_key = os.getenv("OPENAI_API_KEY")
# self.openai_reranker_api_base = os.getenv("OPENAI_API_RERANKER_BASE")
# self.openai_reranker_model = os.getenv("OPENAI_RERANKER_MODEL")
# self.embedding_client = OpenAI(
# api_key=self.openai_api_key,
# base_url=self.openai_embedding_api_base,
# )
# self.chat_client = OpenAI(
# api_key=self.openai_api_key,
# base_url=self.openai_chat_api_base,
# )
# self.rerank_client = OpenAI(
# api_key=self.openai_api_key,
# base_url=self.openai_reranker_api_base,
# )
@staticmethod
def get_embeddings(embedding_texts: List[str]):
base_url = os.getenv("OPENAI_API_EMBEDDING_BASE")
embedding_base_url = os.path.join(base_url, "embeddings")
with httpx.Client(timeout=60) as client:
response = client.post(
embedding_base_url,
json={"model": "bge-m3", "input": embedding_texts},
headers={"Authorization": os.getenv("OPENAI_API_KEY")},
)
embedding_list = response.json()["data"][0]["embedding"]
# response = OpenAI(
# api_key=os.getenv("OPENAI_API_KEY","EMPTY"),
# base_url=os.getenv("OPENAI_API_EMBEDDING_BASE","http://172.18.30.122:40052/v1"),
# ).embeddings.create(
# model=os.getenv("OPENAI_EMBEDDING_MODEL","bge-m3"),
# input=embedding_texts,
# )
# embedding_list = response.data[0].embedding
return embedding_list
@staticmethod
def batch_embeddings(embedding_texts: List[str]):
base_url = os.getenv("OPENAI_API_EMBEDDING_BASE")
embedding_base_url = os.path.join(base_url, "embeddings")
with httpx.Client(timeout=60) as client:
response = client.post(
embedding_base_url,
json={"model": "bge-m3", "input": embedding_texts},
headers={"Authorization": os.getenv("OPENAI_API_KEY")},
)
embedding_list = response.json()["data"]
return embedding_list
@staticmethod
def rerank_query(query: str, documents: List[str], top_n: str = 3):
"""
发送重排序请求到 GPUStack API。
:param query: 查询字符串
:param documents: 文档列表(字符串列表)
:param top_n: 返回的文档数量
:param model: 使用的模型名称
:param api_key: API 密钥
:param server_url: API 服务器地址
:return: API 的响应结果(JSON 格式)
return :
{
"id": "rerank-1446b27f0d8642288578ee26707253a1",
"model": "bge-reranker-v2-m3",
"usage": {
"total_tokens": 53
},
"results": [
{
"index": 2,
"document": {
"text": "I have a dog named Dollor"
},
"relevance_score": 0.297119140625
},
{
"index": 1,
"document": {
"text": "hi"
},
"relevance_score": 0.0009255409240722656
},
{
"index": 0,
"document": {
"text": "what is panda?"
},
"relevance_score": 0.0008897781372070312
}
]
}
"""
# 请求头
headers = {
'Content-Type': 'application/json',
'Authorization': f'Bearer {os.getenv("OPENAI_API_KEY")}'
}
# 请求体
data = {
"model": os.getenv("OPENAI_RERANKER_MODEL"),
"query": query,
"top_n": top_n,
"documents": documents
}
# 发送 POST 请求
response = requests.post(os.getenv("OPENAI_API_RERANKER_BASE") + "/rerank", headers=headers,
data=json.dumps(data))
results_list = []
# 检查响应状态码
if response.status_code == 200:
result_json = response.json() # 返回解析后的 JSON 数据
for result in result_json["results"]:
results_list.append(
{"index": result["index"], "text": result["document"]["text"], "scores": result["relevance_score"]})
# results_list.append(result["document"]["text"])
# relevance_scores.append(result["relevance_score"])
return results_list
else:
raise Exception(f"Error: {response.status_code}\n{response.text}") # 抛出异常
@staticmethod
def embedding_cosine_similarity(text_1, text_2):
# 请求头
headers = {
'Content-Type': 'application/json',
'Authorization': f'Bearer {os.getenv("OPENAI_API_KEY")}'
}
# 请求体
data = {
"text_1": text_1,
"text_2": text_2,
"model": "bge-rerank"
}
# # 发送 POST 请求
# response = requests.post(os.getenv("OPENAI_API_EMBEDDING_BASE") + "/score", headers=headers, data=json.dumps(data))
response = requests.post(os.getenv("OPENAI_API_RERANKER_BASE") + "/score", headers=headers,
data=json.dumps(data))
if response.status_code == 200:
result_json = response.json()
score = result_json["data"][0]["score"]
return score
else:
raise Exception(f"Error: {response.status_code}\n{response.text}") # 抛出异常
@staticmethod
def open_api_chat(query: str, model: str):
response = OpenAI(
api_key=os.getenv("OPENAI_API_KEY"),
base_url=os.getenv("OPENAI_API_BASE"),
).chat.completions.create(
model=model,
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": query},
],
temperature=0.6,
stream=False
)
return response.choices[0].message.content
@staticmethod
def open_api_caht_without_thinking(query: str, model: str):
response = OpenAI(
api_key=os.getenv("OPENAI_API_KEY"),
base_url=os.getenv("OPENAI_API_BASE"),
).chat.completions.create(
model=model,
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": query}, ],
temperature=0.1,
stream=False,
extra_body={"chat_template_kwargs": {"enable_thinking": False}})
return response.choices[0].message.content
@staticmethod
def open_api_chat_stream(query: str, model: str):
response = OpenAI(
api_key=os.getenv("OPENAI_API_KEY"),
base_url=os.getenv("OPENAI_API_BASE"),
).chat.completions.create(
model=model,
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": query},
],
temperature=0.6,
stream=True
)
for chunk in response:
content = chunk.choices[0].delta.content
if content:
yield content.encode('utf-8')
@staticmethod
def context_chunk_process(markdown_context: str, chunks: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""同步调用入口"""
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
result = loop.run_until_complete(
batch_chunk_process(markdown_context, chunks)
)
finally:
loop.close()
return result
async def batch_chunk_process(
markdown_context: str,
chunks: List[Dict[str, Any]],
batch_size: int = 5,
timeout: int = 30
) -> List[Dict[str, Any]]:
"""
批量处理文档块(异步并发)
Args:
markdown_context: 完整文档内容
chunks: 文档块列表,每个块需包含"text"键
batch_size: 并发请求数 (默认5)
timeout: 单个请求超时时间(秒)
Returns:
处理后的文档块列表
"""
# 初始化参数
api_key = os.getenv("OPENAI_API_KEY")
base_url = os.getenv("OPENAI_API_BASE")
model = os.getenv("OPENAI_MODEL", "qwen3-32b")
if not api_key or not base_url:
raise ValueError("Missing OpenAI API configuration")
async def process_single_chunk(
session: aiohttp.ClientSession,
chunk: Dict[str, Any],
semaphore: asyncio.Semaphore
) -> Dict[str, Any]:
"""处理单个chunk的异步函数"""
async with semaphore:
text = chunk.get("text", "").strip()
if not text:
return chunk
prompt = f"""
{markdown_context}
{text}
Please provide a succinct context for this chunk within the document.
Respond ONLY with the context text, no additional commentary.
"""
chinese_prompt = f"""
<文档>
{markdown_context}
文档>
以下是需要定位到整体文档中的片段
<片段>
{text}
片段>
请简要说明该片段在整体文档中的上下文关系,以提升片段搜索效果。只需给出简洁的上下文说明,无需其他内容。
"""
try:
async with session.post(
f"{base_url}/chat/completions",
headers={"Authorization": f"Bearer {api_key}"},
json={
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": chinese_prompt}
],
"temperature": 0.1,
"extra_body": {"chat_template_kwargs": {"enable_thinking": False}}
},
timeout=timeout
) as response:
if response.status == 200:
data = await response.json()
if data.get("choices"):
chunk["text"] = data["choices"][0]["message"]["content"].strip()
else:
print(f"API error for chunk: {await response.text()}")
except Exception as e:
print(f"Error processing chunk: {e}")
return chunk
# 创建信号量控制并发数
semaphore = asyncio.Semaphore(batch_size)
# 异步处理所有chunks
async with aiohttp.ClientSession() as session:
tasks = [
process_single_chunk(session, chunk, semaphore)
for chunk in chunks
]
# 使用tqdm显示进度(可选)
processed_chunks = []
for future in tqdm(asyncio.as_completed(tasks), total=len(tasks)):
processed_chunks.append(await future)
return processed_chunks
def prepare_rag_prompt(query, documents):
docs = ""
for i, doc in enumerate(documents):
docs += f"文本信息{i + 1}:{doc}\n"
return f"""你作为一个文档检索问答专家,请你从以下文本信息中选取有用的内容作为参考,仔细思考后用中文回答用户的问题,如果文本中没有关于问题的答案,则回答不知道:
-----
{docs}
-----
用户提问:{query}
你的回答:"""
def query_rewrite(query):
prompt = f"""
你是一个专业的搜索查询优化助手。你的任务是根据用户输入的原始查询,生成更清晰、更易被搜索引擎理解的改写版本。请遵循以下规则:
1. **保持原意**:改写后的查询必须与用户意图一致,不能曲解原意。
2. **简洁清晰**:去掉冗余词,优化表达,但不要过度缩写。
3. **结构化**:如果查询适合拆解,可以适当结构化(如添加隐含的限定条件)。
4. **自然语言**:确保改写后的查询仍然符合自然语言习惯。
5. **多语言支持**:如果查询是中文,优先用中文改写;如果是英文,用英文优化。
请改写以下查询,直接返回优化后的内容,不要额外解释:
【用户查询开始】
{query}
【用户查询结束】
"""
return OpenaiAPI.open_api_chat(prompt)
def RAG_PROMPT(QUERY, CONTEXT):
PROMPT = f"""
### Task:
Respond to the user query using the provided context, incorporating inline citations in the format [source_id] **only when the tag is explicitly provided** in the context.
### Guidelines:
- If you don't know the answer, clearly state that.
- If uncertain, ask the user for clarification.
- Respond in the same language as the user's query.
- If the context is unreadable or of poor quality, inform the user and provide the best possible answer.
- If the answer isn't present in the context but you possess the knowledge, explain this to the user and provide the answer using your own understanding.
- When you find corresponding content in the provided sources, try to quote it verbatim as much as possible.
- When quoting from sources, reproduce the original text verbatim—do not omit, paraphrase, or add content.
- **Only include inline citations using [source_id] (e.g., [1], [2]) when a `` tag is explicitly provided in the context.**
- If the source file has a .glb extension, directly mark the corresponding content with the source_id in the relevant position.
- Do not cite if the or tag is not provided in the context.
- Do not use XML tags in your response.
- **If none of the provided sources contain relevant information to answer the question, respond only with: "没有找到相关来源" (No relevant sources found). Do not add any other content.
- Ensure citations are concise and directly related to the information provided.
### Example of Citation:
If the user asks about a specific topic and the information is found in "whitepaper.pdf" with a provided , the response should include the citation like so:
* "According to the study, the proposed method increases efficiency by 20% [whitepaper.pdf]."
If no is present, the response should omit the citation.
### Output:
Provide a clear and direct response to the user's query, including inline citations in the format [source_id] only when the tag is present in the context.
{CONTEXT}
{QUERY}
"""
return PROMPT
# for chunk in response:
# content = chunk.choices[0].delta.content
# if content:
# print(content, end='', flush=True)
# print('\n')
# return response
def split_by_period(text):
sentences = re.split('(?<=。)', text)
sentences = [sentence.strip() for sentence in sentences if sentence.strip()]
return sentences
def rerank_text(query,data):
for datainfo in data:
splitdata = split_by_period(datainfo['text'])
current_score = datainfo['score']
text_1 = query
url = os.getenv("OPENAI_API_RERANKER_BASE") + "/score"
headers = {
'Authorization': 'Bearer 123456',
'Content-Type': 'application/json'
}
parse_data = {
"model": "bge-rerank",
"text_1": text_1,
"text_2": splitdata
}
response = requests.post(url, headers=headers, json=parse_data)
parseresult = response.json()
scores = [item['score'] for item in parseresult['data'] if 'score' in item]
max_score = max(scores)
socore = 0.1 * current_score + 0.9 * max_score
datainfo['score'] = socore
return data
def rerank_3d(query,data):
for recall_result in data:
text_2 = recall_result["text"]
current_score = recall_result["score"]
text_1 = query
url = os.getenv("OPENAI_API_RERANKER_BASE") + "/score"
headers = {
'Authorization': 'Bearer 123456',
'Content-Type': 'application/json'
}
parse_data = {
"model": "bge-rerank",
"text_1": text_1,
"text_2": text_2
}
response = requests.post(url, headers=headers, json=parse_data)
parseresult = response.json()
rerank_score = parseresult['data'][0]['score']
score = current_score * 0.2 + rerank_score * 0.8
recall_result["score"] = score
result_dict = {}
for item in data:
resource = item['resource']
# 如果该 resource 不存在,或当前 score 更高,则替换
if resource not in result_dict or item['score'] > result_dict[resource]['score']:
result_dict[resource] = item
# 转换为列表
filtered_result = list(result_dict.values())
for new_index, item in enumerate(filtered_result, start=1):
item['index'] = new_index
return filtered_result
class OPENAIAPI:
@staticmethod
def openai_chat(url,query, model):
"""调用OpenAI API"""
try:
client = OpenAI(
api_key="none",
base_url=url,
)
response = client.chat.completions.create(
model=model,
messages=[
{"role": "system",
"content": "你是一个专业的武器类查询改写助手,将用户问题改写为一个容易理解、且明确的查询。"},
{"role": "user", "content": query},
],
temperature=0.1,
stream=False,
extra_body={"chat_template_kwargs": {"enable_thinking": False}}
)
return response.choices[0].message.content
except Exception as e:
print(f"调用OpenAI API时出错: {e}")
return None
prompt = """
你是一个查询重写工具。请将以下用户查询改写为一个清晰、具体、完整、无歧义的自然语言问题。
要求:
- 保留原始意图,不添加未提及的假设;
- 若查询模糊(如缩写、代词、片段),基于常识明确实体和领域(例如“AR50”应理解为“AR-50狙击步枪”);
- 使用完整陈述句,聚焦信息需求(如定义等);
- **仅输出改写后的问题本身,不要任何其他文字、标点前缀、问候、反问或说明。**
- 不要进行任何思考,直接输出改写后查询
- 请按照下述示例进行改写
原始查询:ar50
改写后查询:AR-50狙击步枪是什么?
原始查询:变速器?
改写后查询:变速器是什么?
原始查询:{user_query}
"""
def resetQuery(query:str):
url = os.getenv("OPENAI_API_MODEL_BASE"+"/v1","http://192.168.0.46:59800/v1")
model = os.getenv("MODEL_NAME","Qwen3-32B")
query_prompt = prompt.replace("{user_query}",query)
answer = OPENAIAPI.openai_chat(url,query_prompt, model)
return answer