diff --git a/src/__init__.py b/src/__init__.py index 20fafbef..acf8ed3e 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -1,3 +1,15 @@ from dotenv import load_dotenv -load_dotenv("src/.env") \ No newline at end of file +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() \ No newline at end of file diff --git a/src/core/database.py b/src/core/database.py index 2a8936fa..1f775c37 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -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: diff --git a/src/core/retriever.py b/src/core/retriever.py index 55d21a9b..92444f1d 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -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 diff --git a/src/core/startup.py b/src/core/startup.py deleted file mode 100644 index f89e263d..00000000 --- a/src/core/startup.py +++ /dev/null @@ -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() \ No newline at end of file diff --git a/src/routers/base_router.py b/src/routers/base_router.py index 0dcbbde7..8e53cb20 100644 --- a/src/routers/base_router.py +++ b/src/routers/base_router.py @@ -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") diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index 35aa5157..5b9b49c7 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -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}) diff --git a/src/routers/data_router.py b/src/routers/data_router.py index d2bb3417..6a5229ff 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -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"} diff --git a/web/src/views/SettingView.vue b/web/src/views/SettingView.vue index 522770ca..3813018f 100644 --- a/web/src/views/SettingView.vue +++ b/web/src/views/SettingView.vue @@ -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') {