优化启动项
This commit is contained in:
parent
e7a52f66df
commit
dd617c246d
@ -1,3 +1,15 @@
|
|||||||
from dotenv import load_dotenv
|
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()
|
||||||
@ -1,6 +1,8 @@
|
|||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
|
|
||||||
|
from src import config
|
||||||
from src.utils import hashstr, logger
|
from src.utils import hashstr, logger
|
||||||
from src.core.indexing import chunk
|
from src.core.indexing import chunk
|
||||||
from src.models.embedding import get_embedding_model
|
from src.models.embedding import get_embedding_model
|
||||||
@ -8,22 +10,23 @@ from src.models.embedding import get_embedding_model
|
|||||||
|
|
||||||
class DataBaseManager:
|
class DataBaseManager:
|
||||||
|
|
||||||
def __init__(self, config=None) -> None:
|
def __init__(self) -> None:
|
||||||
self.config = config
|
|
||||||
self.database_path = os.path.join(config.save_dir, "data", "database.json")
|
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
|
from src.core.knowledgebase import KnowledgeBase
|
||||||
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
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
|
from src.core.graphbase import GraphDatabase
|
||||||
self.graph_base = GraphDatabase(self.config, self.embed_model)
|
self.graph_base = GraphDatabase(config, self.embed_model)
|
||||||
else:
|
else:
|
||||||
self.graph_base = None
|
self.graph_base = None
|
||||||
|
|
||||||
self.data = {"databases": [], "graph": {}}
|
self.data = {"databases": [], "graph": {}}
|
||||||
|
|
||||||
self._load_databases()
|
self._load_databases()
|
||||||
self._update_database()
|
self._update_database()
|
||||||
|
|
||||||
@ -62,7 +65,7 @@ class DataBaseManager:
|
|||||||
|
|
||||||
def get_databases(self):
|
def get_databases(self):
|
||||||
self._update_database()
|
self._update_database()
|
||||||
assert self.config.enable_knowledge_base, "知识库未启用"
|
assert config.enable_knowledge_base, "知识库未启用"
|
||||||
knowledge_base_collections = self.knowledge_base.get_collection_names()
|
knowledge_base_collections = self.knowledge_base.get_collection_names()
|
||||||
if len(self.data["databases"]) != len(knowledge_base_collections):
|
if len(self.data["databases"]) != len(knowledge_base_collections):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@ -82,7 +85,7 @@ class DataBaseManager:
|
|||||||
return {"databases": [db.to_dict() for db in self.data["databases"]]}
|
return {"databases": [db.to_dict() for db in self.data["databases"]]}
|
||||||
|
|
||||||
def get_graph(self):
|
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"))
|
self.data["graph"].update(self.graph_base.get_database_info("neo4j"))
|
||||||
return {"graph": self.data["graph"]}
|
return {"graph": self.data["graph"]}
|
||||||
else:
|
else:
|
||||||
@ -95,7 +98,7 @@ class DataBaseManager:
|
|||||||
bool: 图数据库是否正在运行
|
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
|
return False
|
||||||
|
|
||||||
# 获取图数据库信息,检查状态
|
# 获取图数据库信息,检查状态
|
||||||
@ -104,12 +107,12 @@ class DataBaseManager:
|
|||||||
|
|
||||||
def create_database(self, database_name, description, db_type, dimension):
|
def create_database(self, database_name, description, db_type, dimension):
|
||||||
from src.config import EMBED_MODEL_INFO
|
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,
|
new_database = DataBaseLite(database_name,
|
||||||
description,
|
description,
|
||||||
db_type,
|
db_type,
|
||||||
embed_model=self.config.embed_model,
|
embed_model=config.embed_model,
|
||||||
dimension=dimension)
|
dimension=dimension)
|
||||||
|
|
||||||
self.knowledge_base.add_collection(new_database.metaname, 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):
|
def add_files(self, db_id, files, params=None):
|
||||||
db = self.get_kb_by_id(db_id)
|
db = self.get_kb_by_id(db_id)
|
||||||
|
|
||||||
if db.embed_model != self.config.embed_model:
|
if db.embed_model != config.embed_model:
|
||||||
logger.error(f"Embed model not match, {db.embed_model} != {self.config.embed_model}")
|
logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}")
|
||||||
return {"message": f"Embed model not match, cur: {self.config.embed_model}", "status": "failed"}
|
return {"message": f"Embed model not match, cur: {config.embed_model}", "status": "failed"}
|
||||||
|
|
||||||
# Preprocessing the files to the queue
|
# Preprocessing the files to the queue
|
||||||
new_files = []
|
new_files = []
|
||||||
@ -254,6 +257,11 @@ class DataBaseManager:
|
|||||||
if f["file_id"] == file_id:
|
if f["file_id"] == file_id:
|
||||||
return idx
|
return idx
|
||||||
|
|
||||||
|
def restart(self):
|
||||||
|
self.embed_model = get_embedding_model(config)
|
||||||
|
self._load_databases()
|
||||||
|
self._update_database()
|
||||||
|
|
||||||
|
|
||||||
class DataBaseLite:
|
class DataBaseLite:
|
||||||
def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None:
|
def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None:
|
||||||
|
|||||||
@ -1,24 +1,24 @@
|
|||||||
|
from src import config, dbm
|
||||||
from src.models.rerank_model import get_reranker
|
from src.models.rerank_model import get_reranker
|
||||||
from src.utils.logging_config import logger
|
from src.utils.logging_config import logger
|
||||||
|
from src.models import select_model
|
||||||
|
|
||||||
class Retriever:
|
class Retriever:
|
||||||
|
|
||||||
def __init__(self, config, dbm, model):
|
def __init__(self):
|
||||||
self.config = config
|
self._load_models()
|
||||||
self.dbm = dbm
|
|
||||||
self.model = model
|
|
||||||
|
|
||||||
if self.config.enable_reranker:
|
def _load_models(self):
|
||||||
|
if config.enable_reranker:
|
||||||
self.reranker = get_reranker(config)
|
self.reranker = get_reranker(config)
|
||||||
|
|
||||||
if self.config.enable_web_search:
|
if config.enable_web_search:
|
||||||
from src.utils.web_search import WebSearcher
|
from src.utils.web_search import WebSearcher
|
||||||
self.web_searcher = WebSearcher()
|
self.web_searcher = WebSearcher()
|
||||||
|
|
||||||
def retrieval(self, query, history, meta):
|
def retrieval(self, query, history, meta):
|
||||||
refs = {"query": query, "history": history, "meta": 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["entities"] = self.reco_entities(query, history, refs)
|
||||||
refs["knowledge_base"] = self.query_knowledgebase(query, history, refs)
|
refs["knowledge_base"] = self.query_knowledgebase(query, history, refs)
|
||||||
refs["graph_base"] = self.query_graph(query, history, refs)
|
refs["graph_base"] = self.query_graph(query, history, refs)
|
||||||
@ -26,6 +26,10 @@ class Retriever:
|
|||||||
|
|
||||||
return refs
|
return refs
|
||||||
|
|
||||||
|
def restart(self):
|
||||||
|
"""所有需要重启的模型"""
|
||||||
|
self._load_models()
|
||||||
|
|
||||||
def construct_query(self, query, refs, meta):
|
def construct_query(self, query, refs, meta):
|
||||||
logger.debug(f"{refs=}")
|
logger.debug(f"{refs=}")
|
||||||
if not refs or len(refs) == 0:
|
if not refs or len(refs) == 0:
|
||||||
@ -70,9 +74,9 @@ class Retriever:
|
|||||||
|
|
||||||
def query_graph(self, query, history, refs):
|
def query_graph(self, query, history, refs):
|
||||||
results = []
|
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"]:
|
for entity in refs["entities"]:
|
||||||
result = self.dbm.graph_base.query_by_vector(entity)
|
result = dbm.graph_base.query_by_vector(entity)
|
||||||
if result != []:
|
if result != []:
|
||||||
results.extend(result)
|
results.extend(result)
|
||||||
return {"results": self.format_query_results(results)}
|
return {"results": self.format_query_results(results)}
|
||||||
@ -85,7 +89,7 @@ class Retriever:
|
|||||||
final_res = []
|
final_res = []
|
||||||
|
|
||||||
db_name = refs["meta"].get("db_name")
|
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 {
|
return {
|
||||||
"results": final_res,
|
"results": final_res,
|
||||||
"all_results": kb_res,
|
"all_results": kb_res,
|
||||||
@ -95,7 +99,7 @@ class Retriever:
|
|||||||
|
|
||||||
rw_query = self.rewrite_query(query, history, refs)
|
rw_query = self.rewrite_query(query, history, refs)
|
||||||
|
|
||||||
kb = self.dbm.metaname2db[db_name]
|
kb = dbm.metaname2db[db_name]
|
||||||
logger.debug(f"{refs['meta']=}")
|
logger.debug(f"{refs['meta']=}")
|
||||||
|
|
||||||
meta = refs["meta"]
|
meta = refs["meta"]
|
||||||
@ -105,14 +109,14 @@ class Retriever:
|
|||||||
top_k = meta.get("topK", 5)
|
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:
|
for r in all_kb_res:
|
||||||
r["file"] = kb.id2file(r["entity"]["file_id"])
|
r["file"] = kb.id2file(r["entity"]["file_id"])
|
||||||
|
|
||||||
kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold]
|
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]
|
texts = [r["entity"]["text"] for r in kb_res]
|
||||||
rerank_scores = self.reranker.compute_score([rw_query, texts], normalize=True)
|
rerank_scores = self.reranker.compute_score([rw_query, texts], normalize=True)
|
||||||
for i, r in enumerate(kb_res):
|
for i, r in enumerate(kb_res):
|
||||||
@ -127,7 +131,7 @@ class Retriever:
|
|||||||
def query_web(self, query, history, refs):
|
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"}
|
return {"results": [], "message": "Web search is disabled"}
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@ -140,10 +144,11 @@ class Retriever:
|
|||||||
|
|
||||||
def rewrite_query(self, query, history, refs):
|
def rewrite_query(self, query, history, refs):
|
||||||
"""重写查询"""
|
"""重写查询"""
|
||||||
|
model = select_model(config)
|
||||||
if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置
|
if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置
|
||||||
rewrite_query_span = refs["meta"].get("use_rewrite_query", "off")
|
rewrite_query_span = refs["meta"].get("use_rewrite_query", "off")
|
||||||
else:
|
else:
|
||||||
rewrite_query_span = self.config.use_rewrite_query
|
rewrite_query_span = config.use_rewrite_query
|
||||||
|
|
||||||
if rewrite_query_span == "off":
|
if rewrite_query_span == "off":
|
||||||
rewritten_query = query
|
rewritten_query = query
|
||||||
@ -152,10 +157,10 @@ class Retriever:
|
|||||||
|
|
||||||
history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else ""
|
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_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":
|
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}"
|
rewritten_query = f"{rewritten_query} {hy_doc}"
|
||||||
|
|
||||||
return rewritten_query
|
return rewritten_query
|
||||||
@ -163,6 +168,7 @@ class Retriever:
|
|||||||
def reco_entities(self, query, history, refs):
|
def reco_entities(self, query, history, refs):
|
||||||
"""识别句子中的实体"""
|
"""识别句子中的实体"""
|
||||||
query = refs.get("rewritten_query", query)
|
query = refs.get("rewritten_query", query)
|
||||||
|
model = select_model(config)
|
||||||
|
|
||||||
entities = []
|
entities = []
|
||||||
if refs["meta"].get("use_graph"):
|
if refs["meta"].get("use_graph"):
|
||||||
@ -170,7 +176,7 @@ class Retriever:
|
|||||||
from src.utils.prompts import keywords_prompt_template as entity_template
|
from src.utils.prompts import keywords_prompt_template as entity_template
|
||||||
|
|
||||||
entity_extraction_prompt = entity_template.format(text=query)
|
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)]
|
# entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)]
|
||||||
|
|
||||||
return entities
|
return entities
|
||||||
|
|||||||
@ -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()
|
|
||||||
@ -1,13 +1,9 @@
|
|||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
|
from fastapi import Request, Body
|
||||||
|
|
||||||
base = APIRouter()
|
base = APIRouter()
|
||||||
|
|
||||||
from fastapi import FastAPI, HTTPException
|
from src import config, dbm, retriever
|
||||||
from fastapi.responses import JSONResponse
|
|
||||||
from fastapi import Request, Body
|
|
||||||
|
|
||||||
from src.core import HistoryManager
|
|
||||||
from src.core.startup import startup
|
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
@ -17,17 +13,18 @@ async def route_index():
|
|||||||
|
|
||||||
@base.get("/config")
|
@base.get("/config")
|
||||||
def get_config():
|
def get_config():
|
||||||
return startup.config
|
return config
|
||||||
|
|
||||||
@base.post("/config")
|
@base.post("/config")
|
||||||
async def update_config(key = Body(...), value = Body(...)):
|
async def update_config(key = Body(...), value = Body(...)):
|
||||||
startup.config[key] = value
|
config[key] = value
|
||||||
startup.config.save()
|
config.save()
|
||||||
return startup.config
|
return config
|
||||||
|
|
||||||
@base.post("/restart")
|
@base.post("/restart")
|
||||||
async def restart():
|
async def restart():
|
||||||
startup.restart()
|
dbm.restart()
|
||||||
|
retriever.restart()
|
||||||
return {"message": "Restarted!"}
|
return {"message": "Restarted!"}
|
||||||
|
|
||||||
@base.get("/log")
|
@base.get("/log")
|
||||||
|
|||||||
@ -3,7 +3,8 @@ import asyncio
|
|||||||
from fastapi import APIRouter, Body
|
from fastapi import APIRouter, Body
|
||||||
from fastapi.responses import StreamingResponse, Response
|
from fastapi.responses import StreamingResponse, Response
|
||||||
from src.core import HistoryManager
|
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
|
from src.utils.logging_config import logger
|
||||||
|
|
||||||
chat = APIRouter(prefix="/chat")
|
chat = APIRouter(prefix="/chat")
|
||||||
@ -21,14 +22,15 @@ def chat_post(
|
|||||||
history: list = Body(...),
|
history: list = Body(...),
|
||||||
cur_res_id: str = 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)
|
history_manager = HistoryManager(history)
|
||||||
logger.debug(f"Received query: {query} with meta: {meta}")
|
logger.debug(f"Received query: {query} with meta: {meta}")
|
||||||
|
|
||||||
def make_chunk(content=None, **kwargs):
|
def make_chunk(content=None, **kwargs):
|
||||||
return json.dumps({
|
return json.dumps({
|
||||||
"response": content,
|
"response": content,
|
||||||
"model_name": startup.config.model_name,
|
"model_name": config.model_name,
|
||||||
"meta": meta,
|
"meta": meta,
|
||||||
**kwargs
|
**kwargs
|
||||||
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
||||||
@ -46,7 +48,7 @@ def chat_post(
|
|||||||
yield chunk
|
yield chunk
|
||||||
|
|
||||||
try:
|
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:
|
except Exception as e:
|
||||||
logger.error(f"Retriever error: {e}")
|
logger.error(f"Retriever error: {e}")
|
||||||
yield make_chunk(message=f"Retriever error: {e}", status="error")
|
yield make_chunk(message=f"Retriever error: {e}", status="error")
|
||||||
@ -60,7 +62,7 @@ def chat_post(
|
|||||||
content = ""
|
content = ""
|
||||||
reasoning_content = ""
|
reasoning_content = ""
|
||||||
try:
|
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'):
|
if not delta.content and hasattr(delta, 'reasoning_content'):
|
||||||
reasoning_content += delta.reasoning_content or ""
|
reasoning_content += delta.reasoning_content or ""
|
||||||
chunk = make_chunk(reasoning_content=reasoning_content, status="reasoning")
|
chunk = make_chunk(reasoning_content=reasoning_content, status="reasoning")
|
||||||
@ -91,9 +93,10 @@ def chat_post(
|
|||||||
|
|
||||||
@chat.post("/call")
|
@chat.post("/call")
|
||||||
async def call(query: str = Body(...), meta: dict = Body(None)):
|
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):
|
async def predict_async(query):
|
||||||
loop = asyncio.get_event_loop()
|
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)
|
response = await predict_async(query)
|
||||||
logger.debug({"query": query, "response": response.content})
|
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 call(query: str = Body(...), meta: dict = Body(None)):
|
||||||
async def predict_async(query):
|
async def predict_async(query):
|
||||||
loop = asyncio.get_event_loop()
|
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)
|
response = await predict_async(query)
|
||||||
logger.debug({"query": query, "response": response.content})
|
logger.debug({"query": query, "response": response.content})
|
||||||
|
|||||||
@ -4,7 +4,7 @@ from typing import List, Optional
|
|||||||
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body
|
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body
|
||||||
|
|
||||||
from src.utils import logger, hashstr
|
from src.utils import logger, hashstr
|
||||||
from src.core.startup import startup, executor
|
from src import executor, dbm, retriever, config
|
||||||
|
|
||||||
data = APIRouter(prefix="/data")
|
data = APIRouter(prefix="/data")
|
||||||
|
|
||||||
@ -12,7 +12,7 @@ data = APIRouter(prefix="/data")
|
|||||||
@data.get("/")
|
@data.get("/")
|
||||||
async def get_databases():
|
async def get_databases():
|
||||||
try:
|
try:
|
||||||
database = startup.dbm.get_databases()
|
database = dbm.get_databases()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return {"message": f"获取数据库列表失败 {e}", "databases": []}
|
return {"message": f"获取数据库列表失败 {e}", "databases": []}
|
||||||
return database
|
return database
|
||||||
@ -25,7 +25,7 @@ async def create_database(
|
|||||||
dimension: Optional[int] = Body(None)
|
dimension: Optional[int] = Body(None)
|
||||||
):
|
):
|
||||||
logger.debug(f"Create database {database_name}")
|
logger.debug(f"Create database {database_name}")
|
||||||
database_info = startup.dbm.create_database(
|
database_info = dbm.create_database(
|
||||||
database_name,
|
database_name,
|
||||||
description,
|
description,
|
||||||
db_type,
|
db_type,
|
||||||
@ -36,13 +36,13 @@ async def create_database(
|
|||||||
@data.delete("/")
|
@data.delete("/")
|
||||||
async def delete_database(db_id):
|
async def delete_database(db_id):
|
||||||
logger.debug(f"Delete database {db_id}")
|
logger.debug(f"Delete database {db_id}")
|
||||||
startup.dbm.delete_database(db_id)
|
dbm.delete_database(db_id)
|
||||||
return {"message": "删除成功"}
|
return {"message": "删除成功"}
|
||||||
|
|
||||||
@data.post("/query-test")
|
@data.post("/query-test")
|
||||||
async def query_test(query: str = Body(...), meta: dict = Body(...)):
|
async def query_test(query: str = Body(...), meta: dict = Body(...)):
|
||||||
logger.debug(f"Query test in {meta}: {query}")
|
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
|
return result
|
||||||
|
|
||||||
@data.post("/add-by-file")
|
@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()
|
loop = asyncio.get_event_loop()
|
||||||
await loop.run_in_executor(
|
await loop.run_in_executor(
|
||||||
executor, # 使用与chat_router相同的线程池
|
executor, # 使用与chat_router相同的线程池
|
||||||
lambda: startup.dbm.add_files(db_id, files)
|
lambda: dbm.add_files(db_id, files)
|
||||||
)
|
)
|
||||||
return {"message": "文件添加完成", "status": "success"}
|
return {"message": "文件添加完成", "status": "success"}
|
||||||
except Exception as e:
|
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")
|
@data.get("/info")
|
||||||
async def get_database_info(db_id: str):
|
async def get_database_info(db_id: str):
|
||||||
logger.debug(f"Get database {db_id} info")
|
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:
|
if database is None:
|
||||||
raise HTTPException(status_code=404, detail="Database not found")
|
raise HTTPException(status_code=404, detail="Database not found")
|
||||||
return database
|
return database
|
||||||
@ -71,7 +71,7 @@ async def get_database_info(db_id: str):
|
|||||||
@data.delete("/document")
|
@data.delete("/document")
|
||||||
async def delete_document(db_id: str = Body(...), file_id: str = Body(...)):
|
async def delete_document(db_id: str = Body(...), file_id: str = Body(...)):
|
||||||
logger.debug(f"DELETE document {file_id} info in {db_id}")
|
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": "删除成功"}
|
return {"message": "删除成功"}
|
||||||
|
|
||||||
@data.get("/document")
|
@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}")
|
logger.debug(f"GET document {file_id} info in {db_id}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
info = startup.dbm.get_file_info(db_id, file_id)
|
info = dbm.get_file_info(db_id, file_id)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to get file info, {e}, {db_id=}, {file_id=}")
|
logger.error(f"Failed to get file info, {e}, {db_id=}, {file_id=}")
|
||||||
info = {"message": "Failed to get file info", "status": "failed"}, 500
|
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:
|
if not file.filename:
|
||||||
raise HTTPException(status_code=400, detail="No selected file")
|
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)
|
os.makedirs(upload_dir, exist_ok=True)
|
||||||
basename, ext = os.path.splitext(file.filename)
|
basename, ext = os.path.splitext(file.filename)
|
||||||
filename = f"{basename}_{hashstr(basename, 4, with_salt=True)}{ext}".lower()
|
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")
|
@data.get("/graph")
|
||||||
async def get_graph_info():
|
async def get_graph_info():
|
||||||
graph_info = startup.dbm.get_graph()
|
graph_info = dbm.get_graph()
|
||||||
|
|
||||||
# 获取未索引节点数量
|
# 获取未索引节点数量
|
||||||
unindexed_count = 0
|
unindexed_count = 0
|
||||||
if startup.dbm.is_graph_running():
|
if dbm.is_graph_running():
|
||||||
# 调用GraphDatabase的query_nodes_without_embedding方法
|
# 调用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
|
unindexed_count = len(unindexed_nodes) if unindexed_nodes else 0
|
||||||
|
|
||||||
# 将未索引节点数量添加到返回结果中
|
# 将未索引节点数量添加到返回结果中
|
||||||
@ -120,39 +120,39 @@ async def get_graph_info():
|
|||||||
|
|
||||||
@data.post("/graph/index-nodes")
|
@data.post("/graph/index-nodes")
|
||||||
async def index_nodes(data: dict = Body(default={})):
|
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="图数据库未启动")
|
raise HTTPException(status_code=400, detail="图数据库未启动")
|
||||||
|
|
||||||
# 获取参数或使用默认值
|
# 获取参数或使用默认值
|
||||||
kgdb_name = data.get('kgdb_name', 'neo4j')
|
kgdb_name = data.get('kgdb_name', 'neo4j')
|
||||||
|
|
||||||
# 调用GraphDatabase的add_embedding_to_nodes方法
|
# 调用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}
|
return {"status": "success", "message": f"已成功为{count}个节点添加嵌入向量", "indexed_count": count}
|
||||||
|
|
||||||
@data.get("/graph/node")
|
@data.get("/graph/node")
|
||||||
async def get_graph_node(entity_name: str):
|
async def get_graph_node(entity_name: str):
|
||||||
result = startup.dbm.graph_base.query_node(entity_name=entity_name)
|
result = dbm.graph_base.query_node(entity_name=entity_name)
|
||||||
return {"result": startup.retriever.format_query_results(result), "message": "success"}
|
return {"result": retriever.format_query_results(result), "message": "success"}
|
||||||
|
|
||||||
@data.get("/graph/nodes")
|
@data.get("/graph/nodes")
|
||||||
async def get_graph_nodes(kgdb_name: str, num: int):
|
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")
|
raise HTTPException(status_code=400, detail="Knowledge graph is not enabled")
|
||||||
|
|
||||||
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
||||||
result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
result = dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
||||||
return {"result": startup.retriever.format_general_results(result), "message": "success"}
|
return {"result": retriever.format_general_results(result), "message": "success"}
|
||||||
|
|
||||||
@data.post("/graph/add-by-jsonl")
|
@data.post("/graph/add-by-jsonl")
|
||||||
async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str] = Body(None)):
|
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")
|
raise HTTPException(status_code=400, detail="Knowledge graph is not enabled")
|
||||||
|
|
||||||
if not file_path.endswith('.jsonl'):
|
if not file_path.endswith('.jsonl'):
|
||||||
raise HTTPException(status_code=400, detail="file_path must be a jsonl file")
|
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"}
|
return {"message": "Entity successfully added"}
|
||||||
|
|
||||||
|
|||||||
@ -278,8 +278,6 @@ const handleChange = (key, e) => {
|
|||||||
|| key == 'enable_knowledge_graph'
|
|| key == 'enable_knowledge_graph'
|
||||||
|| key == 'enable_knowledge_base'
|
|| key == 'enable_knowledge_base'
|
||||||
|| key == 'enable_web_search'
|
|| key == 'enable_web_search'
|
||||||
|| key == 'model_provider'
|
|
||||||
|| key == 'model_name'
|
|
||||||
|| key == 'embed_model'
|
|| key == 'embed_model'
|
||||||
|| key == 'reranker'
|
|| key == 'reranker'
|
||||||
|| key == 'model_local_paths') {
|
|| key == 'model_local_paths') {
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user