This commit is contained in:
Wenjie Zhang 2024-07-09 16:38:05 +08:00
parent b9d6d13d72
commit a9124a1881
4 changed files with 10 additions and 10 deletions

View File

@ -4,7 +4,7 @@ from flask import Flask, jsonify, Response, request
from flask_cors import CORS
from dotenv import load_dotenv
from core import HistoryManager
from core import PreRetrival
from core import PreRetrieval
from config import Config
from models import select_model
from utils.logging_config import setup_logger
@ -16,7 +16,7 @@ logger = setup_logger("server")
config = Config("config/base.yaml")
model = select_model(config)
pre_retrival = PreRetrival(config)
pre_retrieval = PreRetrieval(config)
apps = Flask(__name__)# 这段代码是为了解决跨域问题Flask默认不支持跨域
@ -48,7 +48,7 @@ def chat():
external = ""
if config.enable_knowledge_base:
kb_res = pre_retrival.search(query)
kb_res = pre_retrieval.search(query)
if kb_res:
kb_res = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res[0]])
kb_res = f"知识库信息: {kb_res}"

View File

@ -1,7 +1,7 @@
import os
from dotenv import load_dotenv
from core import HistoryManager
from core import PreRetrival
from core import PreRetrieval
from config import Config
from models import select_model
@ -11,8 +11,8 @@ load_dotenv()
if __name__ == "__main__":
config = Config("config/base.yaml")
model = select_model(config)
pre_retrival = PreRetrival(config)
# pre_retrival.add_file("/home/zwj/workspace/ProjectAthena/src/data/file/鉴定工作报告、技术报告-0708.pdf")
pre_retrieval = PreRetrieval(config)
# pre_retrieval.add_file("/home/zwj/workspace/ProjectAthena/src/data/file/鉴定工作报告、技术报告-0708.pdf")
print(f"[{config.model_provider}:{config.get('model_name', 'default')}] Type 'exit' to quit")
@ -25,7 +25,7 @@ if __name__ == "__main__":
external = ""
if config.enable_knowledge_base:
kb_res = pre_retrival.search(message)
kb_res = pre_retrieval.search(message)
if kb_res:
kb_res = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res[0]])
kb_res = f"知识库信息: {kb_res}"

View File

@ -1,2 +1,2 @@
from .history import *
from .preretrival import *
from .preretrieval import *

View File

@ -10,7 +10,7 @@ from utils.logging_config import setup_logger
from pymilvus import MilvusClient
logger = setup_logger("PreRetrival")
logger = setup_logger("PreRetrieval")
def pdfreader(file_path):
"""读取PDF文件并返回text文本"""
@ -32,7 +32,7 @@ def plainreader(file_path):
return text
class PreRetrival:
class PreRetrieval:
def __init__(self, config):
self.config = config