ForcePilot/src/core/retriever.py

269 lines
10 KiB
Python
Raw Normal View History

from src import config, knowledge_base, graph_base
2025-02-23 16:38:56 +08:00
from src.models.rerank_model import get_reranker
2025-02-27 19:35:25 +08:00
from src.utils.logging_config import logger
2025-03-16 01:11:13 +08:00
from src.models import select_model
2024-07-16 23:12:35 +08:00
2024-07-10 12:56:53 +08:00
class Retriever:
2025-03-16 01:11:13 +08:00
def __init__(self):
self._load_models()
2024-07-10 12:56:53 +08:00
2025-03-16 01:11:13 +08:00
def _load_models(self):
if config.enable_reranker:
2025-02-23 16:38:56 +08:00
self.reranker = get_reranker(config)
2024-07-29 01:00:02 +08:00
2025-03-16 01:11:13 +08:00
if config.enable_web_search:
from src.utils.web_search import WebSearcher
self.web_searcher = WebSearcher()
2024-07-17 18:52:20 +08:00
def retrieval(self, query, history, meta):
2024-08-08 19:51:13 +08:00
refs = {"query": query, "history": history, "meta": meta}
2025-03-16 01:11:13 +08:00
refs["model_name"] = config.model_name
2024-07-29 01:00:02 +08:00
refs["entities"] = self.reco_entities(query, history, refs)
refs["knowledge_base"] = self.query_knowledgebase(query, history, refs)
refs["graph_base"] = self.query_graph(query, history, refs)
refs["web_search"] = self.query_web(query, history, refs)
2024-07-10 12:56:53 +08:00
2024-07-17 18:52:20 +08:00
return refs
2024-07-10 12:56:53 +08:00
2025-03-16 01:11:13 +08:00
def restart(self):
"""所有需要重启的模型"""
self._load_models()
2024-07-17 18:52:20 +08:00
def construct_query(self, query, refs, meta):
logger.debug(f"{refs=}")
2024-09-18 21:03:12 +08:00
if not refs or len(refs) == 0:
2024-07-10 12:56:53 +08:00
return query
2024-09-18 21:03:12 +08:00
external_parts = []
2024-07-10 12:56:53 +08:00
2024-07-29 01:00:02 +08:00
# 解析知识库的结果
2024-09-18 21:03:12 +08:00
kb_res = refs.get("knowledge_base", {}).get("results", [])
if kb_res:
kb_text = "\n".join(f"{r['id']}: {r['entity']['text']}" for r in kb_res)
external_parts.extend(["知识库信息:", kb_text])
2024-07-10 12:56:53 +08:00
2024-07-29 01:00:02 +08:00
# 解析图数据库的结果
2024-09-18 21:03:12 +08:00
db_res = refs.get("graph_base", {}).get("results", {})
if db_res.get("nodes") and len(db_res["nodes"]) > 0:
2024-09-06 21:30:38 +08:00
db_text = "\n".join(
2024-09-18 21:03:12 +08:00
[f"{edge['source_name']}{edge['target_name']}的关系是{edge['type']}" for edge in db_res.get("edges", [])]
2024-09-06 21:30:38 +08:00
)
2024-09-18 21:03:12 +08:00
external_parts.extend(["图数据库信息:", db_text])
2024-07-17 18:52:20 +08:00
# 解析网络搜索的结果
web_res = refs.get("web_search", {}).get("results", [])
if web_res:
web_text = "\n".join(f"{r['title']}: {r['content']}" for r in web_res)
external_parts.extend(["网络搜索信息:", web_text])
2024-07-29 01:00:02 +08:00
# 构造查询
2024-09-25 13:46:23 +08:00
from src.utils.prompts import knowbase_qa_template
2024-09-18 21:03:12 +08:00
if external_parts and len(external_parts) > 0:
external = "\n\n".join(external_parts)
2024-09-25 13:46:23 +08:00
query = knowbase_qa_template.format(external=external, query=query)
2024-07-10 12:56:53 +08:00
return query
def query_classification(self, query):
"""判断是否需要查询
2025-02-28 02:41:45 +08:00
- 对于完全基于用户给定信息的任务称之为"足够""sufficient"不需要检索
- 否则称之为"不足""insufficient"可能需要检索
2024-07-10 12:56:53 +08:00
"""
raise NotImplementedError
2024-07-29 01:00:02 +08:00
def query_graph(self, query, history, refs):
2024-07-17 18:50:01 +08:00
results = []
2025-03-16 01:11:13 +08:00
if refs["meta"].get("use_graph") and config.enable_knowledge_base:
2024-07-29 01:00:02 +08:00
for entity in refs["entities"]:
if entity == "":
continue
result = graph_base.query_by_vector(entity)
2024-07-20 20:30:32 +08:00
if result != []:
results.extend(result)
2024-07-18 02:46:58 +08:00
return {"results": self.format_query_results(results)}
2024-07-16 23:12:35 +08:00
2024-07-29 01:00:02 +08:00
def query_knowledgebase(self, query, history, refs):
"""查询知识库"""
2024-07-17 18:52:20 +08:00
2025-03-23 23:10:32 +08:00
response = {
"results": [],
"all_results": [],
"rw_query": query,
"message": "",
}
2024-07-17 18:52:20 +08:00
2024-09-18 21:03:12 +08:00
meta = refs["meta"]
2024-09-09 17:07:03 +08:00
2025-03-23 23:10:32 +08:00
db_id = meta.get("db_id")
if not db_id or not config.enable_knowledge_base:
response["message"] = "知识库未启用、或未指定知识库、或知识库不存在"
return response
2024-08-25 20:29:24 +08:00
2025-03-23 23:10:32 +08:00
rw_query = self.rewrite_query(query, history, refs)
2024-09-09 17:07:03 +08:00
2025-03-23 23:10:32 +08:00
logger.debug(f"{meta=}")
query_result = knowledge_base.query(query=rw_query,
db_id=db_id,
distance_threshold=meta.get("distanceThreshold", 0.5),
rerank_threshold=meta.get("rerankThreshold", 0.1),
max_query_count=meta.get("maxQueryCount", 20),
top_k=meta.get("topK", 10))
2024-08-08 19:51:13 +08:00
2025-03-23 23:10:32 +08:00
response["results"] = query_result["results"]
response["all_results"] = query_result["all_results"]
response["rw_query"] = rw_query
2025-03-23 23:10:32 +08:00
return response
2024-07-17 18:52:20 +08:00
def query_web(self, query, history, refs):
"""查询网络"""
2025-03-16 01:11:13 +08:00
if not (refs["meta"].get("use_web") and config.enable_web_search):
return {"results": [], "message": "Web search is disabled"}
try:
search_results = self.web_searcher.search(query, max_results=5)
except Exception as e:
logger.error(f"Web search error: {str(e)}")
return {"results": [], "message": "Web search error"}
return {"results": search_results}
2024-07-29 01:00:02 +08:00
def rewrite_query(self, query, history, refs):
2024-07-10 12:56:53 +08:00
"""重写查询"""
2025-03-23 23:10:32 +08:00
model_provider = config.model_provider_lite
model_name = config.model_name_lite
2025-03-29 17:33:09 +08:00
model = select_model(model_provider=model_provider, model_name=model_name)
if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置
rewrite_query_span = refs["meta"].get("use_rewrite_query", "off")
else:
2025-03-16 01:11:13 +08:00
rewrite_query_span = config.use_rewrite_query
2024-09-09 17:07:03 +08:00
if rewrite_query_span == "off":
2024-07-17 18:50:01 +08:00
rewritten_query = query
else:
2024-07-29 01:00:02 +08:00
from src.utils.prompts import rewritten_query_prompt_template
2024-09-06 21:30:38 +08:00
history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else ""
2024-08-08 19:51:13 +08:00
rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query)
2025-03-16 01:11:13 +08:00
rewritten_query = model.predict(rewritten_query_prompt).content
2024-07-17 18:50:01 +08:00
2024-09-09 17:07:03 +08:00
if rewrite_query_span == "hyde":
2025-03-16 01:11:13 +08:00
hy_doc = model.predict(rewritten_query).content
2024-08-08 19:51:13 +08:00
rewritten_query = f"{rewritten_query} {hy_doc}"
2024-07-29 01:00:02 +08:00
return rewritten_query
2024-07-21 18:15:28 +08:00
2024-07-29 01:00:02 +08:00
def reco_entities(self, query, history, refs):
"""识别句子中的实体"""
query = refs.get("rewritten_query", query)
2025-03-23 23:10:32 +08:00
model_provider = config.model_provider_lite
model_name = config.model_name_lite
2025-03-29 17:33:09 +08:00
model = select_model(model_provider=model_provider, model_name=model_name)
2024-07-29 01:00:02 +08:00
entities = []
if refs["meta"].get("use_graph"):
2024-09-18 21:03:12 +08:00
from src.utils.prompts import entity_extraction_prompt_template as entity_template
from src.utils.prompts import keywords_prompt_template as entity_template
2024-09-06 21:30:38 +08:00
2024-09-18 21:03:12 +08:00
entity_extraction_prompt = entity_template.format(text=query)
2025-03-16 01:11:13 +08:00
entities = model.predict(entity_extraction_prompt).content.split("<->")
2024-09-18 21:03:12 +08:00
# entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)]
2024-07-29 01:00:02 +08:00
return entities
2024-07-10 12:56:53 +08:00
2025-02-28 02:42:21 +08:00
def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None):
2024-09-06 21:30:38 +08:00
"""
提取关系信息并返回格式化的节点和边信息
"""
rel_id = relationship.element_id
nodes = relationship.nodes
if len(nodes) != 2:
return None, None
source, target = nodes
source_id = source.element_id
target_id = target.element_id
2025-02-28 02:42:21 +08:00
source_name = node_dict[source_id]["name"] if source_name is None else source_name
target_name = node_dict[target_id]["name"] if target_name is None else target_name
2024-09-06 21:30:38 +08:00
relationship_type = relationship._properties.get("type", "unknown")
if relationship_type == "unknown":
relationship_type = relationship.type
edge_info = {
"id": rel_id,
"type": relationship_type,
"source_id": source_id,
"target_id": target_id,
"source_name": source_name,
"target_name": target_name,
}
node_info = [
{"id": source_id, "name": source_name},
{"id": target_id, "name": target_name},
]
return node_info, edge_info
def format_general_results(self, results):
2024-09-06 12:54:17 +08:00
formatted_results = {"nodes": [], "edges": []}
for item in results:
relationship = item[1]
2024-09-06 21:30:38 +08:00
source_name = item[0]._properties.get("name", "unknown")
target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "unknown"
2024-09-06 12:54:17 +08:00
2024-09-06 21:30:38 +08:00
node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name)
if node_info is None or edge_info is None:
continue
2024-09-06 12:54:17 +08:00
2024-09-06 21:30:38 +08:00
for node in node_info:
if node["id"] not in [n["id"] for n in formatted_results["nodes"]]:
formatted_results["nodes"].append(node)
2024-09-06 12:54:17 +08:00
2024-09-06 21:30:38 +08:00
formatted_results["edges"].append(edge_info)
2024-09-06 12:54:17 +08:00
return formatted_results
2024-07-29 01:00:02 +08:00
def format_query_results(self, results):
# logger.debug(f"Graph Query Results: {results}")
2024-07-18 02:46:58 +08:00
formatted_results = {"nodes": [], "edges": []}
2024-07-20 20:30:32 +08:00
node_dict = {}
2024-07-20 20:30:32 +08:00
for item in results:
2025-02-28 02:42:21 +08:00
# 检查数据格式
if len(item) < 2 or not isinstance(item[1], list):
2024-07-29 01:00:02 +08:00
continue
2025-02-28 02:42:21 +08:00
node_dict[item[0].element_id] = dict(id=item[0].element_id, name=item[0]._properties.get("name", "Unknown"))
node_dict[item[2].element_id] = dict(id=item[2].element_id, name=item[2]._properties.get("name", "Unknown"))
2024-09-06 21:30:38 +08:00
2025-02-28 02:42:21 +08:00
# 处理关系列表中的每个关系
for i, relationship in enumerate(item[1]):
try:
# 提取关系信息
node_info, edge_info = self._extract_relationship_info(relationship, node_dict=node_dict)
if node_info is None or edge_info is None:
continue
2024-07-29 01:00:02 +08:00
2025-02-28 02:42:21 +08:00
# 添加边
formatted_results["edges"].append(edge_info)
except Exception as e:
logger.error(f"处理关系时出错: {e}, 关系: {relationship}")
continue
2025-02-28 02:42:21 +08:00
# 将节点字典转换为列表
2024-07-20 20:30:32 +08:00
formatted_results["nodes"] = list(node_dict.values())
2024-07-18 02:46:58 +08:00
return formatted_results
2024-07-17 18:52:20 +08:00
def __call__(self, query, history, meta):
refs = self.retrieval(query, history, meta)
query = self.construct_query(query, refs, meta)
2024-09-06 21:30:38 +08:00
return query, refs