feat: 添加了基于 tavil的web 搜索
This commit is contained in:
commit
31bbaa307d
5
.gitignore
vendored
5
.gitignore
vendored
@ -35,4 +35,7 @@ web/package-lock.json
|
|||||||
saves
|
saves
|
||||||
notebooks
|
notebooks
|
||||||
graphrag
|
graphrag
|
||||||
docker/volumes
|
docker/volumes
|
||||||
|
|
||||||
|
|
||||||
|
.cursorrules
|
||||||
|
|||||||
@ -12,9 +12,6 @@
|
|||||||
> [!NOTE]
|
> [!NOTE]
|
||||||
> 当前项目还处于开发的早期,还存在一些 BUG,有问题随时提 issue。
|
> 当前项目还处于开发的早期,还存在一些 BUG,有问题随时提 issue。
|
||||||
|
|
||||||
已知问题:
|
|
||||||
|
|
||||||
- [ ] 从 Flask 更换到 Fast API 之后,并行命令还存在问题。
|
|
||||||
|
|
||||||
## 概述
|
## 概述
|
||||||
|
|
||||||
|
|||||||
@ -24,4 +24,5 @@ opencv-python-headless
|
|||||||
docx2txt
|
docx2txt
|
||||||
uvicorn[standard]
|
uvicorn[standard]
|
||||||
fastapi
|
fastapi
|
||||||
python-multipart
|
python-multipart
|
||||||
|
tavily-python
|
||||||
@ -53,7 +53,7 @@ class Config(SimpleConfig):
|
|||||||
self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
|
self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
|
||||||
self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱")
|
self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱")
|
||||||
self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎")
|
self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎")
|
||||||
|
self.add_item("enable_web_search", default=False, des="是否开启网页搜索")
|
||||||
# 模型配置
|
# 模型配置
|
||||||
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
||||||
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
|
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
|
||||||
|
|||||||
17
src/config/base.yaml
Normal file
17
src/config/base.yaml
Normal file
@ -0,0 +1,17 @@
|
|||||||
|
# 基础配置
|
||||||
|
stream: true
|
||||||
|
save_dir: saves
|
||||||
|
|
||||||
|
# 功能开关
|
||||||
|
enable_reranker: false
|
||||||
|
enable_knowledge_base: false
|
||||||
|
enable_knowledge_graph: false
|
||||||
|
enable_search_engine: false
|
||||||
|
enable_web_search: false
|
||||||
|
|
||||||
|
# 模型配置
|
||||||
|
model_provider: "deepseek" # 设置为 deepseek
|
||||||
|
model_name: "deepseek-chat" # 设置默认模型名称
|
||||||
|
embed_model: "zhipu-embedding-3"
|
||||||
|
reranker: "bge-reranker-v2-m3"
|
||||||
|
model_local_paths: {}
|
||||||
@ -18,6 +18,7 @@ MODEL_NAMES:
|
|||||||
- DEEPSEEK_API_KEY
|
- DEEPSEEK_API_KEY
|
||||||
models:
|
models:
|
||||||
- deepseek-chat
|
- deepseek-chat
|
||||||
|
- deepseek-reasoner
|
||||||
zhipu:
|
zhipu:
|
||||||
name: 智谱AI (Zhipu)
|
name: 智谱AI (Zhipu)
|
||||||
url: https://open.bigmodel.cn/dev/api
|
url: https://open.bigmodel.cn/dev/api
|
||||||
|
|||||||
@ -1,3 +1,4 @@
|
|||||||
|
import os
|
||||||
from src.core import DataBaseManager
|
from src.core import DataBaseManager
|
||||||
from src.core.retriever import Retriever
|
from src.core.retriever import Retriever
|
||||||
from src.models import select_model
|
from src.models import select_model
|
||||||
@ -9,10 +10,29 @@ logger = setup_logger("Startup")
|
|||||||
|
|
||||||
class Startup:
|
class Startup:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
|
self.config = Config("config/base.yaml")
|
||||||
|
self._check_environment()
|
||||||
self.start()
|
self.start()
|
||||||
|
|
||||||
|
def _check_environment(self):
|
||||||
|
"""检查必要的环境变量"""
|
||||||
|
required_vars = {
|
||||||
|
"zhipu": ["ZHIPUAI_API_KEY"],
|
||||||
|
"openai": ["OPENAI_API_KEY"],
|
||||||
|
"deepseek": ["DEEPSEEK_API_KEY"],
|
||||||
|
}
|
||||||
|
|
||||||
|
provider = self.config.model_provider
|
||||||
|
if provider in required_vars:
|
||||||
|
missing = [var for var in required_vars[provider] if not os.getenv(var)]
|
||||||
|
if missing:
|
||||||
|
logger.error(f"Missing required environment variables for {provider}: {missing}")
|
||||||
|
raise ValueError(f"Missing required environment variables: {missing}")
|
||||||
|
|
||||||
|
if self.config.enable_web_search and not os.getenv("TAVILY_API_KEY"):
|
||||||
|
logger.warning("TAVILY_API_KEY not set, web search will be disabled")
|
||||||
|
|
||||||
def start(self):
|
def start(self):
|
||||||
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,46 +1,23 @@
|
|||||||
from src.utils.logging_config import logger
|
from src.utils.logging_config import logger
|
||||||
|
from src.models.chat_model import OpenModel, DeepSeek, Zhipu, Qianfan, DashScope, SiliconFlow
|
||||||
|
from src.models.embedding import get_embedding_model
|
||||||
|
|
||||||
|
|
||||||
def select_model(config):
|
def select_model(config):
|
||||||
|
"""
|
||||||
model_provider = config.model_provider
|
根据配置选择模型
|
||||||
model_name = config.model_name
|
"""
|
||||||
|
if config.model_provider == "deepseek":
|
||||||
logger.info(f"Selecting model from {model_provider} with {model_name}")
|
return DeepSeek(config.model_name)
|
||||||
|
elif config.model_provider == "zhipu":
|
||||||
if model_provider == "deepseek":
|
return Zhipu(config.model_name)
|
||||||
from src.models.chat_model import DeepSeek
|
elif config.model_provider == "openai":
|
||||||
return DeepSeek(model_name)
|
return OpenModel(config.model_name)
|
||||||
|
elif config.model_provider == "qianfan":
|
||||||
elif model_provider == "zhipu":
|
return Qianfan(config.model_name)
|
||||||
from src.models.chat_model import Zhipu
|
elif config.model_provider == "dashscope":
|
||||||
return Zhipu(model_name)
|
return DashScope(config.model_name)
|
||||||
|
elif config.model_provider == "siliconflow":
|
||||||
elif model_provider == "qianfan":
|
return SiliconFlow(config.model_name)
|
||||||
from src.models.chat_model import Qianfan
|
|
||||||
return Qianfan(model_name)
|
|
||||||
|
|
||||||
elif model_provider == "dashscope":
|
|
||||||
from src.models.chat_model import DashScope
|
|
||||||
return DashScope(model_name)
|
|
||||||
|
|
||||||
elif model_provider == "openai":
|
|
||||||
from src.models.chat_model import OpenModel
|
|
||||||
return OpenModel(model_name)
|
|
||||||
|
|
||||||
elif model_provider == "siliconflow":
|
|
||||||
from src.models.chat_model import SiliconFlow
|
|
||||||
return SiliconFlow(model_name)
|
|
||||||
|
|
||||||
elif model_provider == "custom":
|
|
||||||
model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None)
|
|
||||||
if model_info is None:
|
|
||||||
raise ValueError(f"Model {model_name} not found in custom models")
|
|
||||||
|
|
||||||
from src.models.chat_model import CustomModel
|
|
||||||
return CustomModel(model_info)
|
|
||||||
|
|
||||||
elif model_provider is None:
|
|
||||||
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
|
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Model provider {model_provider} not supported")
|
raise ValueError(f"Unsupported model provider: {config.model_provider}")
|
||||||
|
|||||||
@ -1,6 +1,7 @@
|
|||||||
import os
|
import os
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
from src.utils.logging_config import setup_logger
|
from src.utils.logging_config import setup_logger
|
||||||
|
from zhipuai import ZhipuAI
|
||||||
|
|
||||||
|
|
||||||
logger = setup_logger(__name__)
|
logger = setup_logger(__name__)
|
||||||
@ -50,8 +51,8 @@ class OpenModel(OpenAIBase):
|
|||||||
class DeepSeek(OpenAIBase):
|
class DeepSeek(OpenAIBase):
|
||||||
def __init__(self, model_name=None):
|
def __init__(self, model_name=None):
|
||||||
model_name = model_name or "deepseek-chat"
|
model_name = model_name or "deepseek-chat"
|
||||||
api_key = os.getenv("DEEPSEEK_API_KEY")
|
api_key = os.getenv("DEEPSEEK_API_KEY", "your-default-api-key")
|
||||||
base_url = "https://api.deepseek.com"
|
base_url = os.getenv("DEEPSEEK_API_BASE", "https://api.deepseek.com/v1")
|
||||||
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
|
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
|
||||||
|
|
||||||
|
|
||||||
@ -165,6 +166,14 @@ class DashScope:
|
|||||||
return response.output.choices[0].message
|
return response.output.choices[0].message
|
||||||
|
|
||||||
|
|
||||||
|
class ChatModel:
|
||||||
|
def __init__(self, config):
|
||||||
|
if config.model_provider == "zhipu":
|
||||||
|
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY"))
|
||||||
|
elif config.model_provider == "openai":
|
||||||
|
self.client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
model = SiliconFlow()
|
model = SiliconFlow()
|
||||||
for a in model.predict("你好", stream=True):
|
for a in model.predict("你好", stream=True):
|
||||||
|
|||||||
179
src/models/ollama_embedding.py
Normal file
179
src/models/ollama_embedding.py
Normal file
@ -0,0 +1,179 @@
|
|||||||
|
import os
|
||||||
|
import requests
|
||||||
|
import numpy as np
|
||||||
|
from typing import List, Union, Dict
|
||||||
|
from src.utils.logging_config import setup_logger
|
||||||
|
|
||||||
|
logger = setup_logger("OllamaEmbedding")
|
||||||
|
|
||||||
|
class OllamaEmbedding:
|
||||||
|
"""
|
||||||
|
使用 Ollama API 进行文本嵌入的类
|
||||||
|
"""
|
||||||
|
def __init__(self, model_info: Dict, config) -> None:
|
||||||
|
"""
|
||||||
|
初始化 Ollama Embedding 模型
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_info: 模型信息字典
|
||||||
|
config: 配置对象
|
||||||
|
"""
|
||||||
|
self.config = config
|
||||||
|
self.model_info = model_info
|
||||||
|
self.base_url = os.getenv("OLLAMA_BASE_URL", "http://localhost:11434")
|
||||||
|
self.model_name = model_info.get("name", "nomic-embed-text")
|
||||||
|
self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:"
|
||||||
|
logger.info(f"Ollama Embedding model {self.model_name} initialized")
|
||||||
|
|
||||||
|
def _get_embedding(self, text: str) -> List[float]:
|
||||||
|
"""
|
||||||
|
获取单个文本的嵌入向量
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: 输入文本
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
嵌入向量
|
||||||
|
"""
|
||||||
|
url = f"{self.base_url}/api/embeddings"
|
||||||
|
try:
|
||||||
|
response = requests.post(url, json={
|
||||||
|
"model": self.model_name,
|
||||||
|
"prompt": text
|
||||||
|
})
|
||||||
|
response.raise_for_status()
|
||||||
|
return response.json()["embedding"]
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting embedding: {str(e)}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
def predict(self, messages: List[str]) -> List[List[float]]:
|
||||||
|
"""
|
||||||
|
批量获取文本嵌入向量
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: 文本列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
嵌入向量列表
|
||||||
|
"""
|
||||||
|
embeddings = []
|
||||||
|
batch_size = 20
|
||||||
|
|
||||||
|
for i in range(0, len(messages), batch_size):
|
||||||
|
batch = messages[i:i + batch_size]
|
||||||
|
logger.info(f"Processing batch {i//batch_size + 1}, size: {len(batch)}")
|
||||||
|
|
||||||
|
batch_embeddings = []
|
||||||
|
for text in batch:
|
||||||
|
embedding = self._get_embedding(text)
|
||||||
|
batch_embeddings.append(embedding)
|
||||||
|
|
||||||
|
embeddings.extend(batch_embeddings)
|
||||||
|
|
||||||
|
return embeddings
|
||||||
|
|
||||||
|
def encode(self, messages: Union[str, List[str]]) -> List[List[float]]:
|
||||||
|
"""
|
||||||
|
编码文本
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: 单个文本或文本列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
嵌入向量列表
|
||||||
|
"""
|
||||||
|
if isinstance(messages, str):
|
||||||
|
messages = [messages]
|
||||||
|
return self.predict(messages)
|
||||||
|
|
||||||
|
def encode_queries(self, queries: List[str]) -> List[List[float]]:
|
||||||
|
"""
|
||||||
|
编码查询文本
|
||||||
|
|
||||||
|
Args:
|
||||||
|
queries: 查询文本列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
查询文本的嵌入向量列表
|
||||||
|
"""
|
||||||
|
return self.predict(queries)
|
||||||
|
|
||||||
|
|
||||||
|
class OllamaReranker:
|
||||||
|
"""
|
||||||
|
使用 Ollama API 进行文本重排序的类
|
||||||
|
"""
|
||||||
|
def __init__(self, config) -> None:
|
||||||
|
"""
|
||||||
|
初始化 Ollama Reranker
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: 配置对象
|
||||||
|
"""
|
||||||
|
self.config = config
|
||||||
|
self.base_url = os.getenv("OLLAMA_BASE_URL", "http://localhost:11434")
|
||||||
|
self.model_name = config.reranker
|
||||||
|
logger.info(f"Ollama Reranker model {self.model_name} initialized")
|
||||||
|
|
||||||
|
def compute_score(self, query: str, passage: str) -> float:
|
||||||
|
"""
|
||||||
|
计算查询和文本段落之间的相关性分数
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: 查询文本
|
||||||
|
passage: 段落文本
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
相关性分数
|
||||||
|
"""
|
||||||
|
prompt = f"Query: {query}\nPassage: {passage}\nRate the relevance of the passage to the query on a scale of 0 to 1:"
|
||||||
|
|
||||||
|
try:
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/api/generate",
|
||||||
|
json={
|
||||||
|
"model": self.model_name,
|
||||||
|
"prompt": prompt,
|
||||||
|
"stream": False
|
||||||
|
}
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
# 提取生成的数字作为分数
|
||||||
|
result = response.json()["response"].strip()
|
||||||
|
try:
|
||||||
|
score = float(result)
|
||||||
|
return min(max(score, 0), 1) # 确保分数在 0-1 之间
|
||||||
|
except ValueError:
|
||||||
|
logger.warning(f"Could not parse score from response: {result}")
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error computing rerank score: {str(e)}")
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
def rerank(self, query: str, passages: List[str], top_n: int = None) -> List[Dict]:
|
||||||
|
"""
|
||||||
|
重新排序文本段落
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: 查询文本
|
||||||
|
passages: 段落文本列表
|
||||||
|
top_n: 返回前 n 个结果
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
排序后的结果列表,每个元素包含索引和分数
|
||||||
|
"""
|
||||||
|
scores = []
|
||||||
|
for i, passage in enumerate(passages):
|
||||||
|
score = self.compute_score(query, passage)
|
||||||
|
scores.append({"index": i, "score": score})
|
||||||
|
|
||||||
|
# 按分数降序排序
|
||||||
|
sorted_results = sorted(scores, key=lambda x: x["score"], reverse=True)
|
||||||
|
|
||||||
|
if top_n:
|
||||||
|
sorted_results = sorted_results[:top_n]
|
||||||
|
|
||||||
|
return sorted_results
|
||||||
@ -6,6 +6,7 @@ from concurrent.futures import ThreadPoolExecutor
|
|||||||
from src.core import HistoryManager
|
from src.core import HistoryManager
|
||||||
from src.core.startup import startup
|
from src.core.startup import startup
|
||||||
from src.utils.logging_config import setup_logger
|
from src.utils.logging_config import setup_logger
|
||||||
|
from src.utils.web_search import WebSearcher
|
||||||
|
|
||||||
chat = APIRouter(prefix="/chat")
|
chat = APIRouter(prefix="/chat")
|
||||||
logger = setup_logger("server-chat")
|
logger = setup_logger("server-chat")
|
||||||
@ -13,6 +14,7 @@ logger = setup_logger("server-chat")
|
|||||||
executor = ThreadPoolExecutor()
|
executor = ThreadPoolExecutor()
|
||||||
|
|
||||||
refs_pool = {}
|
refs_pool = {}
|
||||||
|
web_searcher = WebSearcher()
|
||||||
|
|
||||||
@chat.get("/")
|
@chat.get("/")
|
||||||
async def chat_get():
|
async def chat_get():
|
||||||
@ -37,20 +39,50 @@ def chat_post(
|
|||||||
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
||||||
|
|
||||||
def generate_response():
|
def generate_response():
|
||||||
|
modified_query = query
|
||||||
|
|
||||||
|
# 处理网页搜索
|
||||||
|
if meta and meta.get("enable_web_search"):
|
||||||
|
chunk = make_chunk("正在进行网络搜索...", "searching", history=None)
|
||||||
|
yield chunk
|
||||||
|
|
||||||
if meta.get("enable_retrieval"):
|
try:
|
||||||
|
search_results = web_searcher.search(query)
|
||||||
|
if search_results:
|
||||||
|
search_context = web_searcher.format_search_results(search_results)
|
||||||
|
# 将搜索结果添加到查询中
|
||||||
|
modified_query = f"""基于以下网络搜索结果回答问题:
|
||||||
|
|
||||||
|
{search_context}
|
||||||
|
|
||||||
|
用户问题:{query}
|
||||||
|
|
||||||
|
请综合以上搜索结果,给出准确、客观的回答。如果搜索结果与问题相关性不大,请直接基于你的知识回答。
|
||||||
|
"""
|
||||||
|
logger.info(f"Web search results added to query")
|
||||||
|
else:
|
||||||
|
logger.warning("No web search results found")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Web search error: {str(e)}")
|
||||||
|
chunk = make_chunk("网络搜索失败,将直接回答问题。", "loading", history=None)
|
||||||
|
yield chunk
|
||||||
|
|
||||||
|
# 处理知识库检索
|
||||||
|
if meta and meta.get("enable_retrieval"):
|
||||||
chunk = make_chunk("", "searching", history=None)
|
chunk = make_chunk("", "searching", history=None)
|
||||||
yield chunk
|
yield chunk
|
||||||
|
|
||||||
|
<<<<<<< HEAD
|
||||||
meta["config"] = startup.config
|
meta["config"] = startup.config
|
||||||
new_query, refs = startup.retriever(query, history_manager.messages, meta)
|
new_query, refs = startup.retriever(query, history_manager.messages, meta)
|
||||||
|
=======
|
||||||
|
modified_query, refs = startup.retriever(modified_query, history_manager.messages, meta)
|
||||||
|
>>>>>>> feature/web_search
|
||||||
refs_pool[cur_res_id] = refs
|
refs_pool[cur_res_id] = refs
|
||||||
else:
|
|
||||||
new_query = query
|
|
||||||
|
|
||||||
messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round'))
|
messages = history_manager.get_history_with_msg(modified_query, max_rounds=meta.get('history_round'))
|
||||||
history_manager.add_user(query)
|
history_manager.add_user(query) # 注意这里使用原始查询
|
||||||
logger.debug(f"Web history: {history_manager.messages}")
|
logger.debug(f"Final query: {modified_query}")
|
||||||
|
|
||||||
content = ""
|
content = ""
|
||||||
for delta in startup.model.predict(messages, stream=True):
|
for delta in startup.model.predict(messages, stream=True):
|
||||||
|
|||||||
69
src/utils/web_search.py
Normal file
69
src/utils/web_search.py
Normal file
@ -0,0 +1,69 @@
|
|||||||
|
import os
|
||||||
|
from typing import List, Dict
|
||||||
|
from tavily import TavilyClient
|
||||||
|
from src.utils.logging_config import setup_logger
|
||||||
|
|
||||||
|
logger = setup_logger("web-search")
|
||||||
|
|
||||||
|
class WebSearcher:
|
||||||
|
def __init__(self):
|
||||||
|
api_key = os.getenv("TAVILY_API_KEY", "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")
|
||||||
|
if not api_key:
|
||||||
|
raise ValueError("TAVILY_API_KEY environment variable is not set")
|
||||||
|
self.client = TavilyClient(api_key)
|
||||||
|
logger.info("WebSearcher initialized with Tavily client")
|
||||||
|
|
||||||
|
def search(self, query: str, max_results: int = 1) -> List[Dict]:
|
||||||
|
"""
|
||||||
|
使用 Tavily 搜索相关内容
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: 搜索查询
|
||||||
|
max_results: 最大返回结果数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
搜索结果列表
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
search_results = self.client.search(
|
||||||
|
query=query,
|
||||||
|
search_depth="basic",
|
||||||
|
max_results=max_results
|
||||||
|
)
|
||||||
|
|
||||||
|
# 提取需要的信息
|
||||||
|
formatted_results = []
|
||||||
|
for result in search_results['results'][:max_results]:
|
||||||
|
formatted_results.append({
|
||||||
|
'title': result.get('title', ''),
|
||||||
|
'content': result.get('content', ''),
|
||||||
|
'url': result.get('url', ''),
|
||||||
|
'score': result.get('score', 0)
|
||||||
|
})
|
||||||
|
|
||||||
|
return formatted_results
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error during web search: {str(e)}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
def format_search_results(self, results: List[Dict]) -> str:
|
||||||
|
"""
|
||||||
|
将搜索结果格式化为文本
|
||||||
|
|
||||||
|
Args:
|
||||||
|
results: 搜索结果列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
格式化后的文本
|
||||||
|
"""
|
||||||
|
if not results:
|
||||||
|
return "没有找到相关的网络搜索结果。"
|
||||||
|
|
||||||
|
formatted_text = "以下是相关的网络搜索结果:\n\n"
|
||||||
|
for i, result in enumerate(results, 1):
|
||||||
|
formatted_text += f"{i}. {result['title']}\n"
|
||||||
|
formatted_text += f" {result['content']}\n"
|
||||||
|
formatted_text += f" 来源: {result['url']}\n\n"
|
||||||
|
|
||||||
|
return formatted_text
|
||||||
@ -81,6 +81,9 @@
|
|||||||
<div class="flex-center" @click="meta.use_web = !meta.use_web" v-if="configStore.config.enable_search_engine && meta.enable_retrieval">
|
<div class="flex-center" @click="meta.use_web = !meta.use_web" v-if="configStore.config.enable_search_engine && meta.enable_retrieval">
|
||||||
搜索引擎(Bing) <div @click.stop><a-switch v-model:checked="meta.use_web" /></div>
|
搜索引擎(Bing) <div @click.stop><a-switch v-model:checked="meta.use_web" /></div>
|
||||||
</div>
|
</div>
|
||||||
|
<div class="flex-center" @click="meta.enable_web_search = !meta.enable_web_search">
|
||||||
|
网页搜索 <div @click.stop><a-switch v-model:checked="meta.enable_web_search" /></div>
|
||||||
|
</div>
|
||||||
<!-- <div class="flex-center" v-if="configStore.config.enable_knowledge_base && meta.enable_retrieval">
|
<!-- <div class="flex-center" v-if="configStore.config.enable_knowledge_base && meta.enable_retrieval">
|
||||||
重写查询 <a-segmented v-model:value="meta.use_rewrite_query" :options="['off', 'on', 'hyde']"/>
|
重写查询 <a-segmented v-model:value="meta.use_rewrite_query" :options="['off', 'on', 'hyde']"/>
|
||||||
</div> -->
|
</div> -->
|
||||||
@ -166,6 +169,8 @@ import {
|
|||||||
GlobalOutlined,
|
GlobalOutlined,
|
||||||
FileTextOutlined,
|
FileTextOutlined,
|
||||||
RobotOutlined,
|
RobotOutlined,
|
||||||
|
EditOutlined,
|
||||||
|
PlusOutlined,
|
||||||
} from '@ant-design/icons-vue'
|
} from '@ant-design/icons-vue'
|
||||||
import { onClickOutside } from '@vueuse/core'
|
import { onClickOutside } from '@vueuse/core'
|
||||||
import { Marked } from 'marked';
|
import { Marked } from 'marked';
|
||||||
@ -206,12 +211,14 @@ const meta = reactive(JSON.parse(localStorage.getItem('meta')) || {
|
|||||||
enable_retrieval: false,
|
enable_retrieval: false,
|
||||||
use_graph: false,
|
use_graph: false,
|
||||||
use_web: false,
|
use_web: false,
|
||||||
|
enable_web_search: false,
|
||||||
graph_name: "neo4j",
|
graph_name: "neo4j",
|
||||||
// use_rewrite_query: "off",
|
// use_rewrite_query: "off",
|
||||||
selectedKB: null,
|
selectedKB: null,
|
||||||
stream: true,
|
stream: true,
|
||||||
summary_title: true,
|
summary_title: true,
|
||||||
history_round: 5,
|
history_round: 5,
|
||||||
|
db_name: null,
|
||||||
})
|
})
|
||||||
|
|
||||||
const marked = new Marked(
|
const marked = new Marked(
|
||||||
@ -323,35 +330,26 @@ const appendAiMessage = (message, refs=null) => {
|
|||||||
|
|
||||||
const updateMessage = (info) => {
|
const updateMessage = (info) => {
|
||||||
const message = conv.value.messages.find((message) => message.id === info.id);
|
const message = conv.value.messages.find((message) => message.id === info.id);
|
||||||
|
|
||||||
if (message) {
|
if (message) {
|
||||||
// 只有在 text 不为空时更新
|
try {
|
||||||
if (info.text !== null && info.text !== undefined && info.text !== '') {
|
if (info.text !== null && info.text !== undefined && info.text !== '') {
|
||||||
message.text = info.text;
|
message.text = info.text;
|
||||||
}
|
}
|
||||||
|
if (info.status !== null && info.status !== undefined && info.status !== '') {
|
||||||
// 只有在 refs 不为空时更新
|
message.status = info.status;
|
||||||
if (info.refs !== null && info.refs !== undefined) {
|
}
|
||||||
message.refs = info.refs;
|
if (info.meta !== null && info.meta !== undefined) {
|
||||||
}
|
message.meta = info.meta;
|
||||||
|
}
|
||||||
if (info.model_name !== null && info.model_name !== undefined && info.model_name !== '') {
|
scrollToBottom();
|
||||||
message.model_name = info.model_name;
|
} catch (error) {
|
||||||
}
|
console.error('Error updating message:', error);
|
||||||
|
message.status = 'error';
|
||||||
// 只有在 status 不为空时更新
|
message.text = '消息更新失败';
|
||||||
if (info.status !== null && info.status !== undefined && info.status !== '') {
|
|
||||||
message.status = info.status;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (info.meta !== null && info.meta !== undefined) {
|
|
||||||
message.meta = info.meta;
|
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
console.error('Message not found');
|
console.error('Message not found:', info.id);
|
||||||
}
|
}
|
||||||
|
|
||||||
scrollToBottom();
|
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|
||||||
@ -395,21 +393,21 @@ const loadDatabases = () => {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 新函数用于处理 fetch 请求
|
// 新函数用于处理 fetch 请求
|
||||||
const fetchChatResponse = (user_input, cur_res_id) => {
|
const fetchChatResponse = (requestData) => {
|
||||||
fetch('/api/chat/', {
|
fetch('/api/chat/', {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
body: JSON.stringify({
|
|
||||||
query: user_input,
|
|
||||||
history: conv.value.history,
|
|
||||||
meta: meta,
|
|
||||||
cur_res_id: cur_res_id,
|
|
||||||
}),
|
|
||||||
headers: {
|
headers: {
|
||||||
'Content-Type': 'application/json'
|
'Content-Type': 'application/json'
|
||||||
}
|
},
|
||||||
|
body: JSON.stringify(requestData)
|
||||||
})
|
})
|
||||||
.then((response) => {
|
.then(response => {
|
||||||
if (!response.body) throw new Error("ReadableStream not supported.");
|
if (!response.ok) {
|
||||||
|
throw new Error(`HTTP error! status: ${response.status}`);
|
||||||
|
}
|
||||||
|
if (!response.body) {
|
||||||
|
throw new Error("ReadableStream not supported.");
|
||||||
|
}
|
||||||
const reader = response.body.getReader();
|
const reader = response.body.getReader();
|
||||||
const decoder = new TextDecoder("utf-8");
|
const decoder = new TextDecoder("utf-8");
|
||||||
let buffer = '';
|
let buffer = '';
|
||||||
@ -417,22 +415,22 @@ const fetchChatResponse = (user_input, cur_res_id) => {
|
|||||||
const readChunk = () => {
|
const readChunk = () => {
|
||||||
return reader.read().then(({ done, value }) => {
|
return reader.read().then(({ done, value }) => {
|
||||||
if (done) {
|
if (done) {
|
||||||
const message = conv.value.messages.find((message) => message.id === cur_res_id)
|
const message = conv.value.messages.find((message) => message.id === requestData.cur_res_id)
|
||||||
console.log(message)
|
console.log(message)
|
||||||
if (message.meta.enable_retrieval) {
|
if (message.meta.enable_retrieval) {
|
||||||
console.log("fetching refs")
|
console.log("fetching refs")
|
||||||
fetchRefs(cur_res_id).then((data) => {
|
fetchRefs(requestData.cur_res_id).then((data) => {
|
||||||
console.log(data)
|
console.log(data)
|
||||||
updateMessage({
|
updateMessage({
|
||||||
id: cur_res_id,
|
id: requestData.cur_res_id,
|
||||||
refs: data,
|
refs: data,
|
||||||
status: "finished",
|
status: "finished",
|
||||||
});
|
});
|
||||||
groupRefs(cur_res_id);
|
groupRefs(requestData.cur_res_id);
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
updateMessage({
|
updateMessage({
|
||||||
id: cur_res_id,
|
id: requestData.cur_res_id,
|
||||||
status: "finished",
|
status: "finished",
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
@ -451,7 +449,7 @@ const fetchChatResponse = (user_input, cur_res_id) => {
|
|||||||
try {
|
try {
|
||||||
const data = JSON.parse(line);
|
const data = JSON.parse(line);
|
||||||
updateMessage({
|
updateMessage({
|
||||||
id: cur_res_id,
|
id: requestData.cur_res_id,
|
||||||
text: data.response,
|
text: data.response,
|
||||||
model_name: data.model_name,
|
model_name: data.model_name,
|
||||||
status: data.status,
|
status: data.status,
|
||||||
@ -478,12 +476,14 @@ const fetchChatResponse = (user_input, cur_res_id) => {
|
|||||||
readChunk();
|
readChunk();
|
||||||
})
|
})
|
||||||
.catch((error) => {
|
.catch((error) => {
|
||||||
console.error(error);
|
console.error('Error in fetchChatResponse:', error);
|
||||||
updateMessage({
|
updateMessage({
|
||||||
id: cur_res_id,
|
id: requestData.cur_res_id,
|
||||||
status: "error",
|
status: "error",
|
||||||
|
text: `请求错误:${error.message}`,
|
||||||
});
|
});
|
||||||
isStreaming.value = false;
|
isStreaming.value = false;
|
||||||
|
message.error(`请求失败:${error.message}`);
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -514,9 +514,30 @@ const sendMessage = () => {
|
|||||||
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;
|
||||||
conv.value.inputText = '';
|
conv.value.inputText = '';
|
||||||
meta.db_name = dbName;
|
|
||||||
|
// 准备发送的数据
|
||||||
|
const requestData = {
|
||||||
|
query: user_input,
|
||||||
|
history: conv.value.history,
|
||||||
|
cur_res_id: cur_res_id,
|
||||||
|
meta: {
|
||||||
|
enable_retrieval: meta.enable_retrieval,
|
||||||
|
use_graph: meta.use_graph,
|
||||||
|
use_web: meta.use_web,
|
||||||
|
enable_web_search: meta.enable_web_search,
|
||||||
|
graph_name: meta.graph_name,
|
||||||
|
rewriteQuery: meta.rewriteQuery,
|
||||||
|
selectedKB: meta.selectedKB,
|
||||||
|
stream: meta.stream,
|
||||||
|
summary_title: meta.summary_title,
|
||||||
|
history_round: meta.history_round,
|
||||||
|
db_name: dbName,
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
fetchChatResponse(user_input, cur_res_id)
|
console.log('Sending request with data:', requestData); // 添加日志
|
||||||
|
|
||||||
|
fetchChatResponse(requestData);
|
||||||
} else {
|
} else {
|
||||||
console.log('请输入消息');
|
console.log('请输入消息');
|
||||||
}
|
}
|
||||||
@ -644,6 +665,17 @@ watch(
|
|||||||
&:hover {
|
&:hover {
|
||||||
background-color: var(--main-light-3);
|
background-color: var(--main-light-3);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.anticon {
|
||||||
|
margin-right: 8px;
|
||||||
|
font-size: 16px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.ant-switch {
|
||||||
|
&.ant-switch-checked {
|
||||||
|
background-color: var(--main-500);
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -973,7 +1005,15 @@ watch(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.controls {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 8px;
|
||||||
|
|
||||||
|
.search-switch {
|
||||||
|
margin-right: 8px;
|
||||||
|
}
|
||||||
|
}
|
||||||
</style>
|
</style>
|
||||||
|
|
||||||
<style lang="less">
|
<style lang="less">
|
||||||
|
|||||||
@ -11,7 +11,7 @@ const router = createRouter({
|
|||||||
component: BlankLayout,
|
component: BlankLayout,
|
||||||
children: [ {
|
children: [ {
|
||||||
path: '',
|
path: '',
|
||||||
name: 'home',
|
name: 'Home',
|
||||||
component: () => import('../views/HomeView.vue'),
|
component: () => import('../views/HomeView.vue'),
|
||||||
meta: { keepAlive: true }
|
meta: { keepAlive: true }
|
||||||
}
|
}
|
||||||
@ -50,7 +50,7 @@ const router = createRouter({
|
|||||||
children: [
|
children: [
|
||||||
{
|
{
|
||||||
path: '',
|
path: '',
|
||||||
name: 'database',
|
name: 'Database',
|
||||||
component: () => import('../views/DataBaseView.vue'),
|
component: () => import('../views/DataBaseView.vue'),
|
||||||
meta: { keepAlive: true }
|
meta: { keepAlive: true }
|
||||||
},
|
},
|
||||||
@ -69,7 +69,7 @@ const router = createRouter({
|
|||||||
children: [
|
children: [
|
||||||
{
|
{
|
||||||
path: '',
|
path: '',
|
||||||
name: 'setting',
|
name: 'Setting',
|
||||||
component: () => import('../views/SettingView.vue'),
|
component: () => import('../views/SettingView.vue'),
|
||||||
meta: { keepAlive: true }
|
meta: { keepAlive: true }
|
||||||
}
|
}
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user