From 6c5c099132c41f3e984b679337ac3b08c0ec93d5 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Tue, 16 Jul 2024 23:12:35 +0800 Subject: [PATCH] update details --- src/core/__init__.py | 4 +--- src/core/retriever.py | 13 ++++++++++--- src/core/startup.py | 4 +--- src/utils/logging_config.py | 2 +- src/views/common_view.py | 10 ++++++---- src/views/database_view.py | 2 +- 6 files changed, 20 insertions(+), 15 deletions(-) diff --git a/src/core/__init__.py b/src/core/__init__.py index 8f03f119..d59f3c6e 100644 --- a/src/core/__init__.py +++ b/src/core/__init__.py @@ -1,4 +1,2 @@ from .history import * -from .retriever import * -from .database import * -from .graphbase import * \ No newline at end of file +from .database import * \ No newline at end of file diff --git a/src/core/retriever.py b/src/core/retriever.py index 03fa6741..896de456 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,9 +1,11 @@ +from core.startup import dbm, model + class Retriever: def __init__(self, config): self.config = config - def retrieval(self, query): + def retrieval(self, query, history): refs = {} @@ -37,11 +39,16 @@ class Retriever: """ raise NotImplementedError + def query_graph(self, query, history): + # res = model.predict("qiansdgsa, dasdh ashdsakjdk ak ").content + + return {} + def rewrite_query(self, query): """重写查询""" raise NotImplementedError - def __call__(self, query): - refs = self.retrieval(query) + def __call__(self, query, history): + refs = self.retrieval(query, history) query = self.construct_query(query, refs) return query, refs \ No newline at end of file diff --git a/src/core/startup.py b/src/core/startup.py index 4707f2aa..70d0ae7a 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -1,5 +1,4 @@ -from core import Retriever, DataBaseManager -from core.graphbase import GraphDatabase +from core import DataBaseManager from models import select_model from config import Config @@ -7,6 +6,5 @@ from config import Config config = Config("config/base.yaml") model = select_model(config) dbm = DataBaseManager(config) -retriever = Retriever(config) # 启动本地图数据库 \ No newline at end of file diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index 90e10c77..9368aa3e 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -6,7 +6,7 @@ from datetime import datetime # DATETIME = datetime.now().strftime('%Y-%m-%d-%H%M%S') DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件 -def setup_logger(name, log_file=None, level=logging.DEBUG, console=True): +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' diff --git a/src/views/common_view.py b/src/views/common_view.py index 17398b73..cfb9e5a2 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -3,11 +3,13 @@ from flask import Blueprint, jsonify, request, Response from core import HistoryManager from utils.logging_config import setup_logger -from core.startup import config, model, retriever +from core.startup import config, model +from core.retriever import Retriever common = Blueprint('common', __name__) logger = setup_logger("server-common") +retriever = Retriever(config) @common.route('/', methods=["GET"]) def route_index(): @@ -30,10 +32,10 @@ def chat(): request_data = json.loads(request.data) query = request_data['query'] logger.debug(f"Web query: {query}") - - new_query, refs = retriever(query) - history_manager = HistoryManager(request_data['history']) + + new_query, refs = retriever(query, history_manager.messages) + messages = history_manager.get_history_with_msg(new_query) history_manager.add_user(query) logger.debug(f"Web history: {history_manager}") diff --git a/src/views/database_view.py b/src/views/database_view.py index 6022bf3b..4301a1ed 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -5,7 +5,7 @@ from flask import Blueprint, jsonify, request, Response from core import HistoryManager from utils.logging_config import setup_logger -from core.startup import config, model, retriever, dbm +from core.startup import config, model, dbm db = Blueprint('database', __name__, url_prefix="/database")