更新了很多

- 添加了系统提示词
- 添加了知识库的加载方式
This commit is contained in:
Wenjie Zhang 2025-03-20 19:51:46 +08:00
parent b46421787d
commit d26fda4053
19 changed files with 423 additions and 460 deletions

View File

@ -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()

View File

@ -1,2 +1,3 @@
from .history import *
from .database import *
from .knowledgebase import KnowledgeBase
from .graphbase import GraphDatabase

View File

@ -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()

View File

@ -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

View File

@ -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)

View File

@ -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

View File

@ -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}")

View File

@ -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()

View File

@ -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]

View File

@ -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"])

View File

@ -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]

View File

@ -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!"}

View File

@ -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

View File

@ -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"}

View File

@ -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}

View File

@ -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('请输入消息');

View File

@ -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);
}

View File

@ -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,
})
})

View File

@ -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 => {