优化启动项

This commit is contained in:
Wenjie Zhang 2025-03-16 01:11:13 +08:00
parent e7a52f66df
commit dd617c246d
8 changed files with 104 additions and 112 deletions

View File

@ -1,3 +1,15 @@
from dotenv import load_dotenv
load_dotenv("src/.env")
load_dotenv("src/.env")
from concurrent.futures import ThreadPoolExecutor
executor = ThreadPoolExecutor()
from src.config import Config
config = Config()
from src.core import DataBaseManager
dbm = DataBaseManager()
from src.core.retriever import Retriever
retriever = Retriever()

View File

@ -1,6 +1,8 @@
import os
import json
import time
from src import config
from src.utils import hashstr, logger
from src.core.indexing import chunk
from src.models.embedding import get_embedding_model
@ -8,22 +10,23 @@ from src.models.embedding import get_embedding_model
class DataBaseManager:
def __init__(self, config=None) -> None:
self.config = config
def __init__(self) -> None:
self.database_path = os.path.join(config.save_dir, "data", "database.json")
self.embed_model = get_embedding_model(config)
self._load_models()
if self.config.enable_knowledge_base:
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 self.config.enable_knowledge_graph:
if config.enable_knowledge_graph:
from src.core.graphbase import GraphDatabase
self.graph_base = GraphDatabase(self.config, self.embed_model)
self.graph_base = GraphDatabase(config, self.embed_model)
else:
self.graph_base = None
self.data = {"databases": [], "graph": {}}
self._load_databases()
self._update_database()
@ -62,7 +65,7 @@ class DataBaseManager:
def get_databases(self):
self._update_database()
assert self.config.enable_knowledge_base, "知识库未启用"
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(
@ -82,7 +85,7 @@ class DataBaseManager:
return {"databases": [db.to_dict() for db in self.data["databases"]]}
def get_graph(self):
if self.config.enable_knowledge_graph:
if config.enable_knowledge_graph:
self.data["graph"].update(self.graph_base.get_database_info("neo4j"))
return {"graph": self.data["graph"]}
else:
@ -95,7 +98,7 @@ class DataBaseManager:
bool: 图数据库是否正在运行
"""
# 检查是否启用了图数据库
if not self.config.enable_knowledge_graph or not hasattr(self, 'graph_base') or self.graph_base is None:
if not config.enable_knowledge_graph or not hasattr(self, 'graph_base') or self.graph_base is None:
return False
# 获取图数据库信息,检查状态
@ -104,12 +107,12 @@ class DataBaseManager:
def create_database(self, database_name, description, db_type, dimension):
from src.config import EMBED_MODEL_INFO
dimension = dimension or EMBED_MODEL_INFO[self.config.embed_model]["dimension"]
dimension = dimension or EMBED_MODEL_INFO[config.embed_model]["dimension"]
new_database = DataBaseLite(database_name,
description,
db_type,
embed_model=self.config.embed_model,
embed_model=config.embed_model,
dimension=dimension)
self.knowledge_base.add_collection(new_database.metaname, dimension)
@ -120,9 +123,9 @@ class DataBaseManager:
def add_files(self, db_id, files, params=None):
db = self.get_kb_by_id(db_id)
if db.embed_model != self.config.embed_model:
logger.error(f"Embed model not match, {db.embed_model} != {self.config.embed_model}")
return {"message": f"Embed model not match, cur: {self.config.embed_model}", "status": "failed"}
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 = []
@ -254,6 +257,11 @@ class DataBaseManager:
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:

View File

@ -1,24 +1,24 @@
from src import config, dbm
from src.models.rerank_model import get_reranker
from src.utils.logging_config import logger
from src.models import select_model
class Retriever:
def __init__(self, config, dbm, model):
self.config = config
self.dbm = dbm
self.model = model
def __init__(self):
self._load_models()
if self.config.enable_reranker:
def _load_models(self):
if config.enable_reranker:
self.reranker = get_reranker(config)
if self.config.enable_web_search:
if config.enable_web_search:
from src.utils.web_search import WebSearcher
self.web_searcher = WebSearcher()
def retrieval(self, query, history, meta):
refs = {"query": query, "history": history, "meta": meta}
refs["model_name"] = self.config.model_name
refs["model_name"] = config.model_name
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)
@ -26,6 +26,10 @@ class Retriever:
return refs
def restart(self):
"""所有需要重启的模型"""
self._load_models()
def construct_query(self, query, refs, meta):
logger.debug(f"{refs=}")
if not refs or len(refs) == 0:
@ -70,9 +74,9 @@ class Retriever:
def query_graph(self, query, history, refs):
results = []
if refs["meta"].get("use_graph") and self.config.enable_knowledge_base:
if refs["meta"].get("use_graph") and config.enable_knowledge_base:
for entity in refs["entities"]:
result = self.dbm.graph_base.query_by_vector(entity)
result = dbm.graph_base.query_by_vector(entity)
if result != []:
results.extend(result)
return {"results": self.format_query_results(results)}
@ -85,7 +89,7 @@ class Retriever:
final_res = []
db_name = refs["meta"].get("db_name")
if not db_name or not self.config.enable_knowledge_base:
if not db_name or not config.enable_knowledge_base:
return {
"results": final_res,
"all_results": kb_res,
@ -95,7 +99,7 @@ class Retriever:
rw_query = self.rewrite_query(query, history, refs)
kb = self.dbm.metaname2db[db_name]
kb = dbm.metaname2db[db_name]
logger.debug(f"{refs['meta']=}")
meta = refs["meta"]
@ -105,14 +109,14 @@ class Retriever:
top_k = meta.get("topK", 5)
# 检索
all_kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count)
all_kb_res = dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count)
for r in all_kb_res:
r["file"] = kb.id2file(r["entity"]["file_id"])
kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold]
# 重排序
if self.config.enable_reranker and len(kb_res) > 0:
if config.enable_reranker and len(kb_res) > 0:
texts = [r["entity"]["text"] for r in kb_res]
rerank_scores = self.reranker.compute_score([rw_query, texts], normalize=True)
for i, r in enumerate(kb_res):
@ -127,7 +131,7 @@ class Retriever:
def query_web(self, query, history, refs):
"""查询网络"""
if not (refs["meta"].get("use_web") and self.config.enable_web_search):
if not (refs["meta"].get("use_web") and config.enable_web_search):
return {"results": [], "message": "Web search is disabled"}
try:
@ -140,10 +144,11 @@ class Retriever:
def rewrite_query(self, query, history, refs):
"""重写查询"""
model = select_model(config)
if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置
rewrite_query_span = refs["meta"].get("use_rewrite_query", "off")
else:
rewrite_query_span = self.config.use_rewrite_query
rewrite_query_span = config.use_rewrite_query
if rewrite_query_span == "off":
rewritten_query = query
@ -152,10 +157,10 @@ class Retriever:
history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else ""
rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query)
rewritten_query = self.model.predict(rewritten_query_prompt).content
rewritten_query = model.predict(rewritten_query_prompt).content
if rewrite_query_span == "hyde":
hy_doc = self.model.predict(rewritten_query).content
hy_doc = model.predict(rewritten_query).content
rewritten_query = f"{rewritten_query} {hy_doc}"
return rewritten_query
@ -163,6 +168,7 @@ class Retriever:
def reco_entities(self, query, history, refs):
"""识别句子中的实体"""
query = refs.get("rewritten_query", query)
model = select_model(config)
entities = []
if refs["meta"].get("use_graph"):
@ -170,7 +176,7 @@ class Retriever:
from src.utils.prompts import keywords_prompt_template as entity_template
entity_extraction_prompt = entity_template.format(text=query)
entities = self.model.predict(entity_extraction_prompt).content.split("<->")
entities = model.predict(entity_extraction_prompt).content.split("<->")
# entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)]
return entities

View File

@ -1,35 +0,0 @@
import os
from concurrent.futures import ThreadPoolExecutor
from src.core import DataBaseManager
from src.core.retriever import Retriever
from src.models import select_model
from src.config import Config
from src.utils import logger
# 创建线程池
executor = ThreadPoolExecutor()
class Startup:
def __init__(self):
self.start()
def start(self):
self.config = Config()
self.model = select_model(self.config)
self.dbm = DataBaseManager(self.config)
self.retriever = Retriever(self.config, self.dbm, self.model)
logger.info(f"Loading lite model: {self.config.model_name_lite}")
self.model_lite = select_model(self.config,
model_provider=self.config.model_provider_lite,
model_name=self.config.model_name_lite)
def restart(self):
logger.info("Restarting...")
self.start()
logger.info("Restarted")
startup = Startup()

View File

@ -1,13 +1,9 @@
from fastapi import APIRouter
from fastapi import Request, Body
base = APIRouter()
from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse
from fastapi import Request, Body
from src.core import HistoryManager
from src.core.startup import startup
from src import config, dbm, retriever
from src.utils import logger
@ -17,17 +13,18 @@ async def route_index():
@base.get("/config")
def get_config():
return startup.config
return config
@base.post("/config")
async def update_config(key = Body(...), value = Body(...)):
startup.config[key] = value
startup.config.save()
return startup.config
config[key] = value
config.save()
return config
@base.post("/restart")
async def restart():
startup.restart()
dbm.restart()
retriever.restart()
return {"message": "Restarted!"}
@base.get("/log")

View File

@ -3,7 +3,8 @@ import asyncio
from fastapi import APIRouter, Body
from fastapi.responses import StreamingResponse, Response
from src.core import HistoryManager
from src.core.startup import startup, executor
from src import executor, config, retriever
from src.models import select_model
from src.utils.logging_config import logger
chat = APIRouter(prefix="/chat")
@ -21,14 +22,15 @@ def chat_post(
history: list = Body(...),
cur_res_id: str = Body(...)):
meta["server_model_name"] = startup.model.model_name
model = select_model(config)
meta["server_model_name"] = model.model_name
history_manager = HistoryManager(history)
logger.debug(f"Received query: {query} with meta: {meta}")
def make_chunk(content=None, **kwargs):
return json.dumps({
"response": content,
"model_name": startup.config.model_name,
"model_name": config.model_name,
"meta": meta,
**kwargs
}, ensure_ascii=False).encode('utf-8') + b"\n"
@ -46,7 +48,7 @@ def chat_post(
yield chunk
try:
modified_query, refs = startup.retriever(modified_query, history_manager.messages, meta)
modified_query, refs = retriever(modified_query, history_manager.messages, meta)
except Exception as e:
logger.error(f"Retriever error: {e}")
yield make_chunk(message=f"Retriever error: {e}", status="error")
@ -60,7 +62,7 @@ def chat_post(
content = ""
reasoning_content = ""
try:
for delta in startup.model.predict(messages, stream=True):
for delta in model.predict(messages, stream=True):
if not delta.content and hasattr(delta, 'reasoning_content'):
reasoning_content += delta.reasoning_content or ""
chunk = make_chunk(reasoning_content=reasoning_content, status="reasoning")
@ -91,9 +93,10 @@ def chat_post(
@chat.post("/call")
async def call(query: str = Body(...), meta: dict = Body(None)):
model = select_model(config, model_provider=meta.get("model_provider"), model_name=meta.get("model_name"))
async def predict_async(query):
loop = asyncio.get_event_loop()
return await loop.run_in_executor(executor, startup.model.predict, query)
return await loop.run_in_executor(executor, model.predict, query)
response = await predict_async(query)
logger.debug({"query": query, "response": response.content})
@ -104,7 +107,10 @@ async def call(query: str = Body(...), meta: dict = Body(None)):
async def call(query: str = Body(...), meta: dict = Body(None)):
async def predict_async(query):
loop = asyncio.get_event_loop()
return await loop.run_in_executor(executor, startup.model_lite.predict, query)
model_provider = meta.get("model_provider", config.model_provider_lite)
model_name = meta.get("model_name", config.model_name_lite)
model = select_model(config, model_provider=model_provider, model_name=model_name)
return await loop.run_in_executor(executor, model.predict, query)
response = await predict_async(query)
logger.debug({"query": query, "response": response.content})

View File

@ -4,7 +4,7 @@ from typing import List, Optional
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body
from src.utils import logger, hashstr
from src.core.startup import startup, executor
from src import executor, dbm, retriever, config
data = APIRouter(prefix="/data")
@ -12,7 +12,7 @@ data = APIRouter(prefix="/data")
@data.get("/")
async def get_databases():
try:
database = startup.dbm.get_databases()
database = dbm.get_databases()
except Exception as e:
return {"message": f"获取数据库列表失败 {e}", "databases": []}
return database
@ -25,7 +25,7 @@ async def create_database(
dimension: Optional[int] = Body(None)
):
logger.debug(f"Create database {database_name}")
database_info = startup.dbm.create_database(
database_info = dbm.create_database(
database_name,
description,
db_type,
@ -36,13 +36,13 @@ async def create_database(
@data.delete("/")
async def delete_database(db_id):
logger.debug(f"Delete database {db_id}")
startup.dbm.delete_database(db_id)
dbm.delete_database(db_id)
return {"message": "删除成功"}
@data.post("/query-test")
async def query_test(query: str = Body(...), meta: dict = Body(...)):
logger.debug(f"Query test in {meta}: {query}")
result = startup.retriever.query_knowledgebase(query, history=None, refs={"meta": meta})
result = retriever.query_knowledgebase(query, history=None, refs={"meta": meta})
return result
@data.post("/add-by-file")
@ -53,7 +53,7 @@ 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: startup.dbm.add_files(db_id, files)
lambda: dbm.add_files(db_id, files)
)
return {"message": "文件添加完成", "status": "success"}
except Exception as e:
@ -63,7 +63,7 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod
@data.get("/info")
async def get_database_info(db_id: str):
logger.debug(f"Get database {db_id} info")
database = startup.dbm.get_database_info(db_id)
database = dbm.get_database_info(db_id)
if database is None:
raise HTTPException(status_code=404, detail="Database not found")
return database
@ -71,7 +71,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}")
startup.dbm.delete_file(db_id, file_id)
dbm.delete_file(db_id, file_id)
return {"message": "删除成功"}
@data.get("/document")
@ -79,7 +79,7 @@ async def get_document_info(db_id: str, file_id: str):
logger.debug(f"GET document {file_id} info in {db_id}")
try:
info = startup.dbm.get_file_info(db_id, file_id)
info = dbm.get_file_info(db_id, file_id)
except Exception as e:
logger.error(f"Failed to get file info, {e}, {db_id=}, {file_id=}")
info = {"message": "Failed to get file info", "status": "failed"}, 500
@ -91,7 +91,7 @@ async def upload_file(file: UploadFile = File(...)):
if not file.filename:
raise HTTPException(status_code=400, detail="No selected file")
upload_dir = os.path.join(startup.config.save_dir, "data/uploads")
upload_dir = os.path.join(config.save_dir, "data/uploads")
os.makedirs(upload_dir, exist_ok=True)
basename, ext = os.path.splitext(file.filename)
filename = f"{basename}_{hashstr(basename, 4, with_salt=True)}{ext}".lower()
@ -104,13 +104,13 @@ async def upload_file(file: UploadFile = File(...)):
@data.get("/graph")
async def get_graph_info():
graph_info = startup.dbm.get_graph()
graph_info = dbm.get_graph()
# 获取未索引节点数量
unindexed_count = 0
if startup.dbm.is_graph_running():
if dbm.is_graph_running():
# 调用GraphDatabase的query_nodes_without_embedding方法
unindexed_nodes = startup.dbm.graph_base.query_nodes_without_embedding()
unindexed_nodes = dbm.graph_base.query_nodes_without_embedding()
unindexed_count = len(unindexed_nodes) if unindexed_nodes else 0
# 将未索引节点数量添加到返回结果中
@ -120,39 +120,39 @@ async def get_graph_info():
@data.post("/graph/index-nodes")
async def index_nodes(data: dict = Body(default={})):
if not startup.dbm.is_graph_running():
if not dbm.is_graph_running():
raise HTTPException(status_code=400, detail="图数据库未启动")
# 获取参数或使用默认值
kgdb_name = data.get('kgdb_name', 'neo4j')
# 调用GraphDatabase的add_embedding_to_nodes方法
count = startup.dbm.graph_base.add_embedding_to_nodes(kgdb_name=kgdb_name)
count = dbm.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 = startup.dbm.graph_base.query_node(entity_name=entity_name)
return {"result": startup.retriever.format_query_results(result), "message": "success"}
result = dbm.graph_base.query_node(entity_name=entity_name)
return {"result": retriever.format_query_results(result), "message": "success"}
@data.get("/graph/nodes")
async def get_graph_nodes(kgdb_name: str, num: int):
if not startup.config.enable_knowledge_graph:
if not config.enable_knowledge_graph:
raise HTTPException(status_code=400, detail="Knowledge graph is not enabled")
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
return {"result": startup.retriever.format_general_results(result), "message": "success"}
result = dbm.graph_base.get_sample_nodes(kgdb_name, num)
return {"result": retriever.format_general_results(result), "message": "success"}
@data.post("/graph/add-by-jsonl")
async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str] = Body(None)):
if not startup.config.enable_knowledge_graph:
if not config.enable_knowledge_graph:
raise HTTPException(status_code=400, detail="Knowledge graph is not enabled")
if not file_path.endswith('.jsonl'):
raise HTTPException(status_code=400, detail="file_path must be a jsonl file")
startup.dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name)
dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name)
return {"message": "Entity successfully added"}

View File

@ -278,8 +278,6 @@ const handleChange = (key, e) => {
|| key == 'enable_knowledge_graph'
|| key == 'enable_knowledge_base'
|| key == 'enable_web_search'
|| key == 'model_provider'
|| key == 'model_name'
|| key == 'embed_model'
|| key == 'reranker'
|| key == 'model_local_paths') {