修改后端入口为 src, 运行代码不需要使用 cd src, 添加前端查看检索结果
This commit is contained in:
parent
613b33251b
commit
6971e7e1b7
3
.gitignore
vendored
3
.gitignore
vendored
@ -31,5 +31,6 @@ cache
|
||||
src/data
|
||||
neo4j*
|
||||
*/package-lock.json
|
||||
src/config/base.yaml
|
||||
web/package-lock.json
|
||||
saves
|
||||
notebooks
|
||||
19
README.md
19
README.md
@ -1,6 +1,7 @@
|
||||
## Project: Athena
|
||||
<h1 style="text-align: center">Project: Athena</h1>
|
||||
|
||||

|
||||
|
||||
<img src="web/public/home.png" style="border-radius: 16px; margin: 0 auto; max-height: 400px; display: block;"/>
|
||||
|
||||
### 准备
|
||||
|
||||
@ -11,18 +12,22 @@
|
||||
### 启动命令行模式
|
||||
|
||||
```bash
|
||||
cd src
|
||||
python cli.py
|
||||
python -m src.cli
|
||||
```
|
||||
|
||||
### 启动网页模式
|
||||
|
||||
```bash
|
||||
python -m src.api
|
||||
|
||||
cd web
|
||||
npm install
|
||||
npm run server
|
||||
|
||||
cd src
|
||||
python api.py
|
||||
```
|
||||
|
||||
或者
|
||||
|
||||
```bash
|
||||
bash run.sh
|
||||
```
|
||||
|
||||
|
||||
7
run.sh
7
run.sh
@ -4,7 +4,7 @@
|
||||
stop_services() {
|
||||
echo "Stopping services..."
|
||||
pkill -f "npm run server"
|
||||
pkill -f "python api.py"
|
||||
pkill -f "python -m src.api"
|
||||
exit
|
||||
}
|
||||
|
||||
@ -12,11 +12,10 @@ stop_services() {
|
||||
trap stop_services SIGINT SIGTERM
|
||||
|
||||
# Start the server
|
||||
cd src
|
||||
python api.py &
|
||||
python -m src.api &
|
||||
|
||||
# Start the frontend service
|
||||
cd ../web
|
||||
cd web
|
||||
npm run server &
|
||||
|
||||
# Wait for all background jobs to finish
|
||||
|
||||
@ -2,8 +2,8 @@ from dotenv import load_dotenv
|
||||
load_dotenv()
|
||||
|
||||
import os
|
||||
from views import create_app
|
||||
from utils import setup_logger
|
||||
from src.views import create_app
|
||||
from src.utils import setup_logger
|
||||
|
||||
logger = setup_logger("Server")
|
||||
|
||||
|
||||
@ -1,9 +1,9 @@
|
||||
import os
|
||||
from dotenv import load_dotenv
|
||||
from core import HistoryManager
|
||||
from core import Retriever
|
||||
from config import Config
|
||||
from models import select_model
|
||||
from src.core import HistoryManager
|
||||
from src.core import Retriever
|
||||
from src.config import Config
|
||||
from src.models import select_model
|
||||
|
||||
load_dotenv()
|
||||
|
||||
|
||||
@ -1,7 +1,7 @@
|
||||
import os
|
||||
import json
|
||||
import yaml
|
||||
from utils.logging_config import setup_logger
|
||||
from src.utils.logging_config import setup_logger
|
||||
|
||||
logger = setup_logger("Config")
|
||||
|
||||
@ -34,13 +34,13 @@ class Config(SimpleConfig):
|
||||
|
||||
def __init__(self, filename=None):
|
||||
super().__init__()
|
||||
self.filename = filename or "config/base.yaml"
|
||||
self._config_items = {}
|
||||
|
||||
### >>> 默认配置
|
||||
# 可以在 config/base.yaml 中覆盖
|
||||
self.add_item("mode", default="cli", des="运行模式", choices=["cli", "api"])
|
||||
self.add_item("stream", default=True, des="是否开启流式输出")
|
||||
self.add_item("save_dir", default="saves", des="保存目录")
|
||||
# 功能选项
|
||||
self.add_item("enable_query_rewrite", default=True, des="是否开启查询重写")
|
||||
self.add_item("enable_knowledge_base", default=True, des="是否开启知识库")
|
||||
@ -57,6 +57,8 @@ class Config(SimpleConfig):
|
||||
self.add_item("model_local_paths", default={}, des="本地模型路径")
|
||||
### <<< 默认配置结束
|
||||
|
||||
self.filename = filename or os.path.join(self.save_dir, "config", "config.yaml")
|
||||
|
||||
self.load()
|
||||
self.handle_self()
|
||||
|
||||
@ -99,13 +101,13 @@ class Config(SimpleConfig):
|
||||
else:
|
||||
logger.warning(f"Unknown config file type {self.filename}")
|
||||
else:
|
||||
logger.warning(f"Config file {self.filename} not found")
|
||||
logger.warning(f"\n\n{'='*70}\n{'Config file not found':^70}\n{'You can custum your config in `' + self.filename + '`':^70}\n{'='*70}\n\n")
|
||||
|
||||
def save(self):
|
||||
logger.info(f"Saving config to {self.filename}")
|
||||
if self.filename is None:
|
||||
logger.warning("Config file is not specified, save to default config/base.yaml")
|
||||
self.filename = "config/base.yaml"
|
||||
self.filename = os.path.join(self.save_dir, "config", "config.yaml")
|
||||
|
||||
if self.filename.endswith(".json"):
|
||||
with open(self.filename, 'w+') as f:
|
||||
|
||||
@ -1,12 +1,12 @@
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
from utils import hashstr, setup_logger, is_text_pdf
|
||||
from plugins import pdf2txt
|
||||
from core.knowledgebase import KnowledgeBase
|
||||
from core.filereader import pdfreader, plainreader
|
||||
from core.graphbase import GraphDatabase
|
||||
from models.embedding import get_embedding_model
|
||||
from src.utils import hashstr, setup_logger, is_text_pdf
|
||||
from src.plugins import pdf2txt
|
||||
from src.core.knowledgebase import KnowledgeBase
|
||||
from src.core.filereader import pdfreader, plainreader
|
||||
from src.core.graphbase import GraphDatabase
|
||||
from src.models.embedding import get_embedding_model
|
||||
|
||||
logger = setup_logger("DataBaseManager")
|
||||
|
||||
@ -21,6 +21,7 @@ class DataBaseLite:
|
||||
self.metadata = kwargs.get("metaname", {})
|
||||
self.files = kwargs.get("files", [])
|
||||
self.embed_model = kwargs.get("embed_model", None)
|
||||
self.id2file = {f["file_id"]: f for f in self.files}
|
||||
|
||||
|
||||
def update(self, metadata):
|
||||
@ -48,7 +49,7 @@ class DataBaseManager:
|
||||
|
||||
def __init__(self, config=None) -> None:
|
||||
self.config = config
|
||||
self.database_path = "data/databases.json"
|
||||
self.database_path = os.path.join(config.save_dir, "data", "database.json")
|
||||
self.embed_model = get_embedding_model(config)
|
||||
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
||||
self.graph_base = GraphDatabase(self.config, self.embed_model)
|
||||
@ -58,6 +59,7 @@ class DataBaseManager:
|
||||
self.graph_base.start()
|
||||
|
||||
self._load_databases()
|
||||
self._update_database()
|
||||
|
||||
def _load_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
@ -73,14 +75,20 @@ class DataBaseManager:
|
||||
|
||||
def _save_databases(self):
|
||||
"""将数据库的信息保存到本地的文件里面"""
|
||||
self._update_database()
|
||||
with open(self.database_path, "w+") as f:
|
||||
json.dump({
|
||||
"databases": [db.to_dict() for db in self.data["databases"]],
|
||||
"graph": self.data["graph"]
|
||||
}, f, ensure_ascii=False, indent=4)
|
||||
|
||||
def get_databases(self):
|
||||
def _update_database(self):
|
||||
self.id2db = {db.db_id: db for db in self.data["databases"]}
|
||||
self.name2db = {db.name: db for db in self.data["databases"]}
|
||||
self.metaname2db = {db.metaname: db for db in self.data["databases"]}
|
||||
|
||||
def get_databases(self):
|
||||
self._update_database()
|
||||
knowledge_base_collections = self.knowledge_base.get_collection_names()
|
||||
if len(self.data["databases"]) != len(knowledge_base_collections):
|
||||
logger.warning(f"Database number not match, {knowledge_base_collections}")
|
||||
@ -219,4 +227,5 @@ class DataBaseManager:
|
||||
def get_kb_by_id(self, db_id):
|
||||
for db in self.data["databases"]:
|
||||
if db.db_id == db_id:
|
||||
return db
|
||||
return db
|
||||
return None
|
||||
@ -2,13 +2,13 @@ import os
|
||||
import json
|
||||
import torch
|
||||
from neo4j import GraphDatabase as GD
|
||||
# from plugins import pdf2txt, OneKE
|
||||
# from src.plugins import pdf2txt, OneKE
|
||||
from transformers import AutoTokenizer, AutoModel
|
||||
from FlagEmbedding import FlagModel, FlagReranker
|
||||
import warnings
|
||||
|
||||
from plugins import pdf2txt
|
||||
from plugins.oneke import OneKE
|
||||
from src.plugins import pdf2txt
|
||||
from src.plugins.oneke import OneKE
|
||||
|
||||
warnings.filterwarnings("ignore", category=UserWarning)
|
||||
|
||||
@ -289,7 +289,7 @@ class GraphDatabase:
|
||||
ans = []
|
||||
for query in querys:
|
||||
tep = self.query_specific_entity(query, hops) # 这里是只获取第一个 TODO: 优化
|
||||
ans.extend(tep)
|
||||
ans.extend(tep)
|
||||
return ans
|
||||
|
||||
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
from utils.logging_config import logger
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
class HistoryManager():
|
||||
def __init__(self, history=None):
|
||||
|
||||
@ -1,22 +1,25 @@
|
||||
import os
|
||||
|
||||
from models.embedding import EmbeddingModel
|
||||
from src.models.embedding import EmbeddingModel
|
||||
from pymilvus import MilvusClient
|
||||
from utils import setup_logger, hashstr
|
||||
from src.utils import setup_logger, hashstr
|
||||
logger = setup_logger("KnowledgeBase")
|
||||
|
||||
|
||||
class KnowledgeBase:
|
||||
|
||||
def __init__(self, config=None, embed_model=None ) -> None:
|
||||
def __init__(self, config=None, embed_model=None) -> None:
|
||||
self.config = config
|
||||
self._init_config(config)
|
||||
|
||||
assert embed_model, "embed_model=None"
|
||||
self.embed_model = embed_model
|
||||
self.client = MilvusClient("data/vector_base/milvus.db")
|
||||
self.client = MilvusClient(self.milvus_path)
|
||||
|
||||
def _init_config(self, config):
|
||||
self.vector_dim = 1024 # 暂时不知道这个和 embedding model 的 embedding 大小有什么关系
|
||||
self.milvus_path = os.path.join(config.save_dir, "data/vector_base/milvus.db")
|
||||
os.makedirs(os.path.dirname(self.milvus_path), exist_ok=True)
|
||||
|
||||
def get_collection_names(self):
|
||||
return self.client.list_collections()
|
||||
@ -74,7 +77,7 @@ class KnowledgeBase:
|
||||
collection_name=collection_name, # target collection
|
||||
data=query_vectors, # query vectors
|
||||
limit=limit, # number of returned entities
|
||||
output_fields=["text"], # specifies fields to be returned
|
||||
output_fields=["text", "file_id"], # specifies fields to be returned
|
||||
)
|
||||
|
||||
return res[0] # 因为 query 只有一个
|
||||
|
||||
@ -1,5 +1,5 @@
|
||||
from models.embedding import Reranker
|
||||
from utils.logging_config import setup_logger
|
||||
from src.models.embedding import Reranker
|
||||
from src.utils.logging_config import setup_logger
|
||||
logger = setup_logger("server-common")
|
||||
|
||||
|
||||
@ -65,13 +65,16 @@ class Retriever:
|
||||
|
||||
kb_res = []
|
||||
if meta.get("db_name"):
|
||||
kb = self.dbm.metaname2db[meta["db_name"]]
|
||||
kb_res = self.dbm.knowledge_base.search(query, meta["db_name"], limit=5)
|
||||
for r in kb_res:
|
||||
r["file"] = kb.id2file[r["entity"]["file_id"]]
|
||||
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
|
||||
|
||||
kb_res.sort(key=lambda x: x["rerank_score"], reverse=True)
|
||||
|
||||
final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1]
|
||||
|
||||
return {"results": final_res, "all_results": kb_res}
|
||||
|
||||
def rewrite_query(self, query, history, meta):
|
||||
|
||||
@ -1,15 +1,15 @@
|
||||
from core import DataBaseManager
|
||||
from core.retriever import Retriever
|
||||
from models import select_model
|
||||
from config import Config
|
||||
from utils import setup_logger
|
||||
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 setup_logger
|
||||
|
||||
logger = setup_logger("Startup")
|
||||
|
||||
|
||||
class Startup:
|
||||
def __init__(self):
|
||||
self.config = Config("config/base.yaml")
|
||||
self.config = Config()
|
||||
self.model = select_model(self.config)
|
||||
self.dbm = DataBaseManager(self.config)
|
||||
self.retriever = Retriever(self.config, self.dbm, self.model)
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
from utils.logging_config import logger
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
|
||||
def select_model(config):
|
||||
@ -9,19 +9,19 @@ def select_model(config):
|
||||
logger.info(f"Selecting model from {model_provider} with {model_name or 'default'}")
|
||||
|
||||
if model_provider == "deepseek":
|
||||
from models.chat_model import DeepSeek
|
||||
from src.models.chat_model import DeepSeek
|
||||
return DeepSeek(model_name)
|
||||
|
||||
elif model_provider == "zhipu":
|
||||
from models.chat_model import Zhipu
|
||||
from src.models.chat_model import Zhipu
|
||||
return Zhipu(model_name)
|
||||
|
||||
elif model_provider == "qianfan":
|
||||
from models.chat_model import Qianfan
|
||||
from src.models.chat_model import Qianfan
|
||||
return Qianfan(model_name)
|
||||
|
||||
elif model_provider == "vllm":
|
||||
from models.chat_model import VLLM
|
||||
from src.models.chat_model import VLLM
|
||||
return VLLM(model_name)
|
||||
|
||||
elif model_provider is None:
|
||||
|
||||
@ -1,6 +1,6 @@
|
||||
import os
|
||||
from openai import OpenAI
|
||||
from utils.logging_config import setup_logger
|
||||
from src.utils.logging_config import setup_logger
|
||||
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
@ -1,7 +1,7 @@
|
||||
import os
|
||||
from FlagEmbedding import FlagModel, FlagReranker
|
||||
|
||||
from utils.logging_config import setup_logger
|
||||
from src.utils.logging_config import setup_logger
|
||||
|
||||
|
||||
logger = setup_logger("EmbeddingModel")
|
||||
|
||||
@ -1,2 +1,2 @@
|
||||
from plugins.oneke import *
|
||||
from plugins.pdf2txt import *
|
||||
from src.plugins.oneke import *
|
||||
from src.plugins.pdf2txt import *
|
||||
@ -11,7 +11,7 @@ from transformers import (
|
||||
BitsAndBytesConfig
|
||||
)
|
||||
|
||||
from utils import setup_logger
|
||||
from src.utils import setup_logger
|
||||
logger = setup_logger("OneKE")
|
||||
|
||||
dotenv.load_dotenv()
|
||||
@ -143,7 +143,7 @@ class OneKE:
|
||||
|
||||
print(f"预测结果已添加到 {output_path} 文件中。")
|
||||
return output_path
|
||||
|
||||
|
||||
def read_and_process_chars(file_path, char_size=512, overlap_size=100):
|
||||
buffer = ""
|
||||
with open(file_path, 'r', encoding='utf-8') as file:
|
||||
|
||||
@ -1,4 +1,5 @@
|
||||
from utils.logging_config import setup_logger, logger
|
||||
import time
|
||||
from src.utils.logging_config import setup_logger, logger
|
||||
|
||||
def is_text_pdf(pdf_path):
|
||||
import fitz
|
||||
@ -10,7 +11,11 @@ def is_text_pdf(pdf_path):
|
||||
return True
|
||||
return False
|
||||
|
||||
def hashstr(input_string, length=8):
|
||||
def hashstr(input_string, length=8, with_salt=False):
|
||||
import hashlib
|
||||
# 添加时间戳作为干扰
|
||||
if with_salt:
|
||||
input_string += str(time.time())
|
||||
|
||||
hash = hashlib.md5(str(input_string).encode()).hexdigest()
|
||||
return hash[:length]
|
||||
@ -9,8 +9,8 @@ DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件
|
||||
def setup_logger(name, log_file=None, level=logging.DEBUG, console=False):
|
||||
|
||||
if log_file is None:
|
||||
log_file = f'output/log/project-{DATETIME}.log'
|
||||
os.makedirs("output/log", exist_ok=True)
|
||||
log_file = f'log/project-{DATETIME}.log'
|
||||
os.makedirs("log", exist_ok=True)
|
||||
|
||||
"""Function to setup logger with the given name and log file."""
|
||||
logger = logging.getLogger(name)
|
||||
|
||||
@ -1,7 +1,7 @@
|
||||
from flask import Flask
|
||||
from flask_cors import CORS
|
||||
from views.common_view import common
|
||||
from views.database_view import db
|
||||
from src.views.common_view import common
|
||||
from src.views.database_view import db
|
||||
|
||||
|
||||
def create_app():
|
||||
|
||||
@ -1,9 +1,9 @@
|
||||
import json
|
||||
from flask import Blueprint, jsonify, request, Response
|
||||
|
||||
from core import HistoryManager
|
||||
from utils.logging_config import setup_logger
|
||||
from core.startup import startup
|
||||
from src.core import HistoryManager
|
||||
from src.utils.logging_config import setup_logger
|
||||
from src.core.startup import startup
|
||||
|
||||
common = Blueprint('common', __name__)
|
||||
logger = setup_logger("server-common")
|
||||
|
||||
@ -3,8 +3,8 @@ import json
|
||||
import threading
|
||||
from flask import Blueprint, jsonify, request, Response
|
||||
|
||||
from utils.logging_config import setup_logger
|
||||
from core.startup import startup
|
||||
from src.utils import setup_logger, hashstr
|
||||
from src.core.startup import startup
|
||||
|
||||
db = Blueprint('database', __name__, url_prefix="/database")
|
||||
|
||||
@ -89,9 +89,10 @@ def upload_file():
|
||||
# elif file.filename.split('.')[-1] not in ['pdf', 'txt', 'md']:
|
||||
# return jsonify({'message': 'Unsupported file type'}), 400
|
||||
if file:
|
||||
os.makedirs("data/uploads", exist_ok=True)
|
||||
filename = file.filename
|
||||
file_path = os.path.join("data/uploads", filename)
|
||||
upload_dir = os.path.join(startup.config.save_dir, "data/uploads")
|
||||
os.makedirs(upload_dir, exist_ok=True)
|
||||
filename = f"{hashstr(file.filename, 6, with_salt=True)}_{file.filename}"
|
||||
file_path = os.path.join(upload_dir, filename)
|
||||
file.save(file_path)
|
||||
return jsonify({'message': 'File successfully uploaded', 'file_path': file_path}), 200
|
||||
|
||||
|
||||
@ -98,13 +98,36 @@
|
||||
class="message-md"
|
||||
@click="consoleMsg(message)"></p>
|
||||
|
||||
<div class="refs" v-if="message.role=='received' && message.refs?.knowledge_base.results.length > 0">
|
||||
<div class="refs" v-if="message.role=='received' && message.groupedResults && message.status=='finished'">
|
||||
<a-tag
|
||||
v-for="(ref, index) in message.refs?.knowledge_base.results"
|
||||
:key="index"
|
||||
color="blue"
|
||||
class="filetag"
|
||||
v-for="(results, filename) in message.groupedResults"
|
||||
:key="filename"
|
||||
@click="opts.openDetail = true"
|
||||
:bordered="false"
|
||||
>
|
||||
{{ ref.id }}
|
||||
{{ filename }}
|
||||
<a-drawer
|
||||
v-model:open="opts.openDetail"
|
||||
title="检索详情"
|
||||
width="800"
|
||||
:contentWrapperStyle="{ maxWidth: '100%'}"
|
||||
placement="right"
|
||||
class="retrieval-detail"
|
||||
rootClassName="root"
|
||||
>
|
||||
<div class="fileinfo">
|
||||
<p><strong>文件名:</strong> {{ results[0].file.filename }}</p>
|
||||
<p><strong>文件类型:</strong> {{ results[0].file.type }}</p>
|
||||
<p><strong>创建时间:</strong> {{ new Date(results[0].file.created_at * 1000).toLocaleString() }}</p>
|
||||
</div>
|
||||
<div v-for="(res, idx) in results" :key="idx" class="result-item">
|
||||
<p class="result-id"><strong>ID:</strong> #{{ res.id }}</p>
|
||||
<p class="result-distance"><strong>相似度距离:</strong> {{ res.distance }}</p>
|
||||
<p class="result-rerank-score"><strong>重排序分数:</strong> {{ res.rerank_score }}</p>
|
||||
<p class="result-text">{{ res.entity.text }}</p>
|
||||
</div>
|
||||
</a-drawer>
|
||||
</a-tag>
|
||||
</div>
|
||||
</div>
|
||||
@ -153,7 +176,7 @@ const props = defineProps({
|
||||
state: Object
|
||||
})
|
||||
|
||||
const emit = defineEmits(['renameTitle'])
|
||||
const emit = defineEmits(['rename-title', 'newconv']);
|
||||
const configStore = useConfigStore()
|
||||
|
||||
const { conv, state } = toRefs(props)
|
||||
@ -162,12 +185,16 @@ const isStreaming = ref(false)
|
||||
const panel = ref(null)
|
||||
const examples = ref([
|
||||
'写一个冒泡排序',
|
||||
'肉碱是什么?',
|
||||
'洋葱的功效是什么?',
|
||||
'肉碱的分子量是多少?直接回答',
|
||||
'简述大蒜的功效是什么?',
|
||||
'A大于B,B小于C,A和C哪个大?',
|
||||
'今天天气怎么样?'
|
||||
])
|
||||
|
||||
const opts = reactive({
|
||||
openDetail: false
|
||||
})
|
||||
|
||||
const meta = reactive({
|
||||
db_name: computed(() => state.value.databases[state.value.selectedKB]?.metaname),
|
||||
use_graph: false,
|
||||
@ -188,9 +215,7 @@ const handleKeyDown = (e) => {
|
||||
if (e.key === 'Enter' && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
sendMessage()
|
||||
console.log('Enter')
|
||||
} else if (e.key === 'Enter' && e.shiftKey) {
|
||||
console.log('Shift + Enter')
|
||||
// Insert a newline character at the current cursor position
|
||||
const textarea = e.target;
|
||||
const start = textarea.selectionStart;
|
||||
@ -258,22 +283,47 @@ const appendAiMessage = (message, refs=null) => {
|
||||
id: generateRandomHash(16),
|
||||
role: 'received',
|
||||
text: message,
|
||||
refs
|
||||
refs,
|
||||
status: "querying"
|
||||
})
|
||||
scrollToBottom()
|
||||
}
|
||||
|
||||
const updateMessage = (text, id, refs) => {
|
||||
const updateMessage = (text, id, refs, status) => {
|
||||
const message = conv.value.messages.find((message) => message.id === id)
|
||||
if (message) {
|
||||
message.text = text
|
||||
message.refs = refs
|
||||
message.status = status
|
||||
} else {
|
||||
console.error('Message not found')
|
||||
}
|
||||
|
||||
scrollToBottom()
|
||||
}
|
||||
|
||||
const updateStatus = (id, status) => {
|
||||
const message = conv.value.messages.find((message) => message.id === id)
|
||||
if (message) {
|
||||
message.status = status
|
||||
} else {
|
||||
console.error('Message not found')
|
||||
}
|
||||
|
||||
console.log("updateStatus", message, message.refs.knowledge_base.results.length > 0)
|
||||
if (message.refs.knowledge_base.results.length > 0) {
|
||||
message.groupedResults = message.refs.knowledge_base.results.reduce((acc, result) => {
|
||||
const { filename } = result.file;
|
||||
console.log(acc, result, filename)
|
||||
if (!acc[filename]) {
|
||||
acc[filename] = []
|
||||
}
|
||||
acc[filename].push(result)
|
||||
return acc;
|
||||
}, {})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
const simpleCall = (message) => {
|
||||
return new Promise((resolve, reject) => {
|
||||
@ -293,12 +343,12 @@ const simpleCall = (message) => {
|
||||
}
|
||||
|
||||
const sendMessage = () => {
|
||||
if (conv.value.inputText.trim()) {
|
||||
const user_input = conv.value.inputText.trim()
|
||||
if (user_input) {
|
||||
isStreaming.value = true
|
||||
appendUserMessage(conv.value.inputText)
|
||||
appendUserMessage(user_input)
|
||||
appendAiMessage("检索中……", null)
|
||||
const cur_res_id = conv.value.messages[conv.value.messages.length - 1].id
|
||||
const user_input = conv.value.inputText
|
||||
conv.value.inputText = ''
|
||||
fetch('/api/chat', {
|
||||
method: 'POST',
|
||||
@ -319,6 +369,7 @@ const sendMessage = () => {
|
||||
if (done) {
|
||||
console.log(conv.value)
|
||||
console.log('Finished')
|
||||
updateStatus(cur_res_id, "finished")
|
||||
isStreaming.value = false
|
||||
if (conv.value.messages.length === 2) {
|
||||
renameTitle()
|
||||
@ -331,7 +382,7 @@ const sendMessage = () => {
|
||||
|
||||
try {
|
||||
const data = JSON.parse(message)
|
||||
updateMessage(data.response, cur_res_id, data.refs)
|
||||
updateMessage(data.response, cur_res_id, data.refs, "loading")
|
||||
conv.value.history = data.history
|
||||
buffer = ''
|
||||
} catch (e) {
|
||||
@ -558,6 +609,39 @@ onMounted(() => {
|
||||
|
||||
.refs {
|
||||
margin-bottom: 20px;
|
||||
|
||||
.filetag:hover {
|
||||
cursor: pointer;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
.retrieval-detail {
|
||||
.fileinfo {
|
||||
margin-bottom: 20px;
|
||||
padding: 10px;
|
||||
background: var(--main-light-3);
|
||||
border-radius: 8px;
|
||||
p {
|
||||
margin: 0;
|
||||
line-height: 1.5;
|
||||
}
|
||||
}
|
||||
|
||||
.result-item {
|
||||
margin-bottom: 20px;
|
||||
padding: 10px;
|
||||
border: 1px solid #e8e8e8;
|
||||
border-radius: 8px;
|
||||
background: var(--main-light-4);
|
||||
|
||||
.result-id,
|
||||
.result-distance,
|
||||
.result-rerank-score,
|
||||
.result-text-label,
|
||||
.result-text {
|
||||
margin: 5px 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@ -685,8 +769,5 @@ button:disabled {
|
||||
display: none;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
}
|
||||
</style>
|
||||
|
||||
@ -12,8 +12,10 @@ import {
|
||||
} from '@ant-design/icons-vue'
|
||||
import { themeConfig } from '@/assets/theme'
|
||||
import { useConfigStore } from '@/stores/config'
|
||||
import { useDatabaseStore } from '@/stores/database'
|
||||
|
||||
const configStore = useConfigStore()
|
||||
const databaseStore = useDatabaseStore()
|
||||
|
||||
const getRemoteConfig = () => {
|
||||
fetch('/api/config').then(res => res.json()).then(data => {
|
||||
@ -22,7 +24,15 @@ const getRemoteConfig = () => {
|
||||
})
|
||||
}
|
||||
|
||||
const getRemoteDatabase = () => {
|
||||
fetch('/api/database').then(res => res.json()).then(data => {
|
||||
console.log("database", data)
|
||||
databaseStore.setDatabase(data.databases)
|
||||
})
|
||||
}
|
||||
|
||||
onMounted(() => {
|
||||
getRemoteDatabase()
|
||||
getRemoteConfig()
|
||||
})
|
||||
|
||||
|
||||
11
web/src/stores/database.js
Normal file
11
web/src/stores/database.js
Normal file
@ -0,0 +1,11 @@
|
||||
import { ref, computed } from 'vue'
|
||||
import { defineStore } from 'pinia'
|
||||
|
||||
export const useDatabaseStore = defineStore('database', () => {
|
||||
const db = ref({})
|
||||
function setDatabase(newDatabase) {
|
||||
db.value = newDatabase
|
||||
}
|
||||
|
||||
return { db, setDatabase }
|
||||
})
|
||||
Loading…
Reference in New Issue
Block a user