修改后端入口为 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
|
src/data
|
||||||
neo4j*
|
neo4j*
|
||||||
*/package-lock.json
|
*/package-lock.json
|
||||||
src/config/base.yaml
|
|
||||||
web/package-lock.json
|
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
|
```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
7
run.sh
@ -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
|
||||||
|
|||||||
@ -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")
|
||||||
|
|
||||||
|
|||||||
@ -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()
|
||||||
|
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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
|
||||||
@ -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):
|
||||||
|
|||||||
@ -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):
|
||||||
|
|||||||
@ -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 只有一个
|
||||||
|
|||||||
@ -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):
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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__)
|
||||||
|
|||||||
@ -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")
|
||||||
|
|||||||
@ -1,2 +1,2 @@
|
|||||||
from plugins.oneke import *
|
from src.plugins.oneke import *
|
||||||
from plugins.pdf2txt import *
|
from src.plugins.pdf2txt import *
|
||||||
@ -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:
|
||||||
|
|||||||
@ -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]
|
||||||
@ -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)
|
||||||
|
|||||||
@ -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():
|
||||||
|
|||||||
@ -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")
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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大于B,B小于C,A和C哪个大?',
|
'A大于B,B小于C,A和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>
|
||||||
|
|||||||
@ -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()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
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