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 +
### 准备
@@ -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)">
- 文件名: {{ 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 }}
+