diff --git a/.gitignore b/.gitignore index da5dd14d..53dc6d4a 100644 --- a/.gitignore +++ b/.gitignore @@ -31,5 +31,6 @@ cache src/data neo4j* */package-lock.json -src/config/base.yaml web/package-lock.json +saves +notebooks \ No newline at end of file diff --git a/README.md b/README.md index 5f47a457..a391ec99 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,7 @@ -## Project: Athena +

Project: Athena

-![home](web/public/home.png) + + ### 准备 @@ -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 ``` diff --git a/run.sh b/run.sh index 4631b502..a05c6e52 100644 --- a/run.sh +++ b/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 diff --git a/src/api.py b/src/api.py index addf6f34..68c78193 100644 --- a/src/api.py +++ b/src/api.py @@ -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") diff --git a/src/cli.py b/src/cli.py index 816ec1e9..539eb10c 100644 --- a/src/cli.py +++ b/src/cli.py @@ -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() diff --git a/src/config/__init__.py b/src/config/__init__.py index 4fa70ba4..2d362b53 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -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: diff --git a/src/core/database.py b/src/core/database.py index cc8223c5..b10922ac 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -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 \ No newline at end of file + return db + return None \ No newline at end of file diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 1906e009..31e5b377 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -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): diff --git a/src/core/history.py b/src/core/history.py index 42075438..6b9fbced 100644 --- a/src/core/history.py +++ b/src/core/history.py @@ -1,4 +1,4 @@ -from utils.logging_config import logger +from src.utils.logging_config import logger class HistoryManager(): def __init__(self, history=None): diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index fd4946eb..20eb68f0 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -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 只有一个 diff --git a/src/core/retriever.py b/src/core/retriever.py index 98420004..c0750c24 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -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): diff --git a/src/core/startup.py b/src/core/startup.py index 4a862028..bc85bba6 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -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) diff --git a/src/models/__init__.py b/src/models/__init__.py index 33142937..f2bb4c96 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -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: diff --git a/src/models/chat_model.py b/src/models/chat_model.py index 94e4f7f0..a5b2a771 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -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__) diff --git a/src/models/embedding.py b/src/models/embedding.py index 3097dfc8..06a4d8ea 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -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") diff --git a/src/plugins/__init__.py b/src/plugins/__init__.py index cae80c54..9804a77f 100644 --- a/src/plugins/__init__.py +++ b/src/plugins/__init__.py @@ -1,2 +1,2 @@ -from plugins.oneke import * -from plugins.pdf2txt import * \ No newline at end of file +from src.plugins.oneke import * +from src.plugins.pdf2txt import * \ No newline at end of file diff --git a/src/plugins/oneke.py b/src/plugins/oneke.py index 0f1f3757..9f150d26 100644 --- a/src/plugins/oneke.py +++ b/src/plugins/oneke.py @@ -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: diff --git a/src/utils/__init__.py b/src/utils/__init__.py index 2ff4a3c6..1df9e7e6 100644 --- a/src/utils/__init__.py +++ b/src/utils/__init__.py @@ -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] \ No newline at end of file diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index 9368aa3e..6f4849ce 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -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) diff --git a/src/views/__init__.py b/src/views/__init__.py index 784eba5c..6886b6aa 100644 --- a/src/views/__init__.py +++ b/src/views/__init__.py @@ -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(): diff --git a/src/views/common_view.py b/src/views/common_view.py index 58d0d04f..04bbc1e9 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -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") diff --git a/src/views/database_view.py b/src/views/database_view.py index 4114669d..6f67b805 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -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 diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 5e844a45..f9a20cb2 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -98,13 +98,36 @@ class="message-md" @click="consoleMsg(message)">

-
+
- {{ ref.id }} + {{ filename }} + +
+

文件名: {{ results[0].file.filename }}

+

文件类型: {{ results[0].file.type }}

+

创建时间: {{ new Date(results[0].file.created_at * 1000).toLocaleString() }}

+
+
+

ID: #{{ res.id }}

+

相似度距离: {{ res.distance }}

+

重排序分数: {{ res.rerank_score }}

+

{{ res.entity.text }}

+
+
@@ -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; } } - - - } diff --git a/web/src/layouts/AppLayout.vue b/web/src/layouts/AppLayout.vue index 85a3ba03..b66974f4 100644 --- a/web/src/layouts/AppLayout.vue +++ b/web/src/layouts/AppLayout.vue @@ -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() }) diff --git a/web/src/stores/database.js b/web/src/stores/database.js new file mode 100644 index 00000000..aadb3f51 --- /dev/null +++ b/web/src/stores/database.js @@ -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 } +}) \ No newline at end of file