update details
This commit is contained in:
parent
d00056c74e
commit
6c5c099132
@ -1,4 +1,2 @@
|
||||
from .history import *
|
||||
from .retriever import *
|
||||
from .database import *
|
||||
from .graphbase import *
|
||||
@ -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
|
||||
@ -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)
|
||||
|
||||
# 启动本地图数据库
|
||||
@ -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'
|
||||
|
||||
@ -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}")
|
||||
|
||||
@ -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")
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user