更新了很多
- 添加了系统提示词 - 添加了知识库的加载方式
This commit is contained in:
parent
b46421787d
commit
d26fda4053
@ -8,8 +8,11 @@ executor = ThreadPoolExecutor()
|
||||
from src.config import Config
|
||||
config = Config()
|
||||
|
||||
from src.core import DataBaseManager
|
||||
dbm = DataBaseManager()
|
||||
from src.core import KnowledgeBase
|
||||
knowledge_base = KnowledgeBase()
|
||||
|
||||
from src.core import GraphDatabase
|
||||
graph_base = GraphDatabase()
|
||||
|
||||
from src.core.retriever import Retriever
|
||||
retriever = Retriever()
|
||||
@ -1,2 +1,3 @@
|
||||
from .history import *
|
||||
from .database import *
|
||||
from .knowledgebase import KnowledgeBase
|
||||
from .graphbase import GraphDatabase
|
||||
|
||||
@ -1,305 +0,0 @@
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
import traceback
|
||||
|
||||
from src import config
|
||||
from src.utils import hashstr, logger
|
||||
from src.core.indexing import chunk
|
||||
from src.models.embedding import get_embedding_model
|
||||
|
||||
|
||||
class DataBaseManager:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.database_path = os.path.join(config.save_dir, "data", "database.json")
|
||||
self._load_models()
|
||||
|
||||
def _load_models(self):
|
||||
"""所有需要重启的模型"""
|
||||
self.embed_model = get_embedding_model(config)
|
||||
if config.enable_knowledge_base:
|
||||
from src.core.knowledgebase import KnowledgeBase
|
||||
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
||||
if config.enable_knowledge_graph:
|
||||
from src.core.graphbase import GraphDatabase
|
||||
self.graph_base = GraphDatabase(config, self.embed_model)
|
||||
else:
|
||||
self.graph_base = None
|
||||
|
||||
self.data = {"databases": [], "graph": {}}
|
||||
self._load_databases()
|
||||
self._update_database()
|
||||
|
||||
def _load_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
if not os.path.exists(self.database_path):
|
||||
return
|
||||
|
||||
with open(self.database_path, "r") as f:
|
||||
data = json.load(f)
|
||||
self.data = {
|
||||
"databases": [DataBaseLite(**db) for db in data["databases"]],
|
||||
"graph": data["graph"]
|
||||
}
|
||||
|
||||
# 检查所有文件,如果出现状态是 processing 的,那么设置为 failed
|
||||
for db in self.data["databases"]:
|
||||
for file in db.files:
|
||||
if file["status"] == "processing" or file["status"] == "waiting":
|
||||
file["status"] = "failed"
|
||||
|
||||
def _save_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
self._update_database()
|
||||
os.makedirs(os.path.dirname(self.database_path), exist_ok=True)
|
||||
with open(self.database_path, "w+") as f:
|
||||
json.dump({
|
||||
"databases": [db.to_dict() for db in self.data["databases"]],
|
||||
"graph": self.data["graph"]
|
||||
}, f, ensure_ascii=False, indent=4)
|
||||
|
||||
def _update_database(self):
|
||||
self.id2db = {db.db_id: db for db in self.data["databases"]}
|
||||
self.name2db = {db.name: db for db in self.data["databases"]}
|
||||
self.metaname2db = {db.metaname: db for db in self.data["databases"]}
|
||||
|
||||
def get_databases(self):
|
||||
self._update_database()
|
||||
assert config.enable_knowledge_base, "知识库未启用"
|
||||
knowledge_base_collections = self.knowledge_base.get_collection_names()
|
||||
if len(self.data["databases"]) != len(knowledge_base_collections):
|
||||
logger.warning(
|
||||
f"Database number not match, {knowledge_base_collections}, "
|
||||
f"self.data['databases']: {self.get_db_metanames()}, ")
|
||||
|
||||
# 更新每个数据库的状态信息
|
||||
for db in self.data["databases"]:
|
||||
# 获取最新的集合信息
|
||||
db.update(self.knowledge_base.get_collection_info(db.metaname))
|
||||
|
||||
# 检查文件处理状态
|
||||
processing_files = [f for f in db.files if f["status"] in ["processing", "waiting"]]
|
||||
if processing_files:
|
||||
logger.info(f"数据库 {db.name} 有 {len(processing_files)} 个文件正在处理中")
|
||||
|
||||
return {"databases": [db.to_dict() for db in self.data["databases"]]}
|
||||
|
||||
def get_graph(self):
|
||||
if config.enable_knowledge_graph:
|
||||
self.data["graph"].update(self.graph_base.get_database_info("neo4j"))
|
||||
return {"graph": self.data["graph"]}
|
||||
else:
|
||||
return {"message": "Graph base not enabled", "graph": {}}
|
||||
|
||||
def is_graph_running(self):
|
||||
"""检查图数据库是否正在运行
|
||||
|
||||
Returns:
|
||||
bool: 图数据库是否正在运行
|
||||
"""
|
||||
# 检查是否启用了图数据库
|
||||
if not config.enable_knowledge_graph or not hasattr(self, 'graph_base') or self.graph_base is None:
|
||||
return False
|
||||
|
||||
# 获取图数据库信息,检查状态
|
||||
graph_info = self.graph_base.get_database_info("neo4j")
|
||||
return graph_info.get("status") == "open"
|
||||
|
||||
def create_database(self, database_name, description, db_type, dimension):
|
||||
from src.config import EMBED_MODEL_INFO
|
||||
dimension = dimension or EMBED_MODEL_INFO[config.embed_model]["dimension"]
|
||||
|
||||
new_database = DataBaseLite(database_name,
|
||||
description,
|
||||
db_type,
|
||||
embed_model=config.embed_model,
|
||||
dimension=dimension)
|
||||
|
||||
self.knowledge_base.add_collection(new_database.metaname, dimension)
|
||||
self.data["databases"].append(new_database)
|
||||
self._save_databases()
|
||||
return self.get_databases()
|
||||
|
||||
def add_files(self, db_id, files, params=None):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
|
||||
if db.embed_model != config.embed_model:
|
||||
logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}")
|
||||
return {"message": f"Embed model not match, cur: {config.embed_model}", "status": "failed"}
|
||||
|
||||
# Preprocessing the files to the queue
|
||||
new_files = []
|
||||
for file in files:
|
||||
new_file = {
|
||||
"file_id": "file_" + hashstr(file + str(time.time())),
|
||||
"filename": os.path.basename(file),
|
||||
"path": file,
|
||||
"type": file.split(".")[-1].lower(),
|
||||
"status": "waiting",
|
||||
"created_at": time.time()
|
||||
}
|
||||
db.files.append(new_file)
|
||||
new_files.append(new_file)
|
||||
|
||||
# 先保存一次数据库状态,确保waiting状态被记录
|
||||
self._save_databases()
|
||||
|
||||
for new_file in new_files:
|
||||
file_id = new_file["file_id"]
|
||||
idx = self.get_idx_by_fileid(db, file_id)
|
||||
db.files[idx]["status"] = "processing"
|
||||
# 更新处理状态
|
||||
self._save_databases()
|
||||
|
||||
try:
|
||||
if new_file["type"] == "pdf":
|
||||
texts = self.read_text(new_file["path"])
|
||||
nodes = chunk(texts, params=params)
|
||||
else:
|
||||
nodes = chunk(new_file["path"], params=params)
|
||||
|
||||
self.knowledge_base.add_documents(
|
||||
file_id=file_id,
|
||||
collection_name=db.metaname,
|
||||
docs=[node.text for node in nodes],
|
||||
chunk_infos=[node.dict() for node in nodes])
|
||||
|
||||
idx = self.get_idx_by_fileid(db, file_id)
|
||||
db.files[idx]["status"] = "done"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to add documents to collection {db.metaname}, {e}, {traceback.format_exc()}")
|
||||
idx = self.get_idx_by_fileid(db, file_id)
|
||||
db.files[idx]["status"] = "failed"
|
||||
|
||||
# 每个文件处理完成后立即保存数据库状态
|
||||
self._save_databases()
|
||||
|
||||
|
||||
def get_database_info(self, db_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
return None
|
||||
else:
|
||||
db.update(self.knowledge_base.get_collection_info(db.metaname))
|
||||
return db.to_dict()
|
||||
|
||||
def read_text(self, file, params=None):
|
||||
support_format = [".pdf", ".txt", ".md"]
|
||||
assert os.path.exists(file), "File not found"
|
||||
logger.info(f"Try to read file {file}")
|
||||
|
||||
if not os.path.isfile(file):
|
||||
logger.error(f"Directory not supported now!")
|
||||
raise NotImplementedError("Directory not supported now!")
|
||||
|
||||
if file.endswith(".pdf"):
|
||||
from src.plugins import ocr
|
||||
return ocr.process_pdf(file)
|
||||
|
||||
elif file.endswith(".txt") or file.endswith(".md"):
|
||||
from src.core.filereader import plainreader
|
||||
return plainreader(file)
|
||||
|
||||
else:
|
||||
logger.error(f"File format not supported, only support {support_format}")
|
||||
raise Exception(f"File format not supported, only support {support_format}")
|
||||
|
||||
def delete_file(self, db_id, file_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
file_idx_to_delete = self.get_idx_by_fileid(db, file_id)
|
||||
|
||||
self.knowledge_base.client.delete(
|
||||
collection_name=db.metaname,
|
||||
filter=f"file_id == '{file_id}'"),
|
||||
|
||||
del db.files[file_idx_to_delete]
|
||||
self._save_databases()
|
||||
|
||||
def get_file_info(self, db_id, file_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
return {"message": "database not found"}, 404
|
||||
lines = self.knowledge_base.client.query(
|
||||
collection_name=db.metaname,
|
||||
filter=f"file_id == '{file_id}'",
|
||||
output_fields=None
|
||||
)
|
||||
# 删除 vector 字段
|
||||
for line in lines:
|
||||
line.pop("vector")
|
||||
|
||||
lines.sort(key=lambda x: x.get("start_char_idx") or 0)
|
||||
# logger.debug(f"lines[0]: {lines[0]}")
|
||||
return {"lines": lines}
|
||||
|
||||
def get_db_metanames(self):
|
||||
return [db.metaname for db in self.data["databases"]]
|
||||
|
||||
def delete_database(self, db_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
return {"message": "database not found"}, 404
|
||||
|
||||
self.knowledge_base.client.drop_collection(db.metaname)
|
||||
self.data["databases"] = [d for d in self.data["databases"] if d.db_id != db_id]
|
||||
self._save_databases()
|
||||
return {"message": "删除成功"}
|
||||
|
||||
def get_kb_by_id(self, db_id):
|
||||
for db in self.data["databases"]:
|
||||
if db.db_id == db_id:
|
||||
return db
|
||||
return None
|
||||
|
||||
def get_idx_by_fileid(self, db, file_id):
|
||||
for idx, f in enumerate(db.files):
|
||||
if f["file_id"] == file_id:
|
||||
return idx
|
||||
|
||||
def restart(self):
|
||||
self.embed_model = get_embedding_model(config)
|
||||
self._load_databases()
|
||||
self._update_database()
|
||||
|
||||
|
||||
class DataBaseLite:
|
||||
def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None:
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.db_type = db_type
|
||||
self.dimension = dimension
|
||||
self.db_id = kwargs.get("db_id", hashstr(name))
|
||||
self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}")
|
||||
self.metadata = kwargs.get("metadata", {})
|
||||
self.files = kwargs.get("files", [])
|
||||
self.embed_model = kwargs.get("embed_model", None)
|
||||
|
||||
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
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"db_type": self.db_type,
|
||||
"db_id": self.db_id,
|
||||
"embed_model": self.embed_model,
|
||||
"metaname": self.metaname,
|
||||
"metadata": self.metadata,
|
||||
"files": self.files,
|
||||
"dimension": self.dimension
|
||||
}
|
||||
|
||||
def to_json(self):
|
||||
return json.dumps(self.to_dict(), ensure_ascii=False)
|
||||
|
||||
def __str__(self):
|
||||
return self.to_json()
|
||||
@ -1,24 +0,0 @@
|
||||
import os
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
def pdfreader(file_path):
|
||||
"""读取PDF文件并返回text文本"""
|
||||
assert os.path.exists(file_path), "File not found"
|
||||
assert file_path.endswith(".pdf"), "File format not supported"
|
||||
|
||||
from llama_index.readers.file import PDFReader
|
||||
doc = PDFReader().load_data(file=Path(file_path))
|
||||
|
||||
# 简单的拼接起来之后返回纯文本
|
||||
text = "\n\n".join([d.get_content() for d in doc])
|
||||
return text
|
||||
|
||||
def plainreader(file_path):
|
||||
"""读取普通文本文件并返回text文本"""
|
||||
assert os.path.exists(file_path), "File not found"
|
||||
|
||||
with open(file_path, "r") as f:
|
||||
text = f.read()
|
||||
return text
|
||||
|
||||
@ -6,6 +6,7 @@ import traceback
|
||||
import torch
|
||||
from neo4j import GraphDatabase as GD
|
||||
|
||||
from src import config
|
||||
from src.utils import logger
|
||||
|
||||
warnings.filterwarnings("ignore", category=UserWarning)
|
||||
@ -14,16 +15,13 @@ warnings.filterwarnings("ignore", category=UserWarning)
|
||||
UIE_MODEL = None
|
||||
|
||||
class GraphDatabase:
|
||||
def __init__(self, config, embed_model=None, kgdb_name="neo4j"):
|
||||
self.config = config
|
||||
def __init__(self):
|
||||
self.driver = None
|
||||
self.files = []
|
||||
self.status = "closed"
|
||||
self.kgdb_name = kgdb_name
|
||||
assert embed_model, "embed_model=None"
|
||||
self.embed_model = embed_model
|
||||
self.kgdb_name = "neo4j"
|
||||
self.embed_model_name = None
|
||||
self.work_dir = os.path.join(config.save_dir, "knowledge_graph", kgdb_name)
|
||||
self.work_dir = os.path.join(config.save_dir, "knowledge_graph", self.kgdb_name)
|
||||
os.makedirs(self.work_dir, exist_ok=True)
|
||||
|
||||
# 尝试加载已保存的图数据库信息
|
||||
@ -33,6 +31,8 @@ class GraphDatabase:
|
||||
self.start()
|
||||
|
||||
def start(self):
|
||||
if not config.enable_knowledge_graph or not config.enable_knowledge_base:
|
||||
return
|
||||
uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687")
|
||||
username = os.environ.get("NEO4J_USERNAME", "neo4j")
|
||||
password = os.environ.get("NEO4J_PASSWORD", "0123456789")
|
||||
@ -40,9 +40,9 @@ class GraphDatabase:
|
||||
try:
|
||||
self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password))
|
||||
self.status = "open"
|
||||
logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_database_info()}")
|
||||
logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_graph_info(self.kgdb_name)}")
|
||||
# 连接成功后保存图数据库信息
|
||||
self.save_graph_info()
|
||||
self.save_graph_info(self.kgdb_name)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}")
|
||||
self.config.enable_knowledge_graph = False
|
||||
@ -51,6 +51,13 @@ class GraphDatabase:
|
||||
"""关闭数据库连接"""
|
||||
self.driver.close()
|
||||
|
||||
def is_running(self):
|
||||
"""检查图数据库是否正在运行"""
|
||||
if not config.enable_knowledge_graph or not config.enable_knowledge_base:
|
||||
return False
|
||||
else:
|
||||
return self.status == "open"
|
||||
|
||||
def get_sample_nodes(self, kgdb_name='neo4j', num=50):
|
||||
"""获取指定数据库的 num 个节点信息"""
|
||||
self.use_database(kgdb_name)
|
||||
@ -75,30 +82,6 @@ class GraphDatabase:
|
||||
print(f"数据库 '{kgdb_name}' 创建成功.")
|
||||
return kgdb_name # 返回创建的数据库名称
|
||||
|
||||
def get_database_info(self, db_name="neo4j"):
|
||||
"""获取指定数据库的信息"""
|
||||
self.use_database(db_name)
|
||||
def query(tx):
|
||||
entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"]
|
||||
relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"]
|
||||
triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"]
|
||||
|
||||
# 获取所有标签
|
||||
labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"]
|
||||
|
||||
return {
|
||||
"database_name": db_name,
|
||||
"entity_count": entity_count,
|
||||
"relationship_count": relationship_count,
|
||||
"triples_count": triples_count,
|
||||
"labels": labels,
|
||||
"status": self.status,
|
||||
"embed_model_name": self.embed_model_name
|
||||
}
|
||||
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query)
|
||||
|
||||
def use_database(self, kgdb_name="neo4j"):
|
||||
"""切换到指定数据库"""
|
||||
assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致"
|
||||
@ -159,7 +142,7 @@ class GraphDatabase:
|
||||
|
||||
# 判断模型名称是否匹配
|
||||
from src.config import EMBED_MODEL_INFO
|
||||
cur_embed_info = EMBED_MODEL_INFO[self.config.embed_model]
|
||||
cur_embed_info = EMBED_MODEL_INFO[config.embed_model]
|
||||
self.embed_model_name = self.embed_model_name or cur_embed_info.get('name')
|
||||
assert self.embed_model_name == cur_embed_info.get('name') or self.embed_model_name is None, \
|
||||
f"embed_model_name={self.embed_model_name}, {cur_embed_info.get('name')=}"
|
||||
@ -167,7 +150,7 @@ class GraphDatabase:
|
||||
with self.driver.session() as session:
|
||||
logger.info(f"Adding entity to {kgdb_name}")
|
||||
session.execute_write(_create_graph, triples)
|
||||
logger.info(f"Creating vector index for {kgdb_name} with {self.config.embed_model}")
|
||||
logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}")
|
||||
session.execute_write(_create_vector_index, cur_embed_info['dimension'])
|
||||
# NOTE 这里需要异步处理
|
||||
for i, entry in enumerate(triples):
|
||||
@ -329,7 +312,8 @@ class GraphDatabase:
|
||||
|
||||
def get_embedding(self, text):
|
||||
with torch.no_grad():
|
||||
outputs = self.embed_model.encode([text])[0]
|
||||
from src import knowledge_base
|
||||
outputs = knowledge_base.embed_model.encode([text])[0]
|
||||
return outputs
|
||||
|
||||
def set_embedding(self, tx, entity_name, embedding):
|
||||
@ -338,34 +322,53 @@ class GraphDatabase:
|
||||
CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding)
|
||||
""", name=entity_name, embedding=embedding)
|
||||
|
||||
def save_graph_info(self):
|
||||
def get_graph_info(self, graph_name="neo4j"):
|
||||
self.use_database(graph_name)
|
||||
def query(tx):
|
||||
entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"]
|
||||
relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"]
|
||||
triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"]
|
||||
|
||||
# 获取所有标签
|
||||
labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"]
|
||||
|
||||
return {
|
||||
"graph_name": graph_name,
|
||||
"entity_count": entity_count,
|
||||
"relationship_count": relationship_count,
|
||||
"triples_count": triples_count,
|
||||
"labels": labels,
|
||||
"status": self.status,
|
||||
"embed_model_name": self.embed_model_name,
|
||||
"unindexed_node_count": self.query_nodes_without_embedding(graph_name)
|
||||
}
|
||||
|
||||
try:
|
||||
if self.status == "open" and self.driver and self.is_running():
|
||||
# 获取数据库信息
|
||||
with self.driver.session() as session:
|
||||
graph_info = session.execute_read(query)
|
||||
|
||||
# 添加时间戳
|
||||
from datetime import datetime
|
||||
graph_info["last_updated"] = datetime.now().isoformat()
|
||||
return graph_info
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取图数据库信息失败:{e}, {traceback.format_exc()}")
|
||||
return None
|
||||
|
||||
def save_graph_info(self, graph_name="neo4j"):
|
||||
"""
|
||||
将图数据库的基本信息保存到工作目录中的JSON文件
|
||||
保存的信息包括:数据库名称、状态、嵌入模型名称等
|
||||
"""
|
||||
try:
|
||||
# 获取数据库信息
|
||||
db_info = None
|
||||
if self.status == "open" and self.driver:
|
||||
try:
|
||||
db_info = self.get_database_info(self.kgdb_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"无法获取数据库信息:{e}")
|
||||
graph_info = self.get_graph_info(graph_name)
|
||||
if graph_info is None:
|
||||
logger.error(f"图数据库信息为空,无法保存")
|
||||
return False
|
||||
|
||||
# 构建要保存的信息字典
|
||||
graph_info = {
|
||||
"kgdb_name": self.kgdb_name,
|
||||
"status": self.status,
|
||||
"embed_model_name": self.embed_model_name,
|
||||
"last_updated": None, # 这里可以添加时间戳
|
||||
"database_info": db_info
|
||||
}
|
||||
|
||||
# 添加时间戳
|
||||
from datetime import datetime
|
||||
graph_info["last_updated"] = datetime.now().isoformat()
|
||||
|
||||
# 保存到JSON文件
|
||||
info_file_path = os.path.join(self.work_dir, "graph_info.json")
|
||||
with open(info_file_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(graph_info, f, ensure_ascii=False, indent=2)
|
||||
|
||||
@ -1,9 +1,13 @@
|
||||
from src.utils.prompts import get_system_prompt
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
class HistoryManager():
|
||||
def __init__(self, history=None):
|
||||
def __init__(self, history=None, system_prompt=None):
|
||||
self.messages = history or []
|
||||
|
||||
system_prompt = system_prompt or get_system_prompt()
|
||||
self.add_system(system_prompt)
|
||||
|
||||
def add(self, role, content):
|
||||
self.messages.append({"role": role, "content": content})
|
||||
return self.messages
|
||||
|
||||
@ -5,12 +5,13 @@ from llama_index.core.node_parser import SimpleFileNodeParser
|
||||
from llama_index.core.node_parser import SentenceSplitter
|
||||
from llama_index.readers.file import FlatReader, DocxReader
|
||||
|
||||
from src.utils import hashstr
|
||||
from src.utils import hashstr, logger
|
||||
|
||||
|
||||
def chunk(text_or_path, params=None):
|
||||
"""
|
||||
将文本或文件切分成固定大小的块
|
||||
|
||||
|
||||
Args:
|
||||
text_or_path: 文本或文件路径
|
||||
params: 参数
|
||||
@ -48,3 +49,47 @@ def chunk(text_or_path, params=None):
|
||||
nodes = splitter.get_nodes_from_documents(docs)
|
||||
|
||||
return nodes
|
||||
|
||||
|
||||
|
||||
def pdfreader(file_path):
|
||||
"""读取PDF文件并返回text文本"""
|
||||
assert os.path.exists(file_path), "File not found"
|
||||
assert file_path.endswith(".pdf"), "File format not supported"
|
||||
|
||||
from llama_index.readers.file import PDFReader
|
||||
doc = PDFReader().load_data(file=Path(file_path))
|
||||
|
||||
# 简单的拼接起来之后返回纯文本
|
||||
text = "\n\n".join([d.get_content() for d in doc])
|
||||
return text
|
||||
|
||||
def plainreader(file_path):
|
||||
"""读取普通文本文件并返回text文本"""
|
||||
assert os.path.exists(file_path), "File not found"
|
||||
|
||||
with open(file_path, "r") as f:
|
||||
text = f.read()
|
||||
return text
|
||||
|
||||
def read_text(file, params=None):
|
||||
support_format = [".pdf", ".txt", ".md"]
|
||||
assert os.path.exists(file), "File not found"
|
||||
logger.info(f"Try to read file {file}")
|
||||
|
||||
if not os.path.isfile(file):
|
||||
logger.error(f"Directory not supported now!")
|
||||
raise NotImplementedError("Directory not supported now!")
|
||||
|
||||
if file.endswith(".pdf"):
|
||||
from src.plugins import ocr
|
||||
return ocr.process_pdf(file)
|
||||
|
||||
elif file.endswith(".txt") or file.endswith(".md"):
|
||||
return plainreader(file)
|
||||
|
||||
else:
|
||||
logger.error(f"File format not supported, only support {support_format}")
|
||||
raise Exception(f"File format not supported, only support {support_format}")
|
||||
|
||||
|
||||
|
||||
@ -1,27 +1,202 @@
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
import traceback
|
||||
|
||||
from pymilvus import MilvusClient, MilvusException
|
||||
|
||||
from src import config
|
||||
from src.utils import logger, hashstr
|
||||
from src.core.indexing import chunk, read_text
|
||||
|
||||
|
||||
|
||||
class KnowledgeBase:
|
||||
|
||||
def __init__(self, config=None, embed_model=None) -> None:
|
||||
self.config = config or {}
|
||||
assert embed_model, "embed_model=None"
|
||||
self.embed_model = embed_model
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.data = []
|
||||
self.client = None
|
||||
self.database_path = os.path.join(config.save_dir, "data", "database.json")
|
||||
self._load_models()
|
||||
self._load_databases()
|
||||
|
||||
def _load_models(self):
|
||||
"""所有需要重启的模型"""
|
||||
if not config.enable_knowledge_base:
|
||||
return
|
||||
|
||||
from src.models.embedding import get_embedding_model
|
||||
self.embed_model = get_embedding_model(config)
|
||||
|
||||
if not self.connect_to_milvus():
|
||||
raise ConnectionError("Failed to connect to Milvus")
|
||||
|
||||
def _load_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
if not os.path.exists(self.database_path):
|
||||
return
|
||||
|
||||
with open(self.database_path, "r") as f:
|
||||
data = json.load(f)
|
||||
self.data = [DataBaseLite(**db) for db in data["databases"]]
|
||||
|
||||
self._update_database()
|
||||
|
||||
def _save_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
self._update_database()
|
||||
os.makedirs(os.path.dirname(self.database_path), exist_ok=True)
|
||||
with open(self.database_path, "w") as f:
|
||||
json.dump({
|
||||
"databases": [db.to_dict() for db in self.data],
|
||||
}, f, ensure_ascii=False, indent=4)
|
||||
|
||||
def _update_database(self):
|
||||
self.id2db = {db.db_id: db for db in self.data}
|
||||
self.name2db = {db.name: db for db in self.data}
|
||||
|
||||
def create_database(self, database_name, description, dimension=None):
|
||||
"""创建一个数据库"""
|
||||
dimension = dimension or self.embed_model.get_dimension()
|
||||
db = DataBaseLite(database_name,
|
||||
description,
|
||||
embed_model=self.embed_model.embed_model_fullname,
|
||||
dimension=dimension)
|
||||
|
||||
self.add_collection(db.db_id, dimension)
|
||||
self.data.append(db)
|
||||
self._save_databases()
|
||||
|
||||
def get_databases(self):
|
||||
assert config.enable_knowledge_base, "知识库未启用"
|
||||
|
||||
for db in self.data:
|
||||
db.update(self.get_collection_info(db.db_id))
|
||||
processing_files = [f for fid, f in db.files.items() if f["status"] in ["processing", "waiting"]]
|
||||
if processing_files:
|
||||
logger.info(f"数据库 {db.name} 有 {len(processing_files)} 个文件正在处理中")
|
||||
|
||||
self._save_databases()
|
||||
return {"databases": [db.to_dict() for db in self.data]}
|
||||
|
||||
def get_database_info(self, db_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
return None
|
||||
else:
|
||||
db.update(self.get_collection_info(db.db_id))
|
||||
return db.to_dict()
|
||||
|
||||
def get_database_id(self):
|
||||
return [db.db_id for db in self.data]
|
||||
|
||||
def get_file_info(self, db_id, file_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
raise Exception(f"database not found, {db_id}")
|
||||
|
||||
lines = self.client.query(
|
||||
collection_name=db.db_id,
|
||||
filter=f"file_id == '{file_id}'",
|
||||
output_fields=None
|
||||
)
|
||||
# 删除 vector 字段
|
||||
for line in lines:
|
||||
line.pop("vector")
|
||||
|
||||
lines.sort(key=lambda x: x.get("start_char_idx") or 0)
|
||||
# logger.debug(f"lines[0]: {lines[0]}")
|
||||
return {"lines": lines}
|
||||
|
||||
def get_kb_by_id(self, db_id):
|
||||
return next((db for db in self.data if db.db_id == db_id), None)
|
||||
|
||||
def add_files(self, db_id, files, params=None):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
|
||||
if db.embed_model != config.embed_model:
|
||||
logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}")
|
||||
return {"message": f"Embed model not match, cur: {config.embed_model}", "status": "failed"}
|
||||
|
||||
# Preprocessing the files to the queue
|
||||
new_files = {}
|
||||
for file in files:
|
||||
file_id = "file_" + hashstr(file + str(time.time()))
|
||||
new_file = {
|
||||
"file_id": file_id,
|
||||
"filename": os.path.basename(file),
|
||||
"path": file,
|
||||
"type": file.split(".")[-1].lower(),
|
||||
"status": "waiting",
|
||||
"created_at": time.time()
|
||||
}
|
||||
new_files[file_id] = new_file
|
||||
|
||||
db.files.update(new_files) # 更新数据库状态
|
||||
|
||||
# 先保存一次数据库状态,确保waiting状态被记录
|
||||
self._save_databases()
|
||||
|
||||
for file_id, new_file in new_files.items():
|
||||
db.files[file_id]["status"] = "processing"
|
||||
# 更新处理状态
|
||||
self._save_databases()
|
||||
|
||||
try:
|
||||
if new_file["type"] == "pdf":
|
||||
texts = read_text(new_file["path"])
|
||||
nodes = chunk(texts, params=params)
|
||||
else:
|
||||
nodes = chunk(new_file["path"], params=params)
|
||||
|
||||
self.add_documents(
|
||||
file_id=file_id,
|
||||
collection_name=db.db_id,
|
||||
docs=[node.text for node in nodes],
|
||||
chunk_infos=[node.dict() for node in nodes])
|
||||
|
||||
db.files[file_id]["status"] = "done"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to add documents to collection {db.db_id}, {e}, {traceback.format_exc()}")
|
||||
db.files[file_id]["status"] = "failed"
|
||||
|
||||
# 每个文件处理完成后立即保存数据库状态
|
||||
self._save_databases()
|
||||
|
||||
def delete_file(self, db_id, file_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
raise Exception(f"database not found, {db_id}")
|
||||
|
||||
self.client.delete(collection_name=db.db_id, filter=f"file_id == '{file_id}'")
|
||||
del db.files[file_id]
|
||||
self._save_databases()
|
||||
|
||||
def delete_database(self, db_id):
|
||||
db = self.get_kb_by_id(db_id)
|
||||
if db is None:
|
||||
raise Exception(f"database not found, {db_id}")
|
||||
|
||||
self.client.drop_collection(collection_name=db.db_id)
|
||||
self.data.remove(db)
|
||||
self._save_databases()
|
||||
return {"message": "删除成功"}
|
||||
|
||||
def restart(self):
|
||||
self.embed_model = get_embedding_model(config)
|
||||
self._load_databases()
|
||||
|
||||
################################
|
||||
# Below is the code for milvus #
|
||||
################################
|
||||
def connect_to_milvus(self):
|
||||
"""
|
||||
连接到 Milvus 服务。
|
||||
使用配置中的 URI,如果没有配置,则使用默认值。
|
||||
"""
|
||||
try:
|
||||
uri = os.getenv('MILVUS_URI', self.config.get('milvus_uri', "http://milvus:19530"))
|
||||
uri = os.getenv('MILVUS_URI', config.get('milvus_uri', "http://milvus:19530"))
|
||||
self.client = MilvusClient(uri=uri)
|
||||
# 可以添加一个简单的测试来确保连接成功
|
||||
self.client.list_collections()
|
||||
@ -109,3 +284,46 @@ class KnowledgeBase:
|
||||
def search_by_id(self, collection_name, id, output_fields=["id", "text"]):
|
||||
res = self.client.get(collection_name, id, output_fields=output_fields)
|
||||
return res
|
||||
|
||||
|
||||
class DataBaseLite:
|
||||
def __init__(self, name, description, dimension=None, **kwargs) -> None:
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.dimension = dimension
|
||||
self.metadata = kwargs.get("metadata", {})
|
||||
# logger.debug(f"DataBaseLite init: {self.metadata}")
|
||||
self.db_id = self.metadata.get("collection_name", kwargs.get("db_id")) # metaname 的历史遗留问题
|
||||
self.db_id = self.db_id or f"kb_{hashstr(name, with_salt=True)}"
|
||||
self.files = kwargs.get("files", [])
|
||||
|
||||
if isinstance(self.files, list):
|
||||
self.files = {f["file_id"]: f for f in self.files}
|
||||
|
||||
self.embed_model = kwargs.get("embed_model", None)
|
||||
|
||||
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
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"db_id": self.db_id,
|
||||
"embed_model": self.embed_model,
|
||||
"metadata": self.metadata,
|
||||
"files": self.files,
|
||||
"dimension": self.dimension
|
||||
}
|
||||
|
||||
def to_json(self):
|
||||
return json.dumps(self.to_dict(), ensure_ascii=False)
|
||||
|
||||
def __str__(self):
|
||||
return self.to_json()
|
||||
@ -1,4 +1,4 @@
|
||||
from src import config, dbm
|
||||
from src import config, knowledge_base, graph_base
|
||||
from src.models.rerank_model import get_reranker
|
||||
from src.utils.logging_config import logger
|
||||
from src.models import select_model
|
||||
@ -76,7 +76,7 @@ class Retriever:
|
||||
results = []
|
||||
if refs["meta"].get("use_graph") and config.enable_knowledge_base:
|
||||
for entity in refs["entities"]:
|
||||
result = dbm.graph_base.query_by_vector(entity)
|
||||
result = graph_base.query_by_vector(entity)
|
||||
if result != []:
|
||||
results.extend(result)
|
||||
return {"results": self.format_query_results(results)}
|
||||
@ -88,8 +88,8 @@ class Retriever:
|
||||
kb_res = []
|
||||
final_res = []
|
||||
|
||||
db_name = refs["meta"].get("db_name")
|
||||
if not db_name or not config.enable_knowledge_base:
|
||||
db_id = refs["meta"].get("db_id")
|
||||
if not db_id or not config.enable_knowledge_base:
|
||||
return {
|
||||
"results": final_res,
|
||||
"all_results": kb_res,
|
||||
@ -99,7 +99,7 @@ class Retriever:
|
||||
|
||||
rw_query = self.rewrite_query(query, history, refs)
|
||||
|
||||
kb = dbm.metaname2db[db_name]
|
||||
kb = knowledge_base.id2db[db_id]
|
||||
logger.debug(f"{refs['meta']=}")
|
||||
|
||||
meta = refs["meta"]
|
||||
@ -109,9 +109,9 @@ class Retriever:
|
||||
top_k = meta.get("topK", 5)
|
||||
|
||||
# 检索
|
||||
all_kb_res = dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count)
|
||||
all_kb_res = knowledge_base.search(rw_query, db_id, limit=max_query_count)
|
||||
for r in all_kb_res:
|
||||
r["file"] = kb.id2file(r["entity"]["file_id"])
|
||||
r["file"] = kb.files[r["entity"]["file_id"]]
|
||||
|
||||
kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold]
|
||||
|
||||
|
||||
@ -4,13 +4,22 @@ import requests
|
||||
from FlagEmbedding import FlagModel
|
||||
from zhipuai import ZhipuAI
|
||||
|
||||
from src import config
|
||||
from src.config import EMBED_MODEL_INFO
|
||||
from src.utils import hashstr, logger, get_docker_safe_url
|
||||
|
||||
|
||||
class BaseEmbeddingModel:
|
||||
embed_state = {}
|
||||
EMBED_MODEL_INFO = EMBED_MODEL_INFO
|
||||
|
||||
def get_dimension(self):
|
||||
if hasattr(self, "dimension"):
|
||||
return self.dimension
|
||||
|
||||
if hasattr("embed_model_fullname"):
|
||||
return EMBED_MODEL_INFO[self.embed_model_fullname].get("dimension", None)
|
||||
|
||||
return EMBED_MODEL_INFO[self.model].get("dimension", None)
|
||||
|
||||
def encode(self, message):
|
||||
return self.predict(message)
|
||||
@ -49,6 +58,8 @@ class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel):
|
||||
|
||||
self.model = config.model_local_paths.get(info["name"], info.get("local_path"))
|
||||
self.model = self.model or info["name"]
|
||||
self.dimension = info["dimension"]
|
||||
self.embed_model_fullname = config.embed_model
|
||||
|
||||
if os.path.exists(_path := os.path.join(os.getenv("MODEL_DIR"), self.model)):
|
||||
self.model = _path
|
||||
@ -69,7 +80,9 @@ class ZhipuEmbedding(BaseEmbeddingModel):
|
||||
def __init__(self, config) -> None:
|
||||
self.config = config
|
||||
self.model = EMBED_MODEL_INFO[config.embed_model]["name"]
|
||||
self.dimension = EMBED_MODEL_INFO[config.embed_model]["dimension"]
|
||||
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY"))
|
||||
self.embed_model_fullname = config.embed_model
|
||||
|
||||
def predict(self, message):
|
||||
response = self.client.embeddings.create(
|
||||
@ -86,6 +99,8 @@ class OllamaEmbedding(BaseEmbeddingModel):
|
||||
self.model = self.info["name"]
|
||||
self.url = self.info.get("url", "http://localhost:11434/api/embed")
|
||||
self.url = get_docker_safe_url(self.url)
|
||||
self.dimension = self.info.get("dimension", None)
|
||||
self.embed_model_fullname = config.embed_model
|
||||
|
||||
def predict(self, message: list[str] | str):
|
||||
if isinstance(message, str):
|
||||
@ -105,6 +120,8 @@ class OtherEmbedding(BaseEmbeddingModel):
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
self.info = EMBED_MODEL_INFO[config.embed_model]
|
||||
self.embed_model_fullname = config.embed_model
|
||||
self.dimension = self.info.get("dimension", None)
|
||||
self.model = self.info["name"]
|
||||
self.api_key = os.getenv(self.info["api_key"], None)
|
||||
self.url = get_docker_safe_url(self.info["url"])
|
||||
|
||||
@ -40,7 +40,7 @@ class SilconFlowReranker():
|
||||
payload = self.build_payload(query, sentences, max_length)
|
||||
response = requests.request("POST", self.url, json=payload, headers=self.headers)
|
||||
response = json.loads(response.text)
|
||||
logger.debug(f"SiliconFlow Reranker response: {response}")
|
||||
# logger.debug(f"SiliconFlow Reranker response: {response}")
|
||||
|
||||
results = sorted(response["results"], key=lambda x: x["index"])
|
||||
all_scores = [result["relevance_score"] for result in results]
|
||||
|
||||
@ -4,7 +4,7 @@ from fastapi import Request, Body
|
||||
|
||||
base = APIRouter()
|
||||
|
||||
from src import config, dbm, retriever
|
||||
from src import config, retriever, knowledge_base, graph_base
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
@ -27,7 +27,8 @@ async def update_config(key = Body(...), value = Body(...)):
|
||||
|
||||
@base.post("/restart")
|
||||
async def restart():
|
||||
dbm.restart()
|
||||
knowledge_base.restart()
|
||||
graph_base.restart()
|
||||
retriever.restart()
|
||||
return {"message": "Restarted!"}
|
||||
|
||||
|
||||
@ -37,7 +37,7 @@ def chat_post(
|
||||
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
||||
|
||||
def need_retrieve(meta):
|
||||
return meta.get("use_web") or meta.get("use_graph") or meta.get("db_name")
|
||||
return meta.get("use_web") or meta.get("use_graph") or meta.get("db_id")
|
||||
|
||||
def generate_response():
|
||||
modified_query = query
|
||||
|
||||
@ -5,7 +5,7 @@ from typing import List, Optional
|
||||
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body
|
||||
|
||||
from src.utils import logger, hashstr
|
||||
from src import executor, dbm, retriever, config
|
||||
from src import executor, retriever, config, knowledge_base, graph_base
|
||||
|
||||
data = APIRouter(prefix="/data")
|
||||
|
||||
@ -13,8 +13,9 @@ data = APIRouter(prefix="/data")
|
||||
@data.get("/")
|
||||
async def get_databases():
|
||||
try:
|
||||
database = dbm.get_databases()
|
||||
database = knowledge_base.get_databases()
|
||||
except Exception as e:
|
||||
logger.error(f"获取数据库列表失败 {e}, {traceback.format_exc()}")
|
||||
return {"message": f"获取数据库列表失败 {e}", "databases": []}
|
||||
return database
|
||||
|
||||
@ -22,22 +23,24 @@ async def get_databases():
|
||||
async def create_database(
|
||||
database_name: str = Body(...),
|
||||
description: str = Body(...),
|
||||
db_type: str = Body(...),
|
||||
dimension: Optional[int] = Body(None)
|
||||
):
|
||||
logger.debug(f"Create database {database_name}")
|
||||
database_info = dbm.create_database(
|
||||
database_name,
|
||||
description,
|
||||
db_type,
|
||||
dimension=dimension
|
||||
)
|
||||
try:
|
||||
database_info = knowledge_base.create_database(
|
||||
database_name,
|
||||
description,
|
||||
dimension=dimension
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"创建数据库失败 {e}, {traceback.format_exc()}")
|
||||
return {"message": f"创建数据库失败 {e}", "status": "failed"}
|
||||
return database_info
|
||||
|
||||
@data.delete("/")
|
||||
async def delete_database(db_id):
|
||||
logger.debug(f"Delete database {db_id}")
|
||||
dbm.delete_database(db_id)
|
||||
knowledge_base.delete_database(db_id)
|
||||
return {"message": "删除成功"}
|
||||
|
||||
@data.post("/query-test")
|
||||
@ -54,17 +57,17 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod
|
||||
loop = asyncio.get_event_loop()
|
||||
await loop.run_in_executor(
|
||||
executor, # 使用与chat_router相同的线程池
|
||||
lambda: dbm.add_files(db_id, files)
|
||||
lambda: knowledge_base.add_files(db_id, files)
|
||||
)
|
||||
return {"message": "文件添加完成", "status": "success"}
|
||||
except Exception as e:
|
||||
logger.error(f"添加文件失败: {e}")
|
||||
logger.error(f"添加文件失败: {e}, {traceback.format_exc()}")
|
||||
return {"message": f"添加文件失败: {e}", "status": "failed"}
|
||||
|
||||
@data.get("/info")
|
||||
async def get_database_info(db_id: str):
|
||||
logger.debug(f"Get database {db_id} info")
|
||||
database = dbm.get_database_info(db_id)
|
||||
database = knowledge_base.get_database_info(db_id)
|
||||
if database is None:
|
||||
raise HTTPException(status_code=404, detail="Database not found")
|
||||
return database
|
||||
@ -72,7 +75,7 @@ async def get_database_info(db_id: str):
|
||||
@data.delete("/document")
|
||||
async def delete_document(db_id: str = Body(...), file_id: str = Body(...)):
|
||||
logger.debug(f"DELETE document {file_id} info in {db_id}")
|
||||
dbm.delete_file(db_id, file_id)
|
||||
knowledge_base.delete_file(db_id, file_id)
|
||||
return {"message": "删除成功"}
|
||||
|
||||
@data.get("/document")
|
||||
@ -80,10 +83,10 @@ async def get_document_info(db_id: str, file_id: str):
|
||||
logger.debug(f"GET document {file_id} info in {db_id}")
|
||||
|
||||
try:
|
||||
info = dbm.get_file_info(db_id, file_id)
|
||||
info = knowledge_base.get_file_info(db_id, file_id)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get file info, {e}, {db_id=}, {file_id=}, {traceback.format_exc()}")
|
||||
info = {"message": "Failed to get file info", "status": "failed"}, 500
|
||||
info = {"message": "Failed to get file info", "status": "failed"}
|
||||
|
||||
return info
|
||||
|
||||
@ -105,36 +108,27 @@ async def upload_file(file: UploadFile = File(...)):
|
||||
|
||||
@data.get("/graph")
|
||||
async def get_graph_info():
|
||||
graph_info = dbm.get_graph()
|
||||
|
||||
# 获取未索引节点数量
|
||||
unindexed_count = 0
|
||||
if dbm.is_graph_running():
|
||||
# 调用GraphDatabase的query_nodes_without_embedding方法
|
||||
unindexed_nodes = dbm.graph_base.query_nodes_without_embedding()
|
||||
unindexed_count = len(unindexed_nodes) if unindexed_nodes else 0
|
||||
|
||||
# 将未索引节点数量添加到返回结果中
|
||||
graph_info["graph"]["unindexed_node_count"] = unindexed_count
|
||||
|
||||
graph_info = graph_base.get_graph_info()
|
||||
if graph_info is None:
|
||||
raise HTTPException(status_code=400, detail="图数据库获取出错")
|
||||
return graph_info
|
||||
|
||||
@data.post("/graph/index-nodes")
|
||||
async def index_nodes(data: dict = Body(default={})):
|
||||
if not dbm.is_graph_running():
|
||||
if not graph_base.is_running():
|
||||
raise HTTPException(status_code=400, detail="图数据库未启动")
|
||||
|
||||
# 获取参数或使用默认值
|
||||
kgdb_name = data.get('kgdb_name', 'neo4j')
|
||||
|
||||
# 调用GraphDatabase的add_embedding_to_nodes方法
|
||||
count = dbm.graph_base.add_embedding_to_nodes(kgdb_name=kgdb_name)
|
||||
count = graph_base.add_embedding_to_nodes(kgdb_name=kgdb_name)
|
||||
|
||||
return {"status": "success", "message": f"已成功为{count}个节点添加嵌入向量", "indexed_count": count}
|
||||
|
||||
@data.get("/graph/node")
|
||||
async def get_graph_node(entity_name: str):
|
||||
result = dbm.graph_base.query_node(entity_name=entity_name)
|
||||
result = graph_base.query_node(entity_name=entity_name)
|
||||
return {"result": retriever.format_query_results(result), "message": "success"}
|
||||
|
||||
@data.get("/graph/nodes")
|
||||
@ -143,7 +137,7 @@ async def get_graph_nodes(kgdb_name: str, num: int):
|
||||
raise HTTPException(status_code=400, detail="Knowledge graph is not enabled")
|
||||
|
||||
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
||||
result = dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
||||
result = graph_base.get_sample_nodes(kgdb_name, num)
|
||||
return {"result": retriever.format_general_results(result), "message": "success"}
|
||||
|
||||
@data.post("/graph/add-by-jsonl")
|
||||
@ -154,6 +148,6 @@ async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str]
|
||||
if not file_path.endswith('.jsonl'):
|
||||
raise HTTPException(status_code=400, detail="file_path must be a jsonl file")
|
||||
|
||||
dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name)
|
||||
graph_base.jsonl_file_add_entity(file_path, kgdb_name)
|
||||
return {"message": "Entity successfully added"}
|
||||
|
||||
|
||||
@ -1,9 +1,11 @@
|
||||
system_prompt = """
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
def get_system_prompt():
|
||||
return (f"当前时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n")
|
||||
|
||||
|
||||
knowbase_qa_template = """
|
||||
请利用查询到的资料回答问题,回答问题时,不要过度的分点作答。如果非要分点作答,可以使用 一、二、等:
|
||||
请利用查询到的资料回答问题,回答问题时,不要过度的分点作答。
|
||||
|
||||
<参考资料>:
|
||||
{external}
|
||||
|
||||
@ -226,7 +226,7 @@ const meta = reactive(JSON.parse(localStorage.getItem('meta')) || {
|
||||
stream: true,
|
||||
summary_title: false,
|
||||
history_round: 5,
|
||||
db_name: null,
|
||||
db_id: null,
|
||||
})
|
||||
|
||||
const marked = new Marked(
|
||||
@ -524,7 +524,7 @@ const fetchChatResponse = (user_input, cur_res_id) => {
|
||||
// 更新后的 sendMessage 函数
|
||||
const sendMessage = () => {
|
||||
const user_input = conv.value.inputText.trim();
|
||||
const dbName = opts.databases.length > 0 ? opts.databases[meta.selectedKB]?.metaname : null;
|
||||
const dbID = opts.databases.length > 0 ? opts.databases[meta.selectedKB]?.db_id : null;
|
||||
if (isStreaming.value) {
|
||||
message.error('请等待上一条消息处理完成');
|
||||
return
|
||||
@ -537,7 +537,7 @@ const sendMessage = () => {
|
||||
|
||||
const cur_res_id = conv.value.messages[conv.value.messages.length - 1].id;
|
||||
conv.value.inputText = '';
|
||||
meta.db_name = dbName;
|
||||
meta.db_id = dbID;
|
||||
fetchChatResponse(user_input, cur_res_id)
|
||||
} else {
|
||||
console.log('请输入消息');
|
||||
|
||||
@ -54,7 +54,7 @@
|
||||
</a-button>
|
||||
<a-button @click="handleRefresh" :loading="state.refrashing">刷新状态</a-button>
|
||||
</div>
|
||||
<a-table :columns="columns" :data-source="database.files" row-key="filename" class="my-table">
|
||||
<a-table :columns="columns" :data-source="Object.values(database.files || {})" row-key="file_id" class="my-table">
|
||||
<template #bodyCell="{ column, text, record }">
|
||||
<template v-if="column.key === 'filename'">
|
||||
<a-button class="main-btn" type="link" @click="openFileDetail(record)">{{ text }}</a-button>
|
||||
@ -93,8 +93,8 @@
|
||||
placement="right"
|
||||
@after-open-change="afterOpenChange"
|
||||
>
|
||||
<h2>共 {{ selectedFile?.lines.length }} 个片段</h2>
|
||||
<p v-for="line in selectedFile?.lines" :key="line.id" class="line-text">
|
||||
<h2>共 {{ selectedFile?.lines?.length || 0 }} 个片段</h2>
|
||||
<p v-for="line in selectedFile?.lines || []" :key="line.id" class="line-text">
|
||||
{{ line.text }}
|
||||
</p>
|
||||
</a-drawer>
|
||||
@ -302,7 +302,7 @@ const onQuery = () => {
|
||||
state.searchLoading = false
|
||||
return
|
||||
}
|
||||
meta.db_name = database.value.metaname
|
||||
meta.db_id = database.value.db_id
|
||||
fetch('/api/data/query-test', {
|
||||
method: "POST",
|
||||
headers: {
|
||||
@ -402,9 +402,15 @@ const openFileDetail = (record) => {
|
||||
.then(response => response.json())
|
||||
.then(data => {
|
||||
console.log(data)
|
||||
if (data.status == "failed") {
|
||||
message.error(data.message)
|
||||
return
|
||||
}
|
||||
state.lock = false
|
||||
selectedFile.value = record
|
||||
selectedFile.value.lines = data.lines
|
||||
selectedFile.value = {
|
||||
...record,
|
||||
lines: data.lines || []
|
||||
}
|
||||
state.drawer = true
|
||||
})
|
||||
.catch(error => {
|
||||
@ -517,7 +523,6 @@ const addDocumentByFile = () => {
|
||||
})
|
||||
.finally(() => {
|
||||
getDatabaseInfo()
|
||||
// 不在这里清除定时器,而是在定时器内部根据文件状态清除
|
||||
state.loading = false
|
||||
})
|
||||
}
|
||||
@ -875,7 +880,7 @@ onUnmounted(() => {
|
||||
.custom-class .line-text {
|
||||
padding: 10px;
|
||||
border-radius: 4px;
|
||||
|
||||
|
||||
&:hover {
|
||||
background-color: var(--main-light-4);
|
||||
}
|
||||
|
||||
@ -45,7 +45,7 @@
|
||||
<div class="icon"><ReadFilled /></div>
|
||||
<div class="info">
|
||||
<h3>{{ database.name }}</h3>
|
||||
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata.row_count }}行</span></p>
|
||||
<p><span>{{ database.metadata.row_count }} 行</span></p>
|
||||
</div>
|
||||
</div>
|
||||
<p class="description">{{ database.description || '暂无描述' }}</p>
|
||||
@ -138,7 +138,6 @@ const createDatabase = () => {
|
||||
body: JSON.stringify({
|
||||
database_name: newDatabase.name,
|
||||
description: newDatabase.description,
|
||||
db_type: "knowledge",
|
||||
dimension: newDatabase.dimension ? parseInt(newDatabase.dimension) : null,
|
||||
})
|
||||
})
|
||||
|
||||
@ -133,7 +133,7 @@ const loadGraphInfo = () => {
|
||||
.then(response => response.json())
|
||||
.then(data => {
|
||||
console.log(data)
|
||||
graphInfo.value = data.graph
|
||||
graphInfo.value = data
|
||||
state.loadingGraphInfo = false
|
||||
})
|
||||
.catch(error => {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user