修改后端入口为 src, 运行代码不需要使用 cd src, 添加前端查看检索结果

This commit is contained in:
Wenjie Zhang 2024-07-28 16:16:52 +08:00
parent 613b33251b
commit 6971e7e1b7
25 changed files with 223 additions and 93 deletions

3
.gitignore vendored
View File

@ -31,5 +31,6 @@ cache
src/data src/data
neo4j* neo4j*
*/package-lock.json */package-lock.json
src/config/base.yaml
web/package-lock.json web/package-lock.json
saves
notebooks

View File

@ -1,6 +1,7 @@
## Project: Athena <h1 style="text-align: center">Project: Athena</h1>
![home](web/public/home.png)
<img src="web/public/home.png" style="border-radius: 16px; margin: 0 auto; max-height: 400px; display: block;"/>
### 准备 ### 准备
@ -11,18 +12,22 @@
### 启动命令行模式 ### 启动命令行模式
```bash ```bash
cd src python -m src.cli
python cli.py
``` ```
### 启动网页模式 ### 启动网页模式
```bash ```bash
python -m src.api
cd web cd web
npm install npm install
npm run server npm run server
```
cd src
python api.py 或者
```bash
bash run.sh
``` ```

7
run.sh
View File

@ -4,7 +4,7 @@
stop_services() { stop_services() {
echo "Stopping services..." echo "Stopping services..."
pkill -f "npm run server" pkill -f "npm run server"
pkill -f "python api.py" pkill -f "python -m src.api"
exit exit
} }
@ -12,11 +12,10 @@ stop_services() {
trap stop_services SIGINT SIGTERM trap stop_services SIGINT SIGTERM
# Start the server # Start the server
cd src python -m src.api &
python api.py &
# Start the frontend service # Start the frontend service
cd ../web cd web
npm run server & npm run server &
# Wait for all background jobs to finish # Wait for all background jobs to finish

View File

@ -2,8 +2,8 @@ from dotenv import load_dotenv
load_dotenv() load_dotenv()
import os import os
from views import create_app from src.views import create_app
from utils import setup_logger from src.utils import setup_logger
logger = setup_logger("Server") logger = setup_logger("Server")

View File

@ -1,9 +1,9 @@
import os import os
from dotenv import load_dotenv from dotenv import load_dotenv
from core import HistoryManager from src.core import HistoryManager
from core import Retriever from src.core import Retriever
from config import Config from src.config import Config
from models import select_model from src.models import select_model
load_dotenv() load_dotenv()

View File

@ -1,7 +1,7 @@
import os import os
import json import json
import yaml import yaml
from utils.logging_config import setup_logger from src.utils.logging_config import setup_logger
logger = setup_logger("Config") logger = setup_logger("Config")
@ -34,13 +34,13 @@ class Config(SimpleConfig):
def __init__(self, filename=None): def __init__(self, filename=None):
super().__init__() super().__init__()
self.filename = filename or "config/base.yaml"
self._config_items = {} self._config_items = {}
### >>> 默认配置 ### >>> 默认配置
# 可以在 config/base.yaml 中覆盖 # 可以在 config/base.yaml 中覆盖
self.add_item("mode", default="cli", des="运行模式", choices=["cli", "api"]) self.add_item("mode", default="cli", des="运行模式", choices=["cli", "api"])
self.add_item("stream", default=True, des="是否开启流式输出") 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_query_rewrite", default=True, des="是否开启查询重写")
self.add_item("enable_knowledge_base", 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.add_item("model_local_paths", default={}, des="本地模型路径")
### <<< 默认配置结束 ### <<< 默认配置结束
self.filename = filename or os.path.join(self.save_dir, "config", "config.yaml")
self.load() self.load()
self.handle_self() self.handle_self()
@ -99,13 +101,13 @@ class Config(SimpleConfig):
else: else:
logger.warning(f"Unknown config file type {self.filename}") logger.warning(f"Unknown config file type {self.filename}")
else: 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): def save(self):
logger.info(f"Saving config to {self.filename}") logger.info(f"Saving config to {self.filename}")
if self.filename is None: if self.filename is None:
logger.warning("Config file is not specified, save to default config/base.yaml") 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"): if self.filename.endswith(".json"):
with open(self.filename, 'w+') as f: with open(self.filename, 'w+') as f:

View File

@ -1,12 +1,12 @@
import os import os
import json import json
import time import time
from utils import hashstr, setup_logger, is_text_pdf from src.utils import hashstr, setup_logger, is_text_pdf
from plugins import pdf2txt from src.plugins import pdf2txt
from core.knowledgebase import KnowledgeBase from src.core.knowledgebase import KnowledgeBase
from core.filereader import pdfreader, plainreader from src.core.filereader import pdfreader, plainreader
from core.graphbase import GraphDatabase from src.core.graphbase import GraphDatabase
from models.embedding import get_embedding_model from src.models.embedding import get_embedding_model
logger = setup_logger("DataBaseManager") logger = setup_logger("DataBaseManager")
@ -21,6 +21,7 @@ class DataBaseLite:
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) self.embed_model = kwargs.get("embed_model", None)
self.id2file = {f["file_id"]: f for f in self.files}
def update(self, metadata): def update(self, metadata):
@ -48,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 = os.path.join(config.save_dir, "data", "database.json")
self.embed_model = get_embedding_model(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.graph_base = GraphDatabase(self.config, self.embed_model) self.graph_base = GraphDatabase(self.config, self.embed_model)
@ -58,6 +59,7 @@ class DataBaseManager:
self.graph_base.start() self.graph_base.start()
self._load_databases() self._load_databases()
self._update_database()
def _load_databases(self): def _load_databases(self):
"""将数据库的信息保存到本地的文件里面""" """将数据库的信息保存到本地的文件里面"""
@ -73,14 +75,20 @@ class DataBaseManager:
def _save_databases(self): def _save_databases(self):
"""将数据库的信息保存到本地的文件里面""" """将数据库的信息保存到本地的文件里面"""
self._update_database()
with open(self.database_path, "w+") as f: with open(self.database_path, "w+") as f:
json.dump({ json.dump({
"databases": [db.to_dict() for db in self.data["databases"]], "databases": [db.to_dict() for db in self.data["databases"]],
"graph": self.data["graph"] "graph": self.data["graph"]
}, f, ensure_ascii=False, indent=4) }, 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() 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(f"Database number not match, {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): def get_kb_by_id(self, db_id):
for db in self.data["databases"]: for db in self.data["databases"]:
if db.db_id == db_id: if db.db_id == db_id:
return db return db
return None

View File

@ -2,13 +2,13 @@ import os
import json import json
import torch import torch
from neo4j import GraphDatabase as GD from neo4j import GraphDatabase as GD
# from plugins import pdf2txt, OneKE # from src.plugins import pdf2txt, OneKE
from transformers import AutoTokenizer, AutoModel from transformers import AutoTokenizer, AutoModel
from FlagEmbedding import FlagModel, FlagReranker from FlagEmbedding import FlagModel, FlagReranker
import warnings import warnings
from plugins import pdf2txt from src.plugins import pdf2txt
from plugins.oneke import OneKE from src.plugins.oneke import OneKE
warnings.filterwarnings("ignore", category=UserWarning) warnings.filterwarnings("ignore", category=UserWarning)
@ -289,7 +289,7 @@ class GraphDatabase:
ans = [] ans = []
for query in querys: for query in querys:
tep = self.query_specific_entity(query, hops) # 这里是只获取第一个 TODO: 优化 tep = self.query_specific_entity(query, hops) # 这里是只获取第一个 TODO: 优化
ans.extend(tep) ans.extend(tep)
return ans return ans
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):

View File

@ -1,4 +1,4 @@
from utils.logging_config import logger from src.utils.logging_config import logger
class HistoryManager(): class HistoryManager():
def __init__(self, history=None): def __init__(self, history=None):

View File

@ -1,22 +1,25 @@
import os import os
from models.embedding import EmbeddingModel from src.models.embedding import EmbeddingModel
from pymilvus import MilvusClient from pymilvus import MilvusClient
from utils import setup_logger, hashstr from src.utils import setup_logger, hashstr
logger = setup_logger("KnowledgeBase") logger = setup_logger("KnowledgeBase")
class 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.config = config
self._init_config(config) self._init_config(config)
assert embed_model, "embed_model=None" assert embed_model, "embed_model=None"
self.embed_model = embed_model self.embed_model = embed_model
self.client = MilvusClient("data/vector_base/milvus.db") self.client = MilvusClient(self.milvus_path)
def _init_config(self, config): def _init_config(self, config):
self.vector_dim = 1024 # 暂时不知道这个和 embedding model 的 embedding 大小有什么关系 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): def get_collection_names(self):
return self.client.list_collections() return self.client.list_collections()
@ -74,7 +77,7 @@ class KnowledgeBase:
collection_name=collection_name, # target collection collection_name=collection_name, # target collection
data=query_vectors, # query vectors data=query_vectors, # query vectors
limit=limit, # number of returned entities 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 只有一个 return res[0] # 因为 query 只有一个

View File

@ -1,5 +1,5 @@
from models.embedding import Reranker from src.models.embedding import Reranker
from utils.logging_config import setup_logger from src.utils.logging_config import setup_logger
logger = setup_logger("server-common") logger = setup_logger("server-common")
@ -65,13 +65,16 @@ class Retriever:
kb_res = [] kb_res = []
if meta.get("db_name"): 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) kb_res = self.dbm.knowledge_base.search(query, meta["db_name"], limit=5)
for r in kb_res: 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) r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
kb_res.sort(key=lambda x: x["rerank_score"], reverse=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] final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1]
return {"results": final_res, "all_results": kb_res} return {"results": final_res, "all_results": kb_res}
def rewrite_query(self, query, history, meta): def rewrite_query(self, query, history, meta):

View File

@ -1,15 +1,15 @@
from core import DataBaseManager from src.core import DataBaseManager
from core.retriever import Retriever from src.core.retriever import Retriever
from models import select_model from src.models import select_model
from config import Config from src.config import Config
from utils import setup_logger from src.utils import setup_logger
logger = setup_logger("Startup") logger = setup_logger("Startup")
class Startup: class Startup:
def __init__(self): def __init__(self):
self.config = Config("config/base.yaml") self.config = Config()
self.model = select_model(self.config) self.model = select_model(self.config)
self.dbm = DataBaseManager(self.config) self.dbm = DataBaseManager(self.config)
self.retriever = Retriever(self.config, self.dbm, self.model) self.retriever = Retriever(self.config, self.dbm, self.model)

View File

@ -1,4 +1,4 @@
from utils.logging_config import logger from src.utils.logging_config import logger
def select_model(config): 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'}") 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 src.models.chat_model import DeepSeek
return DeepSeek(model_name) return DeepSeek(model_name)
elif model_provider == "zhipu": elif model_provider == "zhipu":
from models.chat_model import Zhipu from src.models.chat_model import Zhipu
return Zhipu(model_name) return Zhipu(model_name)
elif model_provider == "qianfan": elif model_provider == "qianfan":
from models.chat_model import Qianfan from src.models.chat_model import Qianfan
return Qianfan(model_name) return Qianfan(model_name)
elif model_provider == "vllm": elif model_provider == "vllm":
from models.chat_model import VLLM from src.models.chat_model import VLLM
return VLLM(model_name) return VLLM(model_name)
elif model_provider is None: elif model_provider is None:

View File

@ -1,6 +1,6 @@
import os import os
from openai import OpenAI from openai import OpenAI
from utils.logging_config import setup_logger from src.utils.logging_config import setup_logger
logger = setup_logger(__name__) logger = setup_logger(__name__)

View File

@ -1,7 +1,7 @@
import os import os
from FlagEmbedding import FlagModel, FlagReranker from FlagEmbedding import FlagModel, FlagReranker
from utils.logging_config import setup_logger from src.utils.logging_config import setup_logger
logger = setup_logger("EmbeddingModel") logger = setup_logger("EmbeddingModel")

View File

@ -1,2 +1,2 @@
from plugins.oneke import * from src.plugins.oneke import *
from plugins.pdf2txt import * from src.plugins.pdf2txt import *

View File

@ -11,7 +11,7 @@ from transformers import (
BitsAndBytesConfig BitsAndBytesConfig
) )
from utils import setup_logger from src.utils import setup_logger
logger = setup_logger("OneKE") logger = setup_logger("OneKE")
dotenv.load_dotenv() dotenv.load_dotenv()
@ -143,7 +143,7 @@ class OneKE:
print(f"预测结果已添加到 {output_path} 文件中。") print(f"预测结果已添加到 {output_path} 文件中。")
return output_path return output_path
def read_and_process_chars(file_path, char_size=512, overlap_size=100): def read_and_process_chars(file_path, char_size=512, overlap_size=100):
buffer = "" buffer = ""
with open(file_path, 'r', encoding='utf-8') as file: with open(file_path, 'r', encoding='utf-8') as file:

View 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): def is_text_pdf(pdf_path):
import fitz import fitz
@ -10,7 +11,11 @@ def is_text_pdf(pdf_path):
return True return True
return False return False
def hashstr(input_string, length=8): def hashstr(input_string, length=8, with_salt=False):
import hashlib import hashlib
# 添加时间戳作为干扰
if with_salt:
input_string += str(time.time())
hash = hashlib.md5(str(input_string).encode()).hexdigest() hash = hashlib.md5(str(input_string).encode()).hexdigest()
return hash[:length] return hash[:length]

View File

@ -9,8 +9,8 @@ DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件
def setup_logger(name, log_file=None, level=logging.DEBUG, console=False): def setup_logger(name, log_file=None, level=logging.DEBUG, console=False):
if log_file is None: if log_file is None:
log_file = f'output/log/project-{DATETIME}.log' log_file = f'log/project-{DATETIME}.log'
os.makedirs("output/log", exist_ok=True) os.makedirs("log", exist_ok=True)
"""Function to setup logger with the given name and log file.""" """Function to setup logger with the given name and log file."""
logger = logging.getLogger(name) logger = logging.getLogger(name)

View File

@ -1,7 +1,7 @@
from flask import Flask from flask import Flask
from flask_cors import CORS from flask_cors import CORS
from views.common_view import common from src.views.common_view import common
from views.database_view import db from src.views.database_view import db
def create_app(): def create_app():

View File

@ -1,9 +1,9 @@
import json import json
from flask import Blueprint, jsonify, request, Response from flask import Blueprint, jsonify, request, Response
from core import HistoryManager from src.core import HistoryManager
from utils.logging_config import setup_logger from src.utils.logging_config import setup_logger
from core.startup import startup from src.core.startup import startup
common = Blueprint('common', __name__) common = Blueprint('common', __name__)
logger = setup_logger("server-common") logger = setup_logger("server-common")

View File

@ -3,8 +3,8 @@ import json
import threading import threading
from flask import Blueprint, jsonify, request, Response from flask import Blueprint, jsonify, request, Response
from utils.logging_config import setup_logger from src.utils import setup_logger, hashstr
from core.startup import startup from src.core.startup import startup
db = Blueprint('database', __name__, url_prefix="/database") db = Blueprint('database', __name__, url_prefix="/database")
@ -89,9 +89,10 @@ def upload_file():
# elif file.filename.split('.')[-1] not in ['pdf', 'txt', 'md']: # elif file.filename.split('.')[-1] not in ['pdf', 'txt', 'md']:
# return jsonify({'message': 'Unsupported file type'}), 400 # return jsonify({'message': 'Unsupported file type'}), 400
if file: if file:
os.makedirs("data/uploads", exist_ok=True) upload_dir = os.path.join(startup.config.save_dir, "data/uploads")
filename = file.filename os.makedirs(upload_dir, exist_ok=True)
file_path = os.path.join("data/uploads", filename) filename = f"{hashstr(file.filename, 6, with_salt=True)}_{file.filename}"
file_path = os.path.join(upload_dir, filename)
file.save(file_path) file.save(file_path)
return jsonify({'message': 'File successfully uploaded', 'file_path': file_path}), 200 return jsonify({'message': 'File successfully uploaded', 'file_path': file_path}), 200

View File

@ -98,13 +98,36 @@
class="message-md" class="message-md"
@click="consoleMsg(message)"></p> @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 <a-tag
v-for="(ref, index) in message.refs?.knowledge_base.results" class="filetag"
:key="index" v-for="(results, filename) in message.groupedResults"
color="blue" :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> </a-tag>
</div> </div>
</div> </div>
@ -153,7 +176,7 @@ const props = defineProps({
state: Object state: Object
}) })
const emit = defineEmits(['renameTitle']) const emit = defineEmits(['rename-title', 'newconv']);
const configStore = useConfigStore() const configStore = useConfigStore()
const { conv, state } = toRefs(props) const { conv, state } = toRefs(props)
@ -162,12 +185,16 @@ const isStreaming = ref(false)
const panel = ref(null) const panel = ref(null)
const examples = ref([ const examples = ref([
'写一个冒泡排序', '写一个冒泡排序',
'肉碱是什么?', '肉碱的分子量是多少?直接回答',
'洋葱的功效是什么?', '简述大蒜的功效是什么?',
'A大于BB小于CA和C哪个大', 'A大于BB小于CA和C哪个大',
'今天天气怎么样?' '今天天气怎么样?'
]) ])
const opts = reactive({
openDetail: false
})
const meta = reactive({ const meta = reactive({
db_name: computed(() => state.value.databases[state.value.selectedKB]?.metaname), db_name: computed(() => state.value.databases[state.value.selectedKB]?.metaname),
use_graph: false, use_graph: false,
@ -188,9 +215,7 @@ const handleKeyDown = (e) => {
if (e.key === 'Enter' && !e.shiftKey) { if (e.key === 'Enter' && !e.shiftKey) {
e.preventDefault() e.preventDefault()
sendMessage() sendMessage()
console.log('Enter')
} else if (e.key === 'Enter' && e.shiftKey) { } else if (e.key === 'Enter' && e.shiftKey) {
console.log('Shift + Enter')
// Insert a newline character at the current cursor position // Insert a newline character at the current cursor position
const textarea = e.target; const textarea = e.target;
const start = textarea.selectionStart; const start = textarea.selectionStart;
@ -258,22 +283,47 @@ const appendAiMessage = (message, refs=null) => {
id: generateRandomHash(16), id: generateRandomHash(16),
role: 'received', role: 'received',
text: message, text: message,
refs refs,
status: "querying"
}) })
scrollToBottom() scrollToBottom()
} }
const updateMessage = (text, id, refs) => { const updateMessage = (text, id, refs, status) => {
const message = conv.value.messages.find((message) => message.id === id) const message = conv.value.messages.find((message) => message.id === id)
if (message) { if (message) {
message.text = text message.text = text
message.refs = refs message.refs = refs
message.status = status
} else { } else {
console.error('Message not found') console.error('Message not found')
} }
scrollToBottom() 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) => { const simpleCall = (message) => {
return new Promise((resolve, reject) => { return new Promise((resolve, reject) => {
@ -293,12 +343,12 @@ const simpleCall = (message) => {
} }
const sendMessage = () => { const sendMessage = () => {
if (conv.value.inputText.trim()) { const user_input = conv.value.inputText.trim()
if (user_input) {
isStreaming.value = true isStreaming.value = true
appendUserMessage(conv.value.inputText) appendUserMessage(user_input)
appendAiMessage("检索中……", null) appendAiMessage("检索中……", null)
const cur_res_id = conv.value.messages[conv.value.messages.length - 1].id const cur_res_id = conv.value.messages[conv.value.messages.length - 1].id
const user_input = conv.value.inputText
conv.value.inputText = '' conv.value.inputText = ''
fetch('/api/chat', { fetch('/api/chat', {
method: 'POST', method: 'POST',
@ -319,6 +369,7 @@ const sendMessage = () => {
if (done) { if (done) {
console.log(conv.value) console.log(conv.value)
console.log('Finished') console.log('Finished')
updateStatus(cur_res_id, "finished")
isStreaming.value = false isStreaming.value = false
if (conv.value.messages.length === 2) { if (conv.value.messages.length === 2) {
renameTitle() renameTitle()
@ -331,7 +382,7 @@ const sendMessage = () => {
try { try {
const data = JSON.parse(message) 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 conv.value.history = data.history
buffer = '' buffer = ''
} catch (e) { } catch (e) {
@ -558,6 +609,39 @@ onMounted(() => {
.refs { .refs {
margin-bottom: 20px; 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; display: none;
} }
} }
} }
</style> </style>

View File

@ -12,8 +12,10 @@ import {
} from '@ant-design/icons-vue' } from '@ant-design/icons-vue'
import { themeConfig } from '@/assets/theme' import { themeConfig } from '@/assets/theme'
import { useConfigStore } from '@/stores/config' import { useConfigStore } from '@/stores/config'
import { useDatabaseStore } from '@/stores/database'
const configStore = useConfigStore() const configStore = useConfigStore()
const databaseStore = useDatabaseStore()
const getRemoteConfig = () => { const getRemoteConfig = () => {
fetch('/api/config').then(res => res.json()).then(data => { 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(() => { onMounted(() => {
getRemoteDatabase()
getRemoteConfig() getRemoteConfig()
}) })

View 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 }
})