优化启动项
This commit is contained in:
parent
e7a52f66df
commit
dd617c246d
@ -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()
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 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")
|
||||
|
||||
@ -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})
|
||||
|
||||
@ -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"}
|
||||
|
||||
|
||||
@ -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') {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user