修改后端入口为 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
neo4j*
*/package-lock.json
src/config/base.yaml
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
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
View File

@ -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

View File

@ -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")

View File

@ -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()

View File

@ -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:

View File

@ -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

View File

@ -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):

View File

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

View File

@ -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 只有一个

View File

@ -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):

View File

@ -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)

View File

@ -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:

View File

@ -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__)

View File

@ -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")

View File

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

View File

@ -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:

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):
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]

View File

@ -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)

View File

@ -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():

View File

@ -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")

View File

@ -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

View File

@ -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大于BB小于CA和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>

View File

@ -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()
})

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