update details
This commit is contained in:
parent
d00056c74e
commit
6c5c099132
@ -1,4 +1,2 @@
|
|||||||
from .history import *
|
from .history import *
|
||||||
from .retriever import *
|
|
||||||
from .database import *
|
from .database import *
|
||||||
from .graphbase import *
|
|
||||||
@ -1,9 +1,11 @@
|
|||||||
|
from core.startup import dbm, model
|
||||||
|
|
||||||
class Retriever:
|
class Retriever:
|
||||||
|
|
||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
def retrieval(self, query):
|
def retrieval(self, query, history):
|
||||||
|
|
||||||
refs = {}
|
refs = {}
|
||||||
|
|
||||||
@ -37,11 +39,16 @@ class Retriever:
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def query_graph(self, query, history):
|
||||||
|
# res = model.predict("qiansdgsa, dasdh ashdsakjdk ak ").content
|
||||||
|
|
||||||
|
return {}
|
||||||
|
|
||||||
def rewrite_query(self, query):
|
def rewrite_query(self, query):
|
||||||
"""重写查询"""
|
"""重写查询"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __call__(self, query):
|
def __call__(self, query, history):
|
||||||
refs = self.retrieval(query)
|
refs = self.retrieval(query, history)
|
||||||
query = self.construct_query(query, refs)
|
query = self.construct_query(query, refs)
|
||||||
return query, refs
|
return query, refs
|
||||||
@ -1,5 +1,4 @@
|
|||||||
from core import Retriever, DataBaseManager
|
from core import DataBaseManager
|
||||||
from core.graphbase import GraphDatabase
|
|
||||||
from models import select_model
|
from models import select_model
|
||||||
from config import Config
|
from config import Config
|
||||||
|
|
||||||
@ -7,6 +6,5 @@ from config import Config
|
|||||||
config = Config("config/base.yaml")
|
config = Config("config/base.yaml")
|
||||||
model = select_model(config)
|
model = select_model(config)
|
||||||
dbm = DataBaseManager(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 = datetime.now().strftime('%Y-%m-%d-%H%M%S')
|
||||||
DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件
|
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:
|
if log_file is None:
|
||||||
log_file = f'output/log/project-{DATETIME}.log'
|
log_file = f'output/log/project-{DATETIME}.log'
|
||||||
|
|||||||
@ -3,11 +3,13 @@ from flask import Blueprint, jsonify, request, Response
|
|||||||
|
|
||||||
from core import HistoryManager
|
from core import HistoryManager
|
||||||
from utils.logging_config import setup_logger
|
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__)
|
common = Blueprint('common', __name__)
|
||||||
logger = setup_logger("server-common")
|
logger = setup_logger("server-common")
|
||||||
|
retriever = Retriever(config)
|
||||||
|
|
||||||
@common.route('/', methods=["GET"])
|
@common.route('/', methods=["GET"])
|
||||||
def route_index():
|
def route_index():
|
||||||
@ -30,10 +32,10 @@ def chat():
|
|||||||
request_data = json.loads(request.data)
|
request_data = json.loads(request.data)
|
||||||
query = request_data['query']
|
query = request_data['query']
|
||||||
logger.debug(f"Web query: {query}")
|
logger.debug(f"Web query: {query}")
|
||||||
|
|
||||||
new_query, refs = retriever(query)
|
|
||||||
|
|
||||||
history_manager = HistoryManager(request_data['history'])
|
history_manager = HistoryManager(request_data['history'])
|
||||||
|
|
||||||
|
new_query, refs = retriever(query, history_manager.messages)
|
||||||
|
|
||||||
messages = history_manager.get_history_with_msg(new_query)
|
messages = history_manager.get_history_with_msg(new_query)
|
||||||
history_manager.add_user(query)
|
history_manager.add_user(query)
|
||||||
logger.debug(f"Web history: {history_manager}")
|
logger.debug(f"Web history: {history_manager}")
|
||||||
|
|||||||
@ -5,7 +5,7 @@ from flask import Blueprint, jsonify, request, Response
|
|||||||
|
|
||||||
from core import HistoryManager
|
from core import HistoryManager
|
||||||
from utils.logging_config import setup_logger
|
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")
|
db = Blueprint('database', __name__, url_prefix="/database")
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user