modified config embedding model and startup
This commit is contained in:
parent
590dfd1836
commit
3ab20611a2
@ -31,8 +31,7 @@ class Config(SimpleConfig):
|
|||||||
|
|
||||||
def __init__(self, filename=None):
|
def __init__(self, filename=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.filename = filename
|
self.filename = filename or "config/base.yaml"
|
||||||
logger.info(f"Loading config from {filename}")
|
|
||||||
|
|
||||||
### >>> 默认配置
|
### >>> 默认配置
|
||||||
# 可以在 config/base.yaml 中覆盖
|
# 可以在 config/base.yaml 中覆盖
|
||||||
@ -48,6 +47,8 @@ class Config(SimpleConfig):
|
|||||||
# 模型配置
|
# 模型配置
|
||||||
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
||||||
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
|
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
|
||||||
|
self.model_provider = "qianfan"
|
||||||
|
self.model_name = None # 默认使用 provider 的默认模型
|
||||||
self.embed_model = "bge-large-zh-v1.5"
|
self.embed_model = "bge-large-zh-v1.5"
|
||||||
self.reranker = "bge-reranker-v2-m3"
|
self.reranker = "bge-reranker-v2-m3"
|
||||||
### <<< 默认配置结束
|
### <<< 默认配置结束
|
||||||
@ -66,6 +67,7 @@ class Config(SimpleConfig):
|
|||||||
|
|
||||||
def load(self):
|
def load(self):
|
||||||
"""根据传入的文件覆盖掉默认配置"""
|
"""根据传入的文件覆盖掉默认配置"""
|
||||||
|
logger.info(f"Loading config from {self.filename}")
|
||||||
if self.filename is not None and os.path.exists(self.filename):
|
if self.filename is not None and os.path.exists(self.filename):
|
||||||
if self.filename.endswith(".json"):
|
if self.filename.endswith(".json"):
|
||||||
with open(self.filename, 'r') as f:
|
with open(self.filename, 'r') as f:
|
||||||
@ -77,7 +79,20 @@ class Config(SimpleConfig):
|
|||||||
logger.warning(f"Config file {self.filename} not found")
|
logger.warning(f"Config file {self.filename} not found")
|
||||||
|
|
||||||
def save(self):
|
def save(self):
|
||||||
if self.filename is not None:
|
logger.info(f"Saving config to {self.filename}")
|
||||||
with open(self.filename, 'w') as f:
|
if self.filename is None:
|
||||||
|
logger.warning("Config file is not specified, save to default config/base.yaml")
|
||||||
|
self.filename = "config/base.yaml"
|
||||||
|
|
||||||
|
if self.filename.endswith(".json"):
|
||||||
|
with open(self.filename, 'w+') as f:
|
||||||
|
json.dump(self, f, indent=4, ensure_ascii=False)
|
||||||
|
elif self.filename.endswith(".yaml"):
|
||||||
|
with open(self.filename, 'w+') as f:
|
||||||
|
yaml.dump(self, f, indent=2)
|
||||||
|
else:
|
||||||
|
logger.warning(f"Unknown config file type {self.filename}, save as json")
|
||||||
|
with open(self.filename, 'w+') as f:
|
||||||
json.dump(self, f, indent=4)
|
json.dump(self, f, indent=4)
|
||||||
logger.info(f"Config file {self.filename} saved")
|
|
||||||
|
logger.info(f"Config file {self.filename} saved")
|
||||||
@ -6,7 +6,7 @@ from plugins import pdf2txt
|
|||||||
from core.knowledgebase import KnowledgeBase
|
from core.knowledgebase import KnowledgeBase
|
||||||
from core.filereader import pdfreader, plainreader
|
from core.filereader import pdfreader, plainreader
|
||||||
from core.graphbase import GraphDatabase
|
from core.graphbase import GraphDatabase
|
||||||
from models.embedding import EmbeddingModel
|
from models.embedding import get_embedding_model
|
||||||
|
|
||||||
logger = setup_logger("DataBaseManager")
|
logger = setup_logger("DataBaseManager")
|
||||||
|
|
||||||
@ -17,9 +17,10 @@ class DataBaseLite:
|
|||||||
self.description = description
|
self.description = description
|
||||||
self.db_type = db_type
|
self.db_type = db_type
|
||||||
self.db_id = kwargs.get("db_id", hashstr(name))
|
self.db_id = kwargs.get("db_id", hashstr(name))
|
||||||
self.metaname = kwargs.get("metaname", f"{db_type}_{hashstr(name)}")
|
self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}")
|
||||||
self.metadata = kwargs.get("metaname", {})
|
self.metadata = kwargs.get("metaname", {})
|
||||||
self.files = kwargs.get("files", [])
|
self.files = kwargs.get("files", [])
|
||||||
|
self.embed_model = kwargs.get("embed_model", None)
|
||||||
|
|
||||||
|
|
||||||
def update(self, metadata):
|
def update(self, metadata):
|
||||||
@ -31,6 +32,7 @@ class DataBaseLite:
|
|||||||
"description": self.description,
|
"description": self.description,
|
||||||
"db_type": self.db_type,
|
"db_type": self.db_type,
|
||||||
"db_id": self.db_id,
|
"db_id": self.db_id,
|
||||||
|
"embed_model": self.embed_model,
|
||||||
"metaname": self.metaname,
|
"metaname": self.metaname,
|
||||||
"metadata": self.metadata,
|
"metadata": self.metadata,
|
||||||
"files": self.files
|
"files": self.files
|
||||||
@ -47,7 +49,7 @@ class DataBaseManager:
|
|||||||
def __init__(self, config=None) -> None:
|
def __init__(self, config=None) -> None:
|
||||||
self.config = config
|
self.config = config
|
||||||
self.database_path = "data/databases.json"
|
self.database_path = "data/databases.json"
|
||||||
self.embed_model = EmbeddingModel(config)
|
self.embed_model = get_embedding_model(config)
|
||||||
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
||||||
self.data = {"databases": [], "graph": {}}
|
self.data = {"databases": [], "graph": {}}
|
||||||
|
|
||||||
@ -96,7 +98,7 @@ class DataBaseManager:
|
|||||||
return {"graph": {}, "message": "Graph database is not enabled"}
|
return {"graph": {}, "message": "Graph database is not enabled"}
|
||||||
|
|
||||||
def create_database(self, database_name, description, db_type):
|
def create_database(self, database_name, description, db_type):
|
||||||
new_database = DataBaseLite(database_name, description, db_type)
|
new_database = DataBaseLite(database_name, description, db_type, embed_model=self.config.embed_model)
|
||||||
|
|
||||||
self.knowledge_base.add_collection(new_database.metaname)
|
self.knowledge_base.add_collection(new_database.metaname)
|
||||||
self.data["databases"].append(new_database)
|
self.data["databases"].append(new_database)
|
||||||
@ -105,6 +107,11 @@ class DataBaseManager:
|
|||||||
|
|
||||||
def add_files(self, db_id, files):
|
def add_files(self, db_id, files):
|
||||||
db = self.get_kb_by_id(db_id)
|
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": "Embed model not match", "status": "failed"}
|
||||||
|
|
||||||
new_files = []
|
new_files = []
|
||||||
for file in files:
|
for file in files:
|
||||||
# filenames = [f["filename"] for f in db.files]
|
# filenames = [f["filename"] for f in db.files]
|
||||||
@ -138,6 +145,8 @@ class DataBaseManager:
|
|||||||
|
|
||||||
self._save_databases()
|
self._save_databases()
|
||||||
|
|
||||||
|
return {"message": "全部解析完成", "status": "success"}
|
||||||
|
|
||||||
def get_database_info(self, db_id):
|
def get_database_info(self, db_id):
|
||||||
db = self.get_kb_by_id(db_id)
|
db = self.get_kb_by_id(db_id)
|
||||||
if db is None:
|
if db is None:
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
from core.startup import dbm, model
|
|
||||||
from models.embedding import Reranker
|
from models.embedding import Reranker
|
||||||
from utils.logging_config import setup_logger
|
from utils.logging_config import setup_logger
|
||||||
logger = setup_logger("server-common")
|
logger = setup_logger("server-common")
|
||||||
@ -6,9 +5,11 @@ logger = setup_logger("server-common")
|
|||||||
|
|
||||||
class Retriever:
|
class Retriever:
|
||||||
|
|
||||||
def __init__(self, config):
|
def __init__(self, config, dbm, model):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.reranker = Reranker(config)
|
self.reranker = Reranker(config)
|
||||||
|
self.dbm = dbm
|
||||||
|
self.model = model
|
||||||
|
|
||||||
def retrieval(self, query, history, meta):
|
def retrieval(self, query, history, meta):
|
||||||
|
|
||||||
@ -56,7 +57,7 @@ class Retriever:
|
|||||||
results = []
|
results = []
|
||||||
if meta.get("use_graph"):
|
if meta.get("use_graph"):
|
||||||
for entitie in entities:
|
for entitie in entities:
|
||||||
result = dbm.graph_base.query_by_vector(entitie)
|
result = self.dbm.graph_base.query_by_vector(entitie)
|
||||||
if result != []:
|
if result != []:
|
||||||
results.extend(result)
|
results.extend(result)
|
||||||
return {"results": self.format_query_results(results)}
|
return {"results": self.format_query_results(results)}
|
||||||
@ -65,7 +66,7 @@ class Retriever:
|
|||||||
|
|
||||||
kb_res = []
|
kb_res = []
|
||||||
if meta.get("db_name"):
|
if meta.get("db_name"):
|
||||||
kb_res = dbm.knowledge_base.search(query, meta["db_name"], limit=5)
|
kb_res = self.dbm.knowledge_base.search(query, meta["db_name"], limit=5)
|
||||||
for r in kb_res:
|
for r in kb_res:
|
||||||
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
|
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
|
||||||
|
|
||||||
@ -97,7 +98,7 @@ class Retriever:
|
|||||||
# 构建提示词
|
# 构建提示词
|
||||||
rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query)
|
rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query)
|
||||||
# 调用语言模型生成重写的查询(假设使用某个API)
|
# 调用语言模型生成重写的查询(假设使用某个API)
|
||||||
rewritten_query = model.predict(rewritten_query_prompt).content
|
rewritten_query = self.model.predict(rewritten_query_prompt).content
|
||||||
|
|
||||||
|
|
||||||
if meta.get("use_graph"):
|
if meta.get("use_graph"):
|
||||||
@ -113,7 +114,7 @@ class Retriever:
|
|||||||
"""
|
"""
|
||||||
# 构建提示词
|
# 构建提示词
|
||||||
entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query)
|
entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query)
|
||||||
entities = model.predict(entity_extraction_prompt).content.split(",")
|
entities = self.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)]
|
||||||
else:
|
else:
|
||||||
entities = []
|
entities = []
|
||||||
|
|||||||
@ -1,10 +1,25 @@
|
|||||||
from core import DataBaseManager
|
from core import DataBaseManager
|
||||||
|
from core.retriever import Retriever
|
||||||
from models import select_model
|
from models import select_model
|
||||||
from config import Config
|
from config import Config
|
||||||
|
from utils import setup_logger
|
||||||
|
|
||||||
|
logger = setup_logger("Startup")
|
||||||
|
|
||||||
|
|
||||||
config = Config("config/base.yaml")
|
class Startup:
|
||||||
model = select_model(config)
|
def __init__(self):
|
||||||
dbm = DataBaseManager(config)
|
self.config = Config("config/base.yaml")
|
||||||
|
self.model = select_model(self.config)
|
||||||
|
self.dbm = DataBaseManager(self.config)
|
||||||
|
self.retriever = Retriever(self.config, self.dbm, self.model)
|
||||||
|
|
||||||
# 启动本地图数据库
|
def restart(self):
|
||||||
|
logger.info("Restarting...")
|
||||||
|
self.model = select_model(self.config)
|
||||||
|
self.dbm = DataBaseManager(self.config)
|
||||||
|
self.retriever = Retriever(self.config, self.dbm, self.model)
|
||||||
|
logger.info("Restarted")
|
||||||
|
|
||||||
|
|
||||||
|
startup = Startup()
|
||||||
@ -6,7 +6,7 @@ def select_model(config):
|
|||||||
model_provider = config.model_provider
|
model_provider = config.model_provider
|
||||||
model_name = config.model_name
|
model_name = config.model_name
|
||||||
|
|
||||||
logger.info(f"Selecting model from {model_provider} with name {model_name}")
|
logger.info(f"Selecting model from {model_provider} with {model_name or 'default'}")
|
||||||
|
|
||||||
if model_provider == "deepseek":
|
if model_provider == "deepseek":
|
||||||
from models.chat_model import DeepSeek
|
from models.chat_model import DeepSeek
|
||||||
|
|||||||
@ -40,31 +40,39 @@ class Reranker(FlagReranker):
|
|||||||
assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}"
|
assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}"
|
||||||
|
|
||||||
model_name_or_path = config.model_local_paths.get(config.reranker, RERANKER_LIST[config.reranker])
|
model_name_or_path = config.model_local_paths.get(config.reranker, RERANKER_LIST[config.reranker])
|
||||||
logger.info(f"Loading Reranker model {config.re_ranker} from {model_name_or_path}")
|
logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}")
|
||||||
|
|
||||||
super().__init__(model_name_or_path, use_fp16=True, **kwargs)
|
super().__init__(model_name_or_path, use_fp16=True, **kwargs)
|
||||||
logger.info(f"Reranker model {config.re_ranker} loaded")
|
logger.info(f"Reranker model {config.reranker} loaded")
|
||||||
|
|
||||||
|
|
||||||
from zhipuai import ZhipuAI
|
from zhipuai import ZhipuAI
|
||||||
|
|
||||||
client = ZhipuAI(api_key="270ea71e9560c0ff406acbcdd48bfd97.e3XOMdWKuZb7Q1Sk")
|
|
||||||
response = client.embeddings.create(
|
|
||||||
model="embedding-2", #填写需要调用的模型名称
|
|
||||||
input=["你好","woshi"]
|
|
||||||
)
|
|
||||||
|
|
||||||
print(response.data.shape)
|
|
||||||
|
|
||||||
class ZhipuEmbedding:
|
class ZhipuEmbedding:
|
||||||
|
|
||||||
def __init__(self, config) -> None:
|
def __init__(self, config) -> None:
|
||||||
self.config = config
|
self.config = config
|
||||||
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI"))
|
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI"))
|
||||||
|
logger.info("Zhipu Embedding model loaded")
|
||||||
|
self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:"
|
||||||
|
|
||||||
def predict(self, message):
|
def predict(self, message):
|
||||||
response = self.client.embeddings.create(
|
response = self.client.embeddings.create(
|
||||||
model=SUPPORT_LIST[self.config.embed_model],
|
model=SUPPORT_LIST[self.config.embed_model],
|
||||||
input=message
|
input=message
|
||||||
)
|
)
|
||||||
return response.data
|
return [a["embedding"] for a in response["data"]]
|
||||||
|
|
||||||
|
def encode(self, message):
|
||||||
|
return self.predict(message)
|
||||||
|
|
||||||
|
def encode_queries(self, queries):
|
||||||
|
# queries = [self.query_instruction_for_retrieval + query for query in queries]
|
||||||
|
return self.predict(queries)
|
||||||
|
|
||||||
|
|
||||||
|
def get_embedding_model(config):
|
||||||
|
if config.embed_model == "zhipu":
|
||||||
|
return ZhipuEmbedding(config)
|
||||||
|
else:
|
||||||
|
return EmbeddingModel(config)
|
||||||
@ -3,13 +3,10 @@ from flask import Blueprint, jsonify, request, Response
|
|||||||
|
|
||||||
from core import HistoryManager
|
from core import HistoryManager
|
||||||
from utils.logging_config import setup_logger
|
from utils.logging_config import setup_logger
|
||||||
from core.startup import config, model
|
from core.startup import startup
|
||||||
from core.retriever import Retriever
|
|
||||||
|
|
||||||
|
|
||||||
common = Blueprint('common', __name__)
|
common = Blueprint('common', __name__)
|
||||||
logger = setup_logger("server-common")
|
logger = setup_logger("server-common")
|
||||||
retriever = Retriever(config)
|
|
||||||
|
|
||||||
@common.route('/', methods=["GET"])
|
@common.route('/', methods=["GET"])
|
||||||
def route_index():
|
def route_index():
|
||||||
@ -35,7 +32,7 @@ def chat():
|
|||||||
logger.debug(f"Web query: {query}")
|
logger.debug(f"Web query: {query}")
|
||||||
history_manager = HistoryManager(request_data['history'])
|
history_manager = HistoryManager(request_data['history'])
|
||||||
|
|
||||||
new_query, refs = retriever(query, history_manager.messages, meta)
|
new_query, refs = startup.retriever(query, history_manager.messages, meta)
|
||||||
|
|
||||||
messages = history_manager.get_history_with_msg(new_query)
|
messages = history_manager.get_history_with_msg(new_query)
|
||||||
history_manager.add_user(query)
|
history_manager.add_user(query)
|
||||||
@ -43,7 +40,7 @@ def chat():
|
|||||||
|
|
||||||
def generate_response():
|
def generate_response():
|
||||||
content = ""
|
content = ""
|
||||||
for delta in model.predict(messages, stream=True):
|
for delta in startup.model.predict(messages, stream=True):
|
||||||
if delta.content:
|
if delta.content:
|
||||||
content += delta.content
|
content += delta.content
|
||||||
response_chunk = json.dumps({
|
response_chunk = json.dumps({
|
||||||
@ -59,7 +56,7 @@ def chat():
|
|||||||
def call():
|
def call():
|
||||||
request_data = json.loads(request.data)
|
request_data = json.loads(request.data)
|
||||||
query = request_data['query']
|
query = request_data['query']
|
||||||
response = model.predict(query)
|
response = startup.model.predict(query)
|
||||||
logger.debug(f"Call query: {query} Response: {response.content}")
|
logger.debug(f"Call query: {query} Response: {response.content}")
|
||||||
|
|
||||||
return jsonify({
|
return jsonify({
|
||||||
@ -68,4 +65,16 @@ def call():
|
|||||||
|
|
||||||
@common.route('/config', methods=['get'])
|
@common.route('/config', methods=['get'])
|
||||||
def get_config():
|
def get_config():
|
||||||
return jsonify(config)
|
return jsonify(startup.config)
|
||||||
|
|
||||||
|
@common.route('/config', methods=['post'])
|
||||||
|
def update_config():
|
||||||
|
request_data = json.loads(request.data)
|
||||||
|
startup.config.update(request_data)
|
||||||
|
startup.config.save()
|
||||||
|
return jsonify(startup.config)
|
||||||
|
|
||||||
|
@common.route('/restart', methods=['POST'])
|
||||||
|
def restart():
|
||||||
|
startup.restart()
|
||||||
|
return jsonify({"message": "Restarted!"})
|
||||||
@ -3,9 +3,8 @@ import json
|
|||||||
import threading
|
import threading
|
||||||
from flask import Blueprint, jsonify, request, Response
|
from flask import Blueprint, jsonify, request, Response
|
||||||
|
|
||||||
from core import HistoryManager
|
|
||||||
from utils.logging_config import setup_logger
|
from utils.logging_config import setup_logger
|
||||||
from core.startup import config, model, dbm
|
from core.startup import startup
|
||||||
|
|
||||||
db = Blueprint('database', __name__, url_prefix="/database")
|
db = Blueprint('database', __name__, url_prefix="/database")
|
||||||
|
|
||||||
@ -15,7 +14,7 @@ progress = {} # 只针对单个用户的进度
|
|||||||
|
|
||||||
@db.route('/', methods=['GET'])
|
@db.route('/', methods=['GET'])
|
||||||
def get_databases():
|
def get_databases():
|
||||||
database = dbm.get_databases()
|
database = startup.dbm.get_databases()
|
||||||
return jsonify(database)
|
return jsonify(database)
|
||||||
|
|
||||||
@db.route('/', methods=['POST'])
|
@db.route('/', methods=['POST'])
|
||||||
@ -25,7 +24,7 @@ def create_database():
|
|||||||
description = data.get('description')
|
description = data.get('description')
|
||||||
db_type = data.get('db_type')
|
db_type = data.get('db_type')
|
||||||
logger.debug(f"Create database {database_name}")
|
logger.debug(f"Create database {database_name}")
|
||||||
database = dbm.create_database(database_name, description, db_type)
|
database = startup.dbm.create_database(database_name, description, db_type)
|
||||||
return jsonify(database)
|
return jsonify(database)
|
||||||
|
|
||||||
# TODO: 删除数据库
|
# TODO: 删除数据库
|
||||||
@ -34,7 +33,7 @@ def delete_database():
|
|||||||
data = json.loads(request.data)
|
data = json.loads(request.data)
|
||||||
db_id = data.get('db_id')
|
db_id = data.get('db_id')
|
||||||
logger.debug(f"Delete database {db_id}")
|
logger.debug(f"Delete database {db_id}")
|
||||||
dbm.delete_database(db_id)
|
startup.dbm.delete_database(db_id)
|
||||||
return jsonify({"message": "删除成功"})
|
return jsonify({"message": "删除成功"})
|
||||||
|
|
||||||
|
|
||||||
@ -44,8 +43,8 @@ def create_document_by_file():
|
|||||||
db_id = data.get('db_id')
|
db_id = data.get('db_id')
|
||||||
files = data.get('files')
|
files = data.get('files')
|
||||||
logger.debug(f"Add document in {db_id} by file: {files}")
|
logger.debug(f"Add document in {db_id} by file: {files}")
|
||||||
dbm.add_files(db_id, files)
|
msg = startup.dbm.add_files(db_id, files)
|
||||||
return jsonify({"status": "全部解析完成"})
|
return jsonify(msg)
|
||||||
|
|
||||||
|
|
||||||
@db.route('/info', methods=['GET'])
|
@db.route('/info', methods=['GET'])
|
||||||
@ -55,7 +54,7 @@ def get_database_info():
|
|||||||
return jsonify({"message": "db_id is required"}), 400
|
return jsonify({"message": "db_id is required"}), 400
|
||||||
|
|
||||||
logger.debug(f"Get database {db_id} info")
|
logger.debug(f"Get database {db_id} info")
|
||||||
database = dbm.get_database_info(db_id)
|
database = startup.dbm.get_database_info(db_id)
|
||||||
|
|
||||||
if database is None:
|
if database is None:
|
||||||
return jsonify({"message": "database not found"}), 404
|
return jsonify({"message": "database not found"}), 404
|
||||||
@ -69,7 +68,7 @@ def delete_document():
|
|||||||
db_id = data.get('db_id')
|
db_id = data.get('db_id')
|
||||||
file_id = data.get('file_id')
|
file_id = data.get('file_id')
|
||||||
logger.debug(f"DELETE document {file_id} info in {db_id}")
|
logger.debug(f"DELETE document {file_id} info in {db_id}")
|
||||||
dbm.delete_file(db_id, file_id)
|
startup.dbm.delete_file(db_id, file_id)
|
||||||
return jsonify({"message": "删除成功"})
|
return jsonify({"message": "删除成功"})
|
||||||
|
|
||||||
@db.route('/document', methods=['GET'])
|
@db.route('/document', methods=['GET'])
|
||||||
@ -77,7 +76,7 @@ def get_document_info():
|
|||||||
db_id = request.args.get('db_id')
|
db_id = request.args.get('db_id')
|
||||||
file_id = request.args.get('file_id')
|
file_id = request.args.get('file_id')
|
||||||
logger.debug(f"GET document {file_id} info in {db_id}")
|
logger.debug(f"GET document {file_id} info in {db_id}")
|
||||||
info = dbm.get_file_info(db_id, file_id)
|
info = startup.dbm.get_file_info(db_id, file_id)
|
||||||
return jsonify(info)
|
return jsonify(info)
|
||||||
|
|
||||||
@db.route('/upload', methods=['POST'])
|
@db.route('/upload', methods=['POST'])
|
||||||
@ -97,5 +96,5 @@ def upload_file():
|
|||||||
|
|
||||||
@db.route('/graph', methods=['GET'])
|
@db.route('/graph', methods=['GET'])
|
||||||
def get_graph_info():
|
def get_graph_info():
|
||||||
graph_info = dbm.get_graph()
|
graph_info = startup.dbm.get_graph()
|
||||||
return jsonify(graph_info)
|
return jsonify(graph_info)
|
||||||
|
|||||||
@ -1,5 +1,5 @@
|
|||||||
<script setup>
|
<script setup>
|
||||||
import { KeepAlive } from 'vue'
|
import { KeepAlive, onMounted } from 'vue'
|
||||||
import { RouterLink, RouterView, useRoute } from 'vue-router'
|
import { RouterLink, RouterView, useRoute } from 'vue-router'
|
||||||
import {
|
import {
|
||||||
MessageOutlined,
|
MessageOutlined,
|
||||||
@ -10,6 +10,20 @@ import {
|
|||||||
BookFilled
|
BookFilled
|
||||||
} from '@ant-design/icons-vue'
|
} from '@ant-design/icons-vue'
|
||||||
import { themeConfig } from '@/assets/theme'
|
import { themeConfig } from '@/assets/theme'
|
||||||
|
import { useConfigStore } from '@/stores/counter'
|
||||||
|
|
||||||
|
const configStore = useConfigStore()
|
||||||
|
|
||||||
|
const getRemoteConfig = () => {
|
||||||
|
fetch('/api/config').then(res => res.json()).then(data => {
|
||||||
|
console.log(data)
|
||||||
|
configStore.setConfig(data)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
onMounted(() => {
|
||||||
|
getRemoteConfig()
|
||||||
|
})
|
||||||
|
|
||||||
// 打印当前页面的路由信息,使用 vue3 的 setup composition API
|
// 打印当前页面的路由信息,使用 vue3 的 setup composition API
|
||||||
const route = useRoute()
|
const route = useRoute()
|
||||||
@ -72,7 +86,7 @@ div.header, #app-router-view {
|
|||||||
flex: 0 0 80px;
|
flex: 0 0 80px;
|
||||||
justify-content: flex-start;
|
justify-content: flex-start;
|
||||||
align-items: center;
|
align-items: center;
|
||||||
background-color: #F4F8F9;
|
background-color: #F2F6F7;
|
||||||
height: 100%;
|
height: 100%;
|
||||||
width: 80px;
|
width: 80px;
|
||||||
border-right: 1px solid #e2eef3;
|
border-right: 1px solid #e2eef3;
|
||||||
|
|||||||
@ -62,7 +62,7 @@ const router = createRouter({
|
|||||||
{
|
{
|
||||||
path: '',
|
path: '',
|
||||||
name: 'setting',
|
name: 'setting',
|
||||||
component: () => import('../views/EmptyView.vue'),
|
component: () => import('../views/SettingView.vue'),
|
||||||
meta: { keepAlive: true }
|
meta: { keepAlive: true }
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
|
|||||||
@ -10,3 +10,17 @@ export const useCounterStore = defineStore('counter', () => {
|
|||||||
|
|
||||||
return { count, doubleCount, increment }
|
return { count, doubleCount, increment }
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
||||||
|
export const useConfigStore = defineStore('config', () => {
|
||||||
|
const config = ref({})
|
||||||
|
function setConfig(newConfig) {
|
||||||
|
config.value = newConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
function setConfigValue(key, value) {
|
||||||
|
config.value[key] = value
|
||||||
|
}
|
||||||
|
|
||||||
|
return { config, setConfig, setConfigValue }
|
||||||
|
})
|
||||||
@ -10,10 +10,13 @@
|
|||||||
<div class="icon"><ReadFilled /></div>
|
<div class="icon"><ReadFilled /></div>
|
||||||
<div class="info">
|
<div class="info">
|
||||||
<h3>{{ database.name }}</h3>
|
<h3>{{ database.name }}</h3>
|
||||||
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata?.row_count }}</span></p>
|
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata?.row_count }}行</span></p>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<p class="description">{{ database.description }}</p>
|
<p class="description">{{ database.description }}</p>
|
||||||
|
<div class="tags">
|
||||||
|
<a-tag color="blue" v-if="database.embed_model">Embed: {{ database.embed_model }}</a-tag>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="sider-bottom">
|
<div class="sider-bottom">
|
||||||
</div>
|
</div>
|
||||||
@ -292,11 +295,15 @@ const addDocumentByFile = () => {
|
|||||||
.then(data => {
|
.then(data => {
|
||||||
console.log(data)
|
console.log(data)
|
||||||
fileList.value = []
|
fileList.value = []
|
||||||
message.success(data.status)
|
if (data.status === 'failed') {
|
||||||
|
message.error(data.message)
|
||||||
|
} else {
|
||||||
|
message.success(data.message)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
.catch(error => {
|
.catch(error => {
|
||||||
console.error(error)
|
console.error(error)
|
||||||
message.error(error.status)
|
message.error(error.message)
|
||||||
})
|
})
|
||||||
.finally(() => {
|
.finally(() => {
|
||||||
getDatabaseInfo()
|
getDatabaseInfo()
|
||||||
|
|||||||
@ -35,10 +35,13 @@
|
|||||||
<div class="icon"><ReadFilled /></div>
|
<div class="icon"><ReadFilled /></div>
|
||||||
<div class="info">
|
<div class="info">
|
||||||
<h3>{{ database.name }}</h3>
|
<h3>{{ database.name }}</h3>
|
||||||
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata.row_count }}</span></p>
|
<p><span>{{ database.metaname }}</span> · <span>{{ database.metadata.row_count }}行</span></p>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<p class="description">{{ database.description }}</p>
|
<p class="description">{{ database.description }}</p>
|
||||||
|
<div class="tags">
|
||||||
|
<a-tag color="blue" v-if="database.embed_model">Embed: {{ database.embed_model }}</a-tag>
|
||||||
|
</div>
|
||||||
<!-- <button @click="deleteDatabase(database.collection_name)">删除</button> -->
|
<!-- <button @click="deleteDatabase(database.collection_name)">删除</button> -->
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@ -197,8 +200,8 @@ onMounted(() => {
|
|||||||
.dbcard, .database {
|
.dbcard, .database {
|
||||||
padding: 10px;
|
padding: 10px;
|
||||||
border-radius: 12px;
|
border-radius: 12px;
|
||||||
width: 380px;
|
width: 360px;
|
||||||
height: 150px;
|
height: 160px;
|
||||||
padding: 20px;
|
padding: 20px;
|
||||||
cursor: pointer;
|
cursor: pointer;
|
||||||
flex: 1 1 380px;
|
flex: 1 1 380px;
|
||||||
@ -226,6 +229,7 @@ onMounted(() => {
|
|||||||
.info {
|
.info {
|
||||||
h3, p {
|
h3, p {
|
||||||
margin: 0;
|
margin: 0;
|
||||||
|
color: black;
|
||||||
}
|
}
|
||||||
|
|
||||||
p {
|
p {
|
||||||
@ -239,9 +243,10 @@ onMounted(() => {
|
|||||||
color: var(--c-text-light-1);
|
color: var(--c-text-light-1);
|
||||||
overflow: hidden;
|
overflow: hidden;
|
||||||
display: -webkit-box;
|
display: -webkit-box;
|
||||||
-webkit-line-clamp: 2;
|
-webkit-line-clamp: 1;
|
||||||
-webkit-box-orient: vertical;
|
-webkit-box-orient: vertical;
|
||||||
text-overflow: ellipsis;
|
text-overflow: ellipsis;
|
||||||
|
margin-bottom: 10px;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
56
web/src/views/SettingView.vue
Normal file
56
web/src/views/SettingView.vue
Normal file
@ -0,0 +1,56 @@
|
|||||||
|
<template>
|
||||||
|
<div class="not-found">
|
||||||
|
<h1>404 - 页面还没做</h1>
|
||||||
|
<p>Sorry, Yemian has not been zuoed.</p>
|
||||||
|
<a-button @click="sendRestart" :loading="state.loading">重载</a-button>
|
||||||
|
<p>{{ configStore.config }}</p>
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup>
|
||||||
|
import { message } from 'ant-design-vue';
|
||||||
|
import { reactive, ref } from 'vue'
|
||||||
|
import { useConfigStore } from '@/stores/counter';
|
||||||
|
|
||||||
|
const configStore = useConfigStore()
|
||||||
|
const state = reactive({
|
||||||
|
loading: false,
|
||||||
|
})
|
||||||
|
|
||||||
|
const sendRestart = () => {
|
||||||
|
console.log('Restarting...')
|
||||||
|
state.loading = true
|
||||||
|
fetch('/api/restart', {
|
||||||
|
method: 'POST',
|
||||||
|
}).then(() => {
|
||||||
|
console.log('Restarted')
|
||||||
|
state.loading = false
|
||||||
|
message.success('重载成功')
|
||||||
|
setTimeout(() => {
|
||||||
|
window.location.reload()
|
||||||
|
}, 1000)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
</script>
|
||||||
|
|
||||||
|
<style scoped>
|
||||||
|
.not-found {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
align-items: center;
|
||||||
|
justify-content: center;
|
||||||
|
height: 80vh;
|
||||||
|
text-align: center;
|
||||||
|
}
|
||||||
|
|
||||||
|
.not-found h1 {
|
||||||
|
font-size: 2rem;
|
||||||
|
margin-bottom: 1rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
.not-found p {
|
||||||
|
font-size: 1.5rem;
|
||||||
|
margin-bottom: 2rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
</style>
|
||||||
Loading…
Reference in New Issue
Block a user