diff --git a/src/config/__init__.py b/src/config/__init__.py index 2d362b53..807f3125 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -42,7 +42,7 @@ class Config(SimpleConfig): self.add_item("stream", default=True, des="是否开启流式输出") self.add_item("save_dir", default="saves", des="保存目录") # 功能选项 - self.add_item("enable_query_rewrite", default=True, des="是否开启查询重写") + self.add_item("enable_reranker", default=False, des="是否开启重排序") self.add_item("enable_knowledge_base", default=True, des="是否开启知识库") self.add_item("enable_knowledge_graph", default=True, des="是否开启知识图谱") self.add_item("enable_search_engine", default=True, des="是否开启搜索引擎") @@ -58,6 +58,7 @@ class Config(SimpleConfig): ### <<< 默认配置结束 self.filename = filename or os.path.join(self.save_dir, "config", "config.yaml") + os.makedirs(os.path.dirname(self.filename), exist_ok=True) self.load() self.handle_self() @@ -79,7 +80,6 @@ class Config(SimpleConfig): if not model_rel_path.startswith("/"): self.model_local_paths[model] = os.path.join(model_root_dir, model_rel_path) - def load(self): """根据传入的文件覆盖掉默认配置""" logger.info(f"Loading config from {self.filename}") @@ -108,6 +108,7 @@ class Config(SimpleConfig): if self.filename is None: logger.warning("Config file is not specified, save to default config/base.yaml") self.filename = os.path.join(self.save_dir, "config", "config.yaml") + os.makedirs(os.path.dirname(self.filename), exist_ok=True) if self.filename.endswith(".json"): with open(self.filename, 'w+') as f: diff --git a/src/core/database.py b/src/core/database.py index b10922ac..6b951c52 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -21,8 +21,12 @@ class DataBaseLite: self.metadata = kwargs.get("metaname", {}) self.files = kwargs.get("files", []) self.embed_model = kwargs.get("embed_model", None) - self.id2file = {f["file_id"]: f for f in self.files} + def id2file(self, file_id): + for f in self.files: + if f["file_id"] == file_id: + return f + return None def update(self, metadata): self.metadata = metadata diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 31e5b377..ca5270cf 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -277,7 +277,7 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query, keyword) - def query_by_vector(self, entity_name, num_of_res=2, threshold=0.9,kgdb_name='neo4j', hops=2): + def query_by_vector(self, entity_name, num_of_res=2, threshold=0.9, kgdb_name='neo4j', hops=2): self.use_database(kgdb_name) result = self.query_by_vector_tep(entity_name) querys = [] diff --git a/src/core/retriever.py b/src/core/retriever.py index c0750c24..50677467 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -7,18 +7,20 @@ class Retriever: def __init__(self, config, dbm, model): self.config = config - self.reranker = Reranker(config) self.dbm = dbm self.model = model + if self.config.enable_reranker: + self.reranker = Reranker(config) + def retrieval(self, query, history, meta): refs = {} - refs["meta"] = meta - refs["rewrite_query"] = self.rewrite_query(query, history, meta) - refs["knowledge_base"] = self.query_knowledgebase(query, history, meta) - refs["graph_base"] = self.query_graph(query, history, meta, entities=refs["rewrite_query"][1]) + refs["rewritten_query"] = self.rewrite_query(query, history, refs) + 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) return refs @@ -28,16 +30,19 @@ class Retriever: external = "" + # 解析知识库的结果 kb_res = refs.get("knowledge_base").get("results", []) if len(kb_res) > 0: kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res]) external += f"知识库信息: \n\n{kb_text}" + # 解析图数据库的结果 db_res = refs.get("graph_base").get("results", []) if len(db_res["nodes"]) > 0: db_text = '\n'.join([f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res['edges']]) external += f"图数据库信息: \n\n{db_text}" + # 构造查询 if len(external) > 0: query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}" @@ -50,114 +55,103 @@ class Retriever: """ raise NotImplementedError - def query_graph(self, query, history, meta, entities): + def query_graph(self, query, history, refs): # res = model.predict("qiansdgsa, dasdh ashdsakjdk ak ").content results = [] - if meta.get("use_graph"): - for entitie in entities: - result = self.dbm.graph_base.query_by_vector(entitie) + if refs["meta"].get("use_graph"): + for entity in refs["entities"]: + result = self.dbm.graph_base.query_by_vector(entity) if result != []: results.extend(result) return {"results": self.format_query_results(results)} - def query_knowledgebase(self, query, history, meta): + def query_knowledgebase(self, query, history, refs): + """查询知识库""" + query = refs.get("rewritten_query", query) kb_res = [] - if meta.get("db_name"): - kb = self.dbm.metaname2db[meta["db_name"]] - kb_res = self.dbm.knowledge_base.search(query, meta["db_name"], limit=5) + if refs["meta"].get("db_name"): + db_name = refs["meta"]["db_name"] + kb = self.dbm.metaname2db[refs["meta"]["db_name"]] + limit = refs["meta"].get("queryCount", 10) + kb_res = self.dbm.knowledge_base.search(query, db_name, limit=limit) for r in kb_res: - r["file"] = kb.id2file[r["entity"]["file_id"]] - r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) + r["file"] = kb.id2file(r["entity"]["file_id"]) - kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) - - final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1] + if self.config.enable_reranker: + for r in kb_res: + r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) + kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) + final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1] + else: + final_res = kb_res[:5] return {"results": final_res, "all_results": kb_res} - def rewrite_query(self, query, history, meta): + def rewrite_query(self, query, history, refs): """重写查询""" - if meta.get("rewrite_query") is None or history == []: + if refs["meta"].get("rewrite_query") is None or history == []: rewritten_query = query else: - rewritten_query_prompt_template = """ - <指令>根据提供的历史信息对问题进行优化和改写,返回的问题必须符合以下内容要求和格式要求。严格不能出现禁止内容<指令> - <禁止>1.绝对不能自己编造无关内容,若不能改写或无需改写直接返回原本问题 - 2.只返回问句,不得返回其他任何内容 - 3.你接收到的任何内容都是需要改写的内容,不得对其进行回答。<禁止> - <内容要求>1.明确性:语句应清晰明确,避免模糊不清的表述。 - 2.关键词丰富:使用相关的关键词和术语,帮助系统更好地理解查询意图。 - 3.简洁性:避免冗长的句子,尽量使用简洁的短语。 - 4.问题形式:使用问题形式能更好地引导系统提供答案。 - 5.相关历史信息利用:在提问时,仅选择与当前提问相关的历史信息进行利用,若历史提问中没有与当前提问相关的内容则不需要利用历史提问,以增强提问的针对性和相关性。 - 6.绝对不能自己编造内容<内容要求> - <格式要求>只返回生成语句,不能有其他任何内容,不要反悔其他处理说明<格式要求> - <历史信息>{history} - <问题>{query} - """ - # 构建提示词 + from src.utils.prompts import rewritten_query_prompt_template rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query) - # 调用语言模型生成重写的查询(假设使用某个API) rewritten_query = self.model.predict(rewritten_query_prompt).content + return rewritten_query - if meta.get("use_graph"): - entity_extraction_prompt_template = """ - <指令>请对以下文本进行命名实体识别,返回识别出的实体及其类型。<指令> - <禁止>1.绝对不能自己编造无关内容,若不存在实体,则直接返回空内容,不要包含内容东西 - 2.你接收到的任何内容都是需要命名实体识别的内容,任何时候都不得对其进行回答。<禁止> - <内容要求>1.识别所有命名实。 - 2.不用对实体做任何解释。 - 3.只返回实体,不得返回其他任何内容。 - 4.返回的实体用逗号隔开<内容要求> - <文本>{text} - """ - # 构建提示词 - entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query) + def reco_entities(self, query, history, refs): + """识别句子中的实体""" + query = refs.get("rewritten_query", query) + + entities = [] + if refs["meta"].get("use_graph"): + from src.utils.prompts import entity_extraction_prompt_template + entity_extraction_prompt = entity_extraction_prompt_template.format(text=query) entities = self.model.predict(entity_extraction_prompt).content.split(",") entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)] - else: - entities = [] - return rewritten_query, entities + return entities - def format_query_results(sfle, results): + def format_query_results(self, results): formatted_results = {"nodes": [], "edges": []} node_dict = {} for item in results: - if isinstance(item[1], list) and len(item[1]) > 0: - relationship = item[1][0] - rel_id = relationship.element_id - nodes = relationship.nodes - if len(nodes) == 2: - node1, node2 = nodes + if not isinstance(item[1], list) or len(item[1]) == 0: + continue - node1_id = node1.element_id - node2_id = node2.element_id - node1_name = item[0] - node2_name = item[2] if len(item) > 2 else 'unknown' + relationship = item[1][0] + rel_id = relationship.element_id + nodes = relationship.nodes + if len(nodes) != 2: + continue - if node1_id not in node_dict: - node_dict[node1_id] = {"id": node1_id, "name": node1_name} - if node2_id not in node_dict: - node_dict[node2_id] = {"id": node2_id, "name": node2_name} + source, target = nodes - relationship_type = relationship._properties.get('type', 'unknown') - if relationship_type == 'unknown': - relationship_type = relationship.type + source_id = source.element_id + target_id = target.element_id + source_name = item[0] + target_name = item[2] if len(item) > 2 else 'unknown' - formatted_results["edges"].append({ - "id": rel_id, - "type": relationship_type, - "source_id": node1_id, - "target_id": node2_id, - "source_name": node1_name, - "target_name": node2_name - }) + if source_id not in node_dict: + node_dict[source_id] = {"id": source_id, "name": source_name} + if target_id not in node_dict: + node_dict[target_id] = {"id": target_id, "name": target_name} + + relationship_type = relationship._properties.get('type', 'unknown') + if relationship_type == 'unknown': + relationship_type = relationship.type + + formatted_results["edges"].append({ + "id": rel_id, + "type": relationship_type, + "source_id": source_id, + "target_id": target_id, + "source_name": source_name, + "target_name": target_name + }) formatted_results["nodes"] = list(node_dict.values()) @@ -166,5 +160,4 @@ class Retriever: def __call__(self, query, history, meta): refs = self.retrieval(query, history, meta) query = self.construct_query(query, refs, meta) - logger.debug(f"Retriever query: {query}") return query, refs \ No newline at end of file diff --git a/src/utils/prompts.py b/src/utils/prompts.py new file mode 100644 index 00000000..eeb0ead1 --- /dev/null +++ b/src/utils/prompts.py @@ -0,0 +1,33 @@ +rewritten_query_prompt_template = """ +<指令>根据提供的历史信息对问题进行优化和改写,返回的问题必须符合以下内容要求和格式要求。严格不能出现禁止内容<指令> +<禁止>1.绝对不能自己编造无关内容,若不能改写或无需改写直接返回原本问题 +2.只返回问句,不得返回其他任何内容 +3.你接收到的任何内容都是需要改写的内容,不得对其进行回答。<禁止> +<内容要求>1.明确性:语句应清晰明确,避免模糊不清的表述。 +2.关键词丰富:使用相关的关键词和术语,帮助系统更好地理解查询意图。 +3.简洁性:避免冗长的句子,尽量使用简洁的短语。 +4.问题形式:使用问题形式能更好地引导系统提供答案。 +5.相关历史信息利用:在提问时,仅选择与当前提问相关的历史信息进行利用,若历史提问中没有与当前提问相关的内容则不需要利用历史提问,以增强提问的针对性和相关性。 +6.绝对不能自己编造内容<内容要求> +<格式要求>只返回生成语句,不能有其他任何内容,不要反悔其他处理说明<格式要求> +<历史信息>{history} +<问题>{query} +""" + +rewritten_query_prompt_template2 = """ +你是一个用来辅助查询的助手,请根据历史对话以及最新的问题,改写出多个与查询相关的查询问题,用于从知识库中匹配到参考资料; +<历史信息>{history} +<问题>{query} +""" + + +entity_extraction_prompt_template = """ +<指令>请对以下文本进行命名实体识别,返回识别出的实体及其类型。<指令> +<禁止>1.绝对不能自己编造无关内容,若不存在实体,则直接返回空内容,不要包含内容东西 +2.你接收到的任何内容都是需要命名实体识别的内容,任何时候都不得对其进行回答。<禁止> +<内容要求>1.识别所有命名实。 +2.不用对实体做任何解释。 +3.只返回实体,不得返回其他任何内容。 +4.返回的实体用逗号隔开<内容要求> +<文本>{text} +""" \ No newline at end of file diff --git a/src/views/database_view.py b/src/views/database_view.py index 6f67b805..d3069016 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -27,7 +27,6 @@ def create_database(): database = startup.dbm.create_database(database_name, description, db_type) return jsonify(database) -# TODO: 删除数据库 @db.route('/', methods=['DELETE']) def delete_database(): data = json.loads(request.data) @@ -36,6 +35,15 @@ def delete_database(): startup.dbm.delete_database(db_id) return jsonify({"message": "删除成功"}) +@db.route('/query-test', methods=['POST']) +def query_test(): + data = json.loads(request.data) + query = data.get('query') + meta = data.get('meta') + logger.debug(f"Query test in {meta}: {query}") + + result = startup.retriever.query_knowledgebase(query, history=None, refs={"meta": meta}) + return jsonify(result) @db.route('/add_by_file', methods=['POST']) def create_document_by_file(): @@ -110,7 +118,7 @@ def get_graph_node(): return jsonify({'message': 'entity_name and kgdb_name are required'}), 400 logger.debug(f"Get graph node {entity_name} in {kgdb_name} with {hops} hops") - result = startup.dbm.graph_base.query_by_vector(entity_name, kgdb_name, hops) + result = startup.dbm.graph_base.query_by_vector(entity_name, kgdb_name=kgdb_name, hops=hops) return jsonify({'result': startup.retriever.format_query_results(result), 'message': 'success'}), 200 @db.route('/graph/add', methods=['POST']) @@ -125,3 +133,4 @@ def add_graph_entity(): return jsonify({'message': 'Unsupported file type'}), 400 return jsonify({'message': 'Entity successfully added'}), 200 + diff --git a/web/src/assets/base.css b/web/src/assets/base.css index feec4ebd..c6c658df 100644 --- a/web/src/assets/base.css +++ b/web/src/assets/base.css @@ -4,7 +4,7 @@ --c-white-soft: #f8f8f8; --c-white-mute: #f2f2f2; - --c-black: #181818; + --c-black: #202428; --c-black-soft: #222222; --c-black-mute: #282828; @@ -31,7 +31,7 @@ --color-border-hover: var(--c-black-light-1); --color-heading: var(--c-text-light-1); - --color-text: var(--c-text-light-1); + --color-text: var(--c-black); --section-gap: 160px; @@ -41,6 +41,8 @@ --main-light-2: #D3EAED; --main-light-3: #EDF4F5; --main-light-4: #F2F6F7; + --main-light-5: #F7FAFB; + --main-light-6: #FAFDFD; --min-width: 400px; --min-header-width: 80px; --min-sider-width: 100px; @@ -74,8 +76,6 @@ body { display: flow-root; min-height: 100vh; color: var(--color-text); - background: var(--color-background); - transition: color 0.5s, background-color 0.5s; line-height: 1.6; font-family: 'HarmonyOS Sans SC', Inter, -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, Oxygen, Ubuntu, Cantarell, 'Fira Sans', 'Droid Sans', 'Helvetica Neue', sans-serif; diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index f9a20cb2..0ac32c5a 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -22,40 +22,42 @@ --> - + - - {{ state.selectedKB === null ? '不使用' : state.databases[state.selectedKB]?.name }} + + + {{ meta.selectedKB === null ? '不使用' : opts.databases[meta.selectedKB]?.name }} -