feat: 在依赖组中添加 ruff 依赖,并优化代码中的类型提示和异常处理,提升代码可读性和稳定性

This commit is contained in:
Wenjie Zhang 2025-07-02 02:38:36 +08:00
parent c5ca7f7f49
commit 6810f1f84e
14 changed files with 164 additions and 389 deletions

View File

@ -56,6 +56,7 @@ lint.ignore = ["F401"] # 忽略的规则
[dependency-groups] [dependency-groups]
dev = [ dev = [
"ruff>=0.12.1",
"vllm>=0.8.5.post1", "vllm>=0.8.5.post1",
] ]

View File

@ -35,19 +35,19 @@ async def get_subgraph(
""" """
try: try:
logger.info(f"获取子图数据 - db_id: {db_id}, node_label: {node_label}, max_depth: {max_depth}, max_nodes: {max_nodes}") logger.info(f"获取子图数据 - db_id: {db_id}, node_label: {node_label}, max_depth: {max_depth}, max_nodes: {max_nodes}")
# 获取 LightRAG 实例 # 获取 LightRAG 实例
rag_instance = await knowledge_base._get_lightrag_instance(db_id) rag_instance = await knowledge_base._get_lightrag_instance(db_id)
if not rag_instance: if not rag_instance:
raise HTTPException(status_code=404, detail=f"数据库 {db_id} 不存在") raise HTTPException(status_code=404, detail=f"数据库 {db_id} 不存在")
# 使用 LightRAG 的原生 get_knowledge_graph 方法 # 使用 LightRAG 的原生 get_knowledge_graph 方法
knowledge_graph = await rag_instance.get_knowledge_graph( knowledge_graph = await rag_instance.get_knowledge_graph(
node_label=node_label, node_label=node_label,
max_depth=max_depth, max_depth=max_depth,
max_nodes=max_nodes max_nodes=max_nodes
) )
# 将 LightRAG 的 KnowledgeGraph 格式转换为前端需要的格式 # 将 LightRAG 的 KnowledgeGraph 格式转换为前端需要的格式
nodes = [] nodes = []
for node in knowledge_graph.nodes: for node in knowledge_graph.nodes:
@ -57,7 +57,7 @@ async def get_subgraph(
"entity_type": node.properties.get("entity_type", "unknown"), "entity_type": node.properties.get("entity_type", "unknown"),
"properties": node.properties "properties": node.properties
}) })
edges = [] edges = []
for edge in knowledge_graph.edges: for edge in knowledge_graph.edges:
edges.append({ edges.append({
@ -67,7 +67,7 @@ async def get_subgraph(
"type": edge.type, "type": edge.type,
"properties": edge.properties "properties": edge.properties
}) })
result = { result = {
"success": True, "success": True,
"data": { "data": {
@ -78,10 +78,10 @@ async def get_subgraph(
"total_edges": len(edges) "total_edges": len(edges)
} }
} }
logger.info(f"成功获取子图 - 节点数: {len(nodes)}, 边数: {len(edges)}") logger.info(f"成功获取子图 - 节点数: {len(nodes)}, 边数: {len(edges)}")
return result return result
except Exception as e: except Exception as e:
logger.error(f"获取子图数据失败: {e}") logger.error(f"获取子图数据失败: {e}")
logger.error(f"Traceback: {traceback.format_exc()}") logger.error(f"Traceback: {traceback.format_exc()}")
@ -104,22 +104,22 @@ async def get_graph_labels(
""" """
try: try:
logger.info(f"获取图谱标签 - db_id: {db_id}") logger.info(f"获取图谱标签 - db_id: {db_id}")
# 获取 LightRAG 实例 # 获取 LightRAG 实例
rag_instance = await knowledge_base._get_lightrag_instance(db_id) rag_instance = await knowledge_base._get_lightrag_instance(db_id)
if not rag_instance: if not rag_instance:
raise HTTPException(status_code=404, detail=f"数据库 {db_id} 不存在") raise HTTPException(status_code=404, detail=f"数据库 {db_id} 不存在")
# 使用 LightRAG 的原生方法获取所有标签 # 使用 LightRAG 的原生方法获取所有标签
labels = await rag_instance.get_graph_labels() labels = await rag_instance.get_graph_labels()
return { return {
"success": True, "success": True,
"data": { "data": {
"labels": labels "labels": labels
} }
} }
except Exception as e: except Exception as e:
logger.error(f"获取图谱标签失败: {e}") logger.error(f"获取图谱标签失败: {e}")
logger.error(f"Traceback: {traceback.format_exc()}") logger.error(f"Traceback: {traceback.format_exc()}")
@ -142,7 +142,7 @@ async def get_available_databases(
"success": True, "success": True,
"data": databases "data": databases
} }
except Exception as e: except Exception as e:
logger.error(f"获取数据库列表失败: {e}") logger.error(f"获取数据库列表失败: {e}")
raise HTTPException(status_code=500, detail=f"获取数据库列表失败: {str(e)}") raise HTTPException(status_code=500, detail=f"获取数据库列表失败: {str(e)}")
@ -154,8 +154,8 @@ async def get_graph_nodes_legacy(
db_id: str = Query(..., description="数据库ID"), db_id: str = Query(..., description="数据库ID"),
limit: int = Query(500, description="最大节点数量", ge=1, le=2000), limit: int = Query(500, description="最大节点数量", ge=1, le=2000),
offset: int = Query(0, description="偏移量", ge=0), offset: int = Query(0, description="偏移量", ge=0),
entity_type: Optional[str] = Query(None, description="实体类型筛选"), entity_type: str | None = Query(None, description="实体类型筛选"),
search: Optional[str] = Query(None, description="搜索关键词"), search: str | None = Query(None, description="搜索关键词"),
current_user: User = Depends(get_admin_user) current_user: User = Depends(get_admin_user)
): ):
""" """
@ -173,7 +173,7 @@ async def get_graph_nodes_legacy(
"total": 0 "total": 0
} }
} }
except Exception as e: except Exception as e:
logger.error(f"获取图节点数据失败: {e}") logger.error(f"获取图节点数据失败: {e}")
raise HTTPException(status_code=500, detail=f"获取图节点数据失败: {str(e)}") raise HTTPException(status_code=500, detail=f"获取图节点数据失败: {str(e)}")
@ -184,7 +184,7 @@ async def get_graph_edges_legacy(
db_id: str = Query(..., description="数据库ID"), db_id: str = Query(..., description="数据库ID"),
limit: int = Query(500, description="最大边数量", ge=1, le=2000), limit: int = Query(500, description="最大边数量", ge=1, le=2000),
offset: int = Query(0, description="偏移量", ge=0), offset: int = Query(0, description="偏移量", ge=0),
min_weight: Optional[float] = Query(None, description="最小权重筛选"), min_weight: float | None = Query(None, description="最小权重筛选"),
current_user: User = Depends(get_admin_user) current_user: User = Depends(get_admin_user)
): ):
""" """
@ -202,7 +202,7 @@ async def get_graph_edges_legacy(
"total": 0 "total": 0
} }
} }
except Exception as e: except Exception as e:
logger.error(f"获取图边数据失败: {e}") logger.error(f"获取图边数据失败: {e}")
raise HTTPException(status_code=500, detail=f"获取图边数据失败: {str(e)}") raise HTTPException(status_code=500, detail=f"获取图边数据失败: {str(e)}")
@ -218,30 +218,30 @@ async def get_graph_stats(
""" """
try: try:
logger.info(f"获取图谱统计信息 - db_id: {db_id}") logger.info(f"获取图谱统计信息 - db_id: {db_id}")
# 获取 LightRAG 实例 # 获取 LightRAG 实例
rag_instance = await knowledge_base._get_lightrag_instance(db_id) rag_instance = await knowledge_base._get_lightrag_instance(db_id)
if not rag_instance: if not rag_instance:
raise HTTPException(status_code=404, detail=f"数据库 {db_id} 不存在") raise HTTPException(status_code=404, detail=f"数据库 {db_id} 不存在")
# 通过获取全图来统计节点和边的数量 # 通过获取全图来统计节点和边的数量
knowledge_graph = await rag_instance.get_knowledge_graph( knowledge_graph = await rag_instance.get_knowledge_graph(
node_label="*", node_label="*",
max_depth=1, max_depth=1,
max_nodes=10000 # 设置较大值以获取完整统计 max_nodes=10000 # 设置较大值以获取完整统计
) )
# 统计实体类型分布 # 统计实体类型分布
entity_types = {} entity_types = {}
for node in knowledge_graph.nodes: for node in knowledge_graph.nodes:
entity_type = node.properties.get("entity_type", "unknown") entity_type = node.properties.get("entity_type", "unknown")
entity_types[entity_type] = entity_types.get(entity_type, 0) + 1 entity_types[entity_type] = entity_types.get(entity_type, 0) + 1
entity_types_list = [ entity_types_list = [
{"type": k, "count": v} {"type": k, "count": v}
for k, v in sorted(entity_types.items(), key=lambda x: x[1], reverse=True) for k, v in sorted(entity_types.items(), key=lambda x: x[1], reverse=True)
] ]
return { return {
"success": True, "success": True,
"data": { "data": {
@ -251,8 +251,8 @@ async def get_graph_stats(
"is_truncated": knowledge_graph.is_truncated "is_truncated": knowledge_graph.is_truncated
} }
} }
except Exception as e: except Exception as e:
logger.error(f"获取图谱统计信息失败: {e}") logger.error(f"获取图谱统计信息失败: {e}")
logger.error(f"Traceback: {traceback.format_exc()}") logger.error(f"Traceback: {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"获取图谱统计信息失败: {str(e)}") raise HTTPException(status_code=500, detail=f"获取图谱统计信息失败: {str(e)}")

View File

@ -29,11 +29,7 @@ class ChatbotConfiguration(Configuration):
metadata={ metadata={
"name": "智能体模型", "name": "智能体模型",
"configurable": True, "configurable": True,
"options": [ "options": [],
"zhipu/glm-4-plus",
"siliconflow/Qwen/Qwen2.5-72B-Instruct",
"siliconflow/Qwen/Qwen2.5-7B-Instruct",
],
"description": "智能体的驱动模型" "description": "智能体的驱动模型"
}, },
) )

View File

@ -1,29 +1,73 @@
from datetime import datetime, timezone, UTC from datetime import datetime, timezone, UTC
import asyncio import asyncio
import os
import traceback
from src.models import select_model from src import config
from src.utils import logger, get_docker_safe_url
from src.models import get_custom_model
from src.agents.registry import BaseAgent from src.agents.registry import BaseAgent
from langchain_core.language_models import BaseChatModel from langchain_core.language_models import BaseChatModel
from langchain_core.runnables import RunnableConfig from langchain_core.runnables import RunnableConfig
from langchain_core.messages import AIMessageChunk, ToolMessage from langchain_core.messages import AIMessageChunk, ToolMessage
from pydantic import SecretStr
def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel:
"""Load a chat model from a fully specified name. """
Load a chat model from a fully specified name.
Args:
fully_specified_name (str): String in the format 'provider/model'.
**kwargs: Additional parameters to pass to the model.
""" """
provider, model = fully_specified_name.split("/", maxsplit=1) provider, model = fully_specified_name.split("/", maxsplit=1)
model_instance = select_model(model_name=model, model_provider=provider)
return model_instance.chat_open_ai if provider == "custom":
from langchain_openai import ChatOpenAI
model_info = get_custom_model(model)
api_key = model_info.get("api_key") or "custom_model"
base_url = get_docker_safe_url(model_info["api_base"])
model_name = model_info.get("name") or "custom_model"
return ChatOpenAI(
model=model_name,
api_key=SecretStr(api_key),
base_url=base_url,
)
model_info = config.model_names.get(provider, {})
api_key = os.getenv(model_info["env"][0], model_info["env"][0])
base_url = get_docker_safe_url(model_info["base_url"])
if provider in ["deepseek", "dashscope"]:
from langchain_deepseek import ChatDeepSeek
return ChatDeepSeek(
model=model,
api_key=SecretStr(api_key),
base_url=base_url,
api_base=base_url,
)
elif provider == "together":
from langchain_together import ChatTogether
return ChatTogether(
model=model,
api_key=SecretStr(api_key),
base_url=base_url,
)
else:
try: # 其他模型默认使用OpenAIBase, like openai, zhipuai
from langchain_openai import ChatOpenAI
return ChatOpenAI(
model=model,
api_key=SecretStr(api_key),
base_url=base_url,
)
except Exception as e:
raise ValueError(f"Model provider {provider} load failed, {e} \n {traceback.format_exc()}")
async def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
async def agent_cli(agent: BaseAgent, config: RunnableConfig | None = None):
config = config or {} config = config or {}
if "configurable" not in config: if "configurable" not in config:
config["configurable"] = {} config["configurable"] = {}

View File

@ -59,9 +59,6 @@ class Config(SimpleConfig):
self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(self.embed_model_names.keys())) self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(self.embed_model_names.keys()))
self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(self.reranker_names.keys())) # noqa: E501 self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(self.reranker_names.keys())) # noqa: E501
self.add_item("model_local_paths", default={}, des="本地模型路径")
self.add_item("use_rewrite_query", default="on", des="重写查询", choices=["off", "on", "hyde"])
self.add_item("device", default="cuda", des="运行本地模型的设备", choices=["cpu", "cuda"])
### <<< 默认配置结束 ### <<< 默认配置结束
self.load() self.load()
@ -120,7 +117,6 @@ class Config(SimpleConfig):
""" """
处理配置 处理配置
""" """
model_provider_info = self.model_names.get(self.model_provider, {})
self.model_dir = os.environ.get("MODEL_DIR", "") self.model_dir = os.environ.get("MODEL_DIR", "")
if self.model_dir: if self.model_dir:
@ -130,26 +126,6 @@ class Config(SimpleConfig):
logger.warning(f"Warning: The model directory {self.model_dir} does not exist. If not configured, please ignore it. If configured, please check if the configuration is correct;" logger.warning(f"Warning: The model directory {self.model_dir} does not exist. If not configured, please ignore it. If configured, please check if the configuration is correct;"
"For example, the mapping in the docker-compose file") "For example, the mapping in the docker-compose file")
# 检查模型提供商是否存在
if self.model_provider != "custom":
if self.model_name not in model_provider_info["models"]:
logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name")
self.model_name = model_provider_info["default"]
default_model_name = model_provider_info["default"]
self.model_name = self.get("model_name") or default_model_name
else:
self.model_name = self.get("model_name")
if self.model_name not in [item["custom_id"] for item in self.get("custom_models", [])]:
logger.warning(f"Model name {self.model_name} not in custom models, using default model name")
if self.get("custom_models", []):
self.model_name = self.get("custom_models", [])[0]["custom_id"]
else:
self.model_name = self._config_items["model_name"]["default"]
self.model_provider = self._config_items["model_provider"]["default"]
logger.error(f"No custom models found, using default model {self.model_name} from {self.model_provider}")
# 检查模型提供商的环境变量 # 检查模型提供商的环境变量
conds = {} conds = {}
self.model_provider_status = {} self.model_provider_status = {}
@ -158,12 +134,6 @@ class Config(SimpleConfig):
conds_bool = [bool(os.getenv(_k)) for _k in conds[provider]] conds_bool = [bool(os.getenv(_k)) for _k in conds[provider]]
self.model_provider_status[provider] = all(conds_bool) self.model_provider_status[provider] = all(conds_bool)
# 检查web_search的环境变量
# if self.enable_web_search and not os.getenv("TAVILY_API_KEY"):
# logger.warning("TAVILY_API_KEY not set, web search will be disabled")
# self.enable_web_search = False
# 2025.04.08 修改为不手动配置只要配置了TAVILY_API_KEY就默认开启web_search
if os.getenv("TAVILY_API_KEY"): if os.getenv("TAVILY_API_KEY"):
self.enable_web_search = True self.enable_web_search = True

View File

@ -3,7 +3,6 @@ import json
import warnings import warnings
import traceback import traceback
import torch
from neo4j import GraphDatabase as GD from neo4j import GraphDatabase as GD
from src import config from src import config

View File

@ -5,7 +5,7 @@ import traceback
import shutil import shutil
import asyncio import asyncio
from pathlib import Path from pathlib import Path
from typing import Dict, List, Any, Optional from typing import Optional
from datetime import datetime from datetime import datetime
from lightrag import LightRAG, QueryParam from lightrag import LightRAG, QueryParam
@ -26,11 +26,11 @@ class LightRagBasedKB:
def __init__(self) -> None: def __init__(self) -> None:
# 存储 LightRAG 实例映射 {db_id: LightRAG} # 存储 LightRAG 实例映射 {db_id: LightRAG}
self.instances: Dict[str, LightRAG] = {} self.instances: dict[str, LightRAG] = {}
# 数据库元信息存储 {db_id: metadata} # 数据库元信息存储 {db_id: metadata}
self.databases_meta: Dict[str, dict] = {} self.databases_meta: dict[str, dict] = {}
# 文件信息存储 {file_id: file_info} # 文件信息存储 {file_id: file_info}
self.files_meta: Dict[str, dict] = {} self.files_meta: dict[str, dict] = {}
# 工作目录 # 工作目录
self.work_dir = os.path.join(config.save_dir, "lightrag_data") self.work_dir = os.path.join(config.save_dir, "lightrag_data")
os.makedirs(self.work_dir, exist_ok=True) os.makedirs(self.work_dir, exist_ok=True)
@ -45,7 +45,7 @@ class LightRagBasedKB:
meta_file = os.path.join(self.work_dir, "metadata.json") meta_file = os.path.join(self.work_dir, "metadata.json")
if os.path.exists(meta_file): if os.path.exists(meta_file):
try: try:
with open(meta_file, 'r', encoding='utf-8') as f: with open(meta_file, encoding='utf-8') as f:
data = json.load(f) data = json.load(f)
self.databases_meta = data.get("databases", {}) self.databases_meta = data.get("databases", {})
self.files_meta = data.get("files", {}) self.files_meta = data.get("files", {})
@ -69,7 +69,7 @@ class LightRagBasedKB:
except Exception as e: except Exception as e:
logger.error(f"Failed to save metadata: {e}") logger.error(f"Failed to save metadata: {e}")
async def _get_lightrag_instance(self, db_id: str) -> Optional[LightRAG]: async def _get_lightrag_instance(self, db_id: str) -> LightRAG | None:
"""获取或创建 LightRAG 实例""" """获取或创建 LightRAG 实例"""
logger.info(f"Getting or creating LightRAG instance for {db_id}") logger.info(f"Getting or creating LightRAG instance for {db_id}")
@ -170,13 +170,13 @@ class LightRagBasedKB:
elif file_ext in ['.txt', '.md']: elif file_ext in ['.txt', '.md']:
# 直接读取文本文件 # 直接读取文本文件
with open(file_path_obj, 'r', encoding='utf-8') as f: with open(file_path_obj, encoding='utf-8') as f:
content = f.read() content = f.read()
return f"# {file_path_obj.name}\n\n{content}" return f"# {file_path_obj.name}\n\n{content}"
elif file_ext in ['.doc', '.docx']: elif file_ext in ['.doc', '.docx']:
# 处理 Word 文档 # 处理 Word 文档
from docx import Document # type: ignore from docx import Document # type: ignore
doc = Document(file_path_obj) doc = Document(file_path_obj)
text = '\n'.join([para.text for para in doc.paragraphs]) text = '\n'.join([para.text for para in doc.paragraphs])
@ -416,7 +416,7 @@ class LightRagBasedKB:
try: try:
# 获取文档的所有 chunks # 获取文档的所有 chunks
assert hasattr(rag.text_chunks, 'get_all'), "text_chunks does not have get_all method" assert hasattr(rag.text_chunks, 'get_all'), "text_chunks does not have get_all method"
all_chunks = await rag.text_chunks.get_all() all_chunks = await rag.text_chunks.get_all() # type: ignore
# 筛选属于该文档的 chunks # 筛选属于该文档的 chunks
doc_chunks = [] doc_chunks = []
@ -506,4 +506,4 @@ class LightRagBasedKB:
"description": meta["description"], "description": meta["description"],
"retriever": make_retriever(db_id), "retriever": make_retriever(db_id),
} }
return retrievers return retrievers

View File

@ -5,19 +5,14 @@ from src import config
from src.utils.logging_config import logger from src.utils.logging_config import logger
from src.models.chat_model import OpenAIBase from src.models.chat_model import OpenAIBase
def select_model(model_provider=None, model_name=None): def select_model(model_provider, model_name=None):
"""根据模型提供者选择模型""" """根据模型提供者选择模型"""
model_provider = model_provider or config.model_provider assert model_provider is not None, "Model provider not specified"
model_info = config.model_names.get(model_provider, {}) model_info = config.model_names.get(model_provider, {})
model_name = model_name or config.model_name or model_info.get("default", "") model_name = model_name or model_info.get("default", "")
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`") logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
if model_provider is None:
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
if model_provider == "qianfan": if model_provider == "qianfan":
from src.models.chat_model import Qianfan from src.models.chat_model import Qianfan
return Qianfan(model_name) return Qianfan(model_name)
@ -26,33 +21,6 @@ def select_model(model_provider=None, model_name=None):
from src.models.chat_model import OpenModel from src.models.chat_model import OpenModel
return OpenModel(model_name) return OpenModel(model_name)
if model_provider == "deepseek" or model_provider == "dashscope":
from langchain_deepseek import ChatDeepSeek
return OpenAIBase(
api_key=os.getenv(model_info["env"][0]),
base_url=model_info["base_url"],
model_name=model_name,
chat_open_ai=ChatDeepSeek(
model=model_name,
api_key=os.getenv(model_info["env"][0]),
base_url=model_info["base_url"],
api_base=model_info["base_url"],
)
)
if model_provider == "together":
from langchain_together import ChatTogether
return OpenAIBase(
api_key=os.getenv(model_info["env"][0]),
base_url=model_info["base_url"],
model_name=model_name,
chat_open_ai=ChatTogether(
model=model_name,
api_key=os.getenv(model_info["env"][0]),
base_url=model_info["base_url"],
)
)
if model_provider == "custom": if model_provider == "custom":
model_info = get_custom_model(model_name) model_info = get_custom_model(model_name)
@ -69,8 +37,7 @@ def select_model(model_provider=None, model_name=None):
return model return model
except Exception as e: except Exception as e:
raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}") raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}")
def get_custom_model(model_id): def get_custom_model(model_id):
"""return model_info""" """return model_info"""
assert config.custom_models is not None, "custom_models is not set" assert config.custom_models is not None, "custom_models is not set"

View File

@ -5,15 +5,12 @@ from src.utils import logger, get_docker_safe_url
from langchain_openai import ChatOpenAI from langchain_openai import ChatOpenAI
class OpenAIBase: class OpenAIBase:
def __init__(self, api_key, base_url, model_name, chat_open_ai=None, **kwargs): def __init__(self, api_key, base_url, model_name, **kwargs):
self.api_key = api_key self.api_key = api_key
self.base_url = base_url self.base_url = base_url
self.client = OpenAI(api_key=api_key, base_url=base_url) self.client = OpenAI(api_key=api_key, base_url=base_url)
self.model_name = model_name self.model_name = model_name
self.info = kwargs self.info = kwargs
self.chat_open_ai = chat_open_ai or ChatOpenAI(model=model_name,
api_key=api_key,
base_url=base_url)
def predict(self, message, stream=False): def predict(self, message, stream=False):
if isinstance(message, str): if isinstance(message, str):

View File

@ -13,19 +13,18 @@ from src.utils import hashstr, logger, get_docker_safe_url
class BaseEmbeddingModel: class BaseEmbeddingModel:
embed_state = {} embed_state = {}
def __init__(self, model_id):
self.model_id = model_id
self.info = config.embed_model_names[model_id]
self.model = self.info["name"]
self.dimension = self.info.get("dimension", None)
self.url = get_docker_safe_url(self.info["base_url"])
self.api_key = os.getenv(self.info["api_key"], self.info["api_key"])
@abstractmethod @abstractmethod
def predict(self, message): def predict(self, message):
raise NotImplementedError("Subclasses must implement this method") raise NotImplementedError("Subclasses must implement this method")
def get_dimension(self):
if hasattr(self, "dimension"):
return self.dimension
if hasattr(self, "embed_model_fullname"):
return config.embed_model_names[self.embed_model_fullname].get("dimension", None)
return config.embed_model_names[self.model].get("dimension", None)
def encode(self, message): def encode(self, message):
return self.predict(message) return self.predict(message)
@ -66,77 +65,14 @@ class BaseEmbeddingModel:
return data return data
class LocalEmbeddingModel(BaseEmbeddingModel):
def __init__(self, **kwargs):
info = config.embed_model_names[config.embed_model]
self.model = config.model_local_paths.get(info["name"], info.get("local_path"))
self.model = self.model or info["name"]
self.dimension = info["dimension"]
self.embed_model_fullname = config.embed_model
if os.getenv("MODEL_DIR"):
if os.path.exists(_path := os.path.join(os.getenv("MODEL_DIR"), self.model)):
self.model = _path
else:
logger.warning(f"Local model `{info['name']}` not found in `{self.model}`, using `{info['name']}`")
logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`")
logger.debug("如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
f"如果还是不行,建议手动下载到某个文件夹,比如 {os.getenv('MODEL_DIR', '/models')}/BAAI/bge-m3 目录下;")
self.model = HuggingFaceEmbeddings(
model_name=self.model,
model_kwargs={'device': config.device},
encode_kwargs={
'normalize_embeddings': True,
'prompt_name': info.get("query_instruction", None),
},
)
logger.info(f"Embedding model {info['name']} loaded, {self.model=}")
def predict(self, message):
return self.model.embed_documents(message)
async def aencode(self, message):
return await self.model.aembed_documents(message)
def encode_queries(self, queries):
logger.warning("Huggingface Model 不支持批量 encode queries因此使用训练实现")
data = []
for q in queries:
data.append(self.predict(q))
return data
class ZhipuEmbedding(BaseEmbeddingModel):
def __init__(self) -> None:
self.config = config
self.model = config.embed_model_names[config.embed_model]["name"]
self.dimension = config.embed_model_names[config.embed_model]["dimension"]
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY"))
self.embed_model_fullname = config.embed_model
def predict(self, message):
response = self.client.embeddings.create(
model=self.model,
input=message,
)
data = [a.embedding for a in response.data]
return data
class OllamaEmbedding(BaseEmbeddingModel): class OllamaEmbedding(BaseEmbeddingModel):
def __init__(self) -> None: """
self.info = config.embed_model_names[config.embed_model] Ollama Embedding Model
self.model = self.info["name"] """
self.url = self.info.get("base_url", "http://localhost:11434/api/embed")
self.url = get_docker_safe_url(self.url) def __init__(self, model_id) -> None:
self.dimension = self.info.get("dimension", None) super().__init__(model_id)
self.embed_model_fullname = config.embed_model self.url = self.url or get_docker_safe_url("http://localhost:11434/api/embed")
def predict(self, message: list[str] | str): def predict(self, message: list[str] | str):
if isinstance(message, str): if isinstance(message, str):
@ -154,14 +90,8 @@ class OllamaEmbedding(BaseEmbeddingModel):
class OtherEmbedding(BaseEmbeddingModel): class OtherEmbedding(BaseEmbeddingModel):
def __init__(self) -> None: def __init__(self, model_id) -> None:
self.info = config.embed_model_names[config.embed_model] super().__init__(model_id)
self.embed_model_fullname = config.embed_model
self.dimension = self.info.get("dimension", None)
self.model = self.info["name"]
self.api_key = os.getenv(self.info["api_key"], self.info["api_key"])
self.url = get_docker_safe_url(self.info["base_url"])
assert self.url and self.model, f"URL and model are required. Cur embed model: {config.embed_model}"
self.headers = { self.headers = {
"Authorization": f"Bearer {self.api_key}", "Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json" "Content-Type": "application/json"
@ -181,26 +111,18 @@ class OtherEmbedding(BaseEmbeddingModel):
"input": message, "input": message,
} }
def get_embedding_model(): def get_embedding_model(model_id):
provider, model_name = config.embed_model.split('/', 1) provider, model_name = model_id.split('/', 1) if model_id else ("", "")
support_embed_models = config.embed_model_names.keys() support_embed_models = config.embed_model_names.keys()
assert config.embed_model in support_embed_models, f"Unsupported embed model: {config.embed_model}, only support {support_embed_models}" assert model_id in support_embed_models, f"Unsupported embed model: {model_id}, only support {support_embed_models}"
logger.debug(f"Loading embedding model {config.embed_model}") logger.debug(f"Loading embedding model {model_id}")
if provider == "local": if provider == "local":
logger.warning("[DEPRECATED] Local embedding model will be removed in v0.2, please use other embedding models") raise ValueError("Local embedding model is not supported, please use other embedding models")
model = LocalEmbeddingModel()
elif provider == "zhipu":
model = ZhipuEmbedding()
elif provider == "ollama": elif provider == "ollama":
model = OllamaEmbedding() model = OllamaEmbedding(model_id)
else: else:
model = OtherEmbedding() model = OtherEmbedding(model_id)
return model return model
def handle_local_model(paths, model_name, default_path):
model_path = paths.get(model_name, default_path)
return model_path

View File

@ -2,34 +2,18 @@ import os
import json import json
import requests import requests
import numpy as np import numpy as np
from FlagEmbedding import FlagReranker
from src import config from src import config
from src.utils import logger, get_docker_safe_url from src.utils import logger, get_docker_safe_url
class LocalReranker(FlagReranker):
def __init__(self, **kwargs):
model_info = config.reranker_names[config.reranker]
model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("local_path"))
model_name_or_path = model_name_or_path or model_info["name"]
logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}")
super().__init__(model_name_or_path, use_fp16=True, device=config.device, **kwargs)
logger.info(f"Reranker model {config.reranker} loaded")
def sigmoid(x): def sigmoid(x):
return 1 / (1 + np.exp(-x)) return 1 / (1 + np.exp(-x))
class OnlineRerank: class OnlineReranker:
def __init__(self, **kwargs): def __init__(self, model_name, api_key, base_url, **kwargs):
model_info = config.reranker_names[config.reranker] self.url = get_docker_safe_url(base_url)
self.url = get_docker_safe_url(model_info["base_url"]) self.model = model_name
self.model = model_info["name"] self.api_key = api_key
api_key = os.getenv(model_info["api_key"], model_info["api_key"])
assert api_key, f"{model_info['name']} api_key is required"
self.headers = { self.headers = {
"Authorization": f"Bearer {api_key}", "Authorization": f"Bearer {api_key}",
"Content-Type": "application/json" "Content-Type": "application/json"
@ -59,15 +43,17 @@ class OnlineRerank:
"max_chunks_per_doc": max_length, "max_chunks_per_doc": max_length,
} }
def get_reranker(): def get_reranker(model_id, **kwargs):
support_rerankers = config.reranker_names.keys() support_rerankers = config.reranker_names.keys()
assert config.reranker in support_rerankers, f"Unsupported Reranker: {config.reranker}, only support {support_rerankers}" assert model_id in support_rerankers, f"Unsupported Reranker: {model_id}, only support {support_rerankers}"
provider, model_name = config.reranker.split('/', 1)
if provider == "local":
logger.warning("[DEPRECATED] Local reranker will be removed in v0.2, please use other reranker")
return LocalReranker()
elif provider == "siliconflow":
return OnlineRerank()
else:
raise ValueError(f"Unsupported Reranker: {config.reranker}, only support {config.reranker_names.keys()}")
model_info = config.reranker_names[model_id]
base_url = model_info["base_url"]
api_key = os.getenv(model_info["api_key"], model_info["api_key"])
assert api_key, f"{model_info['name']} api_key is required"
return OnlineReranker(
model_name=model_info["name"],
api_key=api_key,
base_url=base_url,
**kwargs
)

View File

@ -6,12 +6,11 @@
# #
##################################################### #####################################################
MODEL_NAMES: MODEL_NAMES:
openai: openai:
name: OpenAI name: OpenAI
url: https://platform.openai.com/docs/models url: https://platform.openai.com/docs/models
base_url: https://api.openai.com/v1
default: gpt-3.5-turbo default: gpt-3.5-turbo
env: env:
- OPENAI_API_KEY - OPENAI_API_KEY
@ -20,21 +19,23 @@ MODEL_NAMES:
- gpt-4o - gpt-4o
- gpt-4o-mini - gpt-4o-mini
- gpt-3.5-turbo - gpt-3.5-turbo
deepseek: deepseek:
name: DeepSeek name: DeepSeek
url: https://platform.deepseek.com/api-docs/zh-cn/pricing url: https://platform.deepseek.com/api-docs/zh-cn/pricing
default: deepseek-chat
base_url: https://api.deepseek.com/v1 base_url: https://api.deepseek.com/v1
default: deepseek-chat
env: env:
- DEEPSEEK_API_KEY - DEEPSEEK_API_KEY
models: models:
- deepseek-chat - deepseek-chat
- deepseek-reasoner - deepseek-reasoner
zhipu: zhipu:
name: 智谱AI (Zhipu) name: 智谱AI (Zhipu)
url: https://open.bigmodel.cn/dev/api url: https://open.bigmodel.cn/dev/api
default: glm-4-flash
base_url: https://open.bigmodel.cn/api/paas/v4/ base_url: https://open.bigmodel.cn/api/paas/v4/
default: glm-4-flash
env: env:
- ZHIPUAI_API_KEY - ZHIPUAI_API_KEY
models: models:
@ -43,17 +44,19 @@ MODEL_NAMES:
- glm-4-air - glm-4-air
- glm-4-flash - glm-4-flash
- glm-z1-air - glm-z1-air
siliconflow: siliconflow:
name: SiliconFlow name: SiliconFlow
url: https://cloud.siliconflow.cn/models url: https://cloud.siliconflow.cn/models
default: Qwen/Qwen2.5-7B-Instruct
base_url: https://api.siliconflow.cn/v1 base_url: https://api.siliconflow.cn/v1
default: Qwen/Qwen3-8B
env: env:
- SILICONFLOW_API_KEY - SILICONFLOW_API_KEY
models: models:
- Pro/deepseek-ai/DeepSeek-R1 - Pro/deepseek-ai/DeepSeek-R1
- Pro/deepseek-ai/DeepSeek-V3 - Pro/deepseek-ai/DeepSeek-V3
- Qwen/QwQ-32B - Qwen/QwQ-32B
- Qwen/Qwen3-8B
together.ai: together.ai:
name: Together.ai name: Together.ai
@ -71,8 +74,8 @@ MODEL_NAMES:
dashscope: dashscope:
name: 阿里百炼 (DashScope) name: 阿里百炼 (DashScope)
url: https://bailian.console.aliyun.com/?switchAgent=10226727&productCode=p_efm#/model-market url: https://bailian.console.aliyun.com/?switchAgent=10226727&productCode=p_efm#/model-market
default: qwen3-235b-a22b
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1 base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
default: qwen3-235b-a22b
env: env:
- DASHSCOPE_API_KEY - DASHSCOPE_API_KEY
models: models:
@ -83,8 +86,8 @@ MODEL_NAMES:
ark: ark:
name: 豆包Ark name: 豆包Ark
url: https://console.volcengine.com/ark/region:ark+cn-beijing/model url: https://console.volcengine.com/ark/region:ark+cn-beijing/model
default: doubao-1-5-pro-32k-250115
base_url: https://ark.cn-beijing.volces.com/api/v3 base_url: https://ark.cn-beijing.volces.com/api/v3
default: doubao-1-5-pro-32k-250115
env: env:
- ARK_API_KEY - ARK_API_KEY
models: models:

View File

@ -15,9 +15,8 @@
okText="确认删除" okText="确认删除"
cancelText="取消" cancelText="取消"
ok-type="danger" ok-type="danger"
:disabled="configStore.config.model_name == item.name"
> >
<a-button type="text" :disabled="configStore.config.model_name == item.name" @click.stop><DeleteOutlined /></a-button> <a-button type="text" @click.stop><DeleteOutlined /></a-button>
</a-popconfirm> </a-popconfirm>
<a-button type="text" @click.stop="prepareToEditCustomModel(item)"><EditOutlined /></a-button> <a-button type="text" @click.stop="prepareToEditCustomModel(item)"><EditOutlined /></a-button>
</div> </div>

View File

@ -15,7 +15,6 @@
</a-button> </a-button>
</template> </template>
</HeaderComponent> </HeaderComponent>
<!-- <a-alert v-if="configStore.config.embed_model &&database.embed_model != configStore.config.embed_model" message="向量模型不匹配,请重新选择" type="warning" style="margin: 10px 20px;" /> -->
<!-- 添加编辑对话框 --> <!-- 添加编辑对话框 -->
<a-modal v-model:open="editModalVisible" title="编辑知识库信息"> <a-modal v-model:open="editModalVisible" title="编辑知识库信息">
@ -282,12 +281,12 @@
<a-switch v-model:checked="meta.only_need_context" /> <a-switch v-model:checked="meta.only_need_context" />
</div> </div>
<div class="params-item"> <div class="params-item">
<p>只使用上下文</p> <p>只使用提示</p>
<a-switch v-model:checked="meta.only_need_prompt" /> <a-switch v-model:checked="meta.only_need_prompt" />
</div> </div>
<div class="params-item"> <div class="params-item">
<p>筛选 TopK</p> <p>筛选 TopK</p>
<a-input-number size="small" v-model:value="meta.top_k" :min="1" :max="meta.maxQueryCount" /> <a-input-number size="small" v-model:value="meta.top_k" :min="1" :max="100" />
</div> </div>
</div> </div>
<div class="params-group"> <div class="params-group">
@ -349,7 +348,6 @@ import {
CloseCircleFilled, CloseCircleFilled,
ClockCircleFilled, ClockCircleFilled,
DeleteOutlined, DeleteOutlined,
CloudUploadOutlined,
SearchOutlined, SearchOutlined,
LoadingOutlined, LoadingOutlined,
FileOutlined, FileOutlined,
@ -357,7 +355,6 @@ import {
EditOutlined, EditOutlined,
PlusOutlined, PlusOutlined,
SettingOutlined, SettingOutlined,
ApartmentOutlined,
} from '@ant-design/icons-vue' } from '@ant-design/icons-vue'
import HeaderComponent from '@/components/HeaderComponent.vue'; import HeaderComponent from '@/components/HeaderComponent.vue';
import KnowledgeGraphViewer from '@/components/KnowledgeGraphViewer.vue'; import KnowledgeGraphViewer from '@/components/KnowledgeGraphViewer.vue';
@ -376,7 +373,6 @@ const selectedFile = ref(null);
// //
const queryText = ref(''); const queryText = ref('');
const queryResult = ref(null) const queryResult = ref(null)
const filteredResults = ref([])
const configStore = useConfigStore() const configStore = useConfigStore()
const state = reactive({ const state = reactive({
@ -431,49 +427,7 @@ const pagination = ref({
onShowSizeChange: (current, pageSize) => pagination.value.pageSize = pageSize, onShowSizeChange: (current, pageSize) => pagination.value.pageSize = pageSize,
}) })
const filterQueryResults = () => {
if (!queryResult.value || !queryResult.value.all_results) {
return;
}
let results = toRaw(queryResult.value.all_results);
console.log("results", results);
if (meta.filter) {
results = results.filter(r => r.distance >= meta.distanceThreshold);
console.log("before", results);
//
if (configStore.config.enable_reranker) {
//
results = results.filter(r => r.rerank_score >= meta.rerankThreshold);
//
if (meta.sortBy === 'rerank_score') {
results = results.sort((a, b) => b.rerank_score - a.rerank_score);
} else {
// ()
results = results.sort((a, b) => b.distance - a.distance);
}
} else {
//
results = results.sort((a, b) => b.distance - a.distance);
}
console.log("after", results);
results = results.slice(0, meta.topK);
}
filteredResults.value = results;
}
const onQuery = () => { const onQuery = () => {
// if (database.value.embed_model != configStore.config.embed_model) {
// message.error('')
// return
// }
console.log(queryText.value) console.log(queryText.value)
state.searchLoading = true state.searchLoading = true
if (!queryText.value.trim()) { if (!queryText.value.trim()) {
@ -491,7 +445,6 @@ const onQuery = () => {
.then(data => { .then(data => {
console.log(data) console.log(data)
queryResult.value = data queryResult.value = data
filterQueryResults()
}) })
.catch(error => { .catch(error => {
console.error(error) console.error(error)
@ -665,7 +618,7 @@ const getDatabaseInfo = () => {
const deleteFile = (fileId) => { const deleteFile = (fileId) => {
state.lock = true state.lock = true
console.debug("deleteFile", databaseId.value, fileId) console.debug("deleteFile", databaseId.value, fileId)
knowledgeBaseApi.deleteFile(databaseId.value, fileId) return knowledgeBaseApi.deleteFile(databaseId.value, fileId)
.then(data => { .then(data => {
console.log(data) console.log(data)
message.success(data.message || '删除成功') message.success(data.message || '删除成功')
@ -674,6 +627,7 @@ const deleteFile = (fileId) => {
.catch(error => { .catch(error => {
console.error(error) console.error(error)
message.error(error.message || '删除失败') message.error(error.message || '删除失败')
throw error
}) })
.finally(() => { .finally(() => {
state.lock = false state.lock = false
@ -802,27 +756,7 @@ const addFiles = (items, contentType = 'file') => {
}); });
}; };
// "" -
const handleFiles = () => {
console.log(fileList.value);
const files = fileList.value.filter(file => file.status === 'done').map(file => file.response.file_path);
console.log(files);
addFiles(files, 'file');
}
// "" - URL
const handleUrls = () => {
const urls = urlList.value.split('\n')
.map(url => url.trim())
.filter(url => url.length > 0 && (url.startsWith('http://') || url.startsWith('https://')));
if (urls.length === 0) {
message.error('请输入有效的网页链接必须以http://或https://开头)');
return;
}
addFiles(urls, 'url');
};
const columns = [ const columns = [
// { title: 'ID', dataIndex: 'file_id', key: 'file_id' }, // { title: 'ID', dataIndex: 'file_id', key: 'file_id' },
@ -843,10 +777,6 @@ watch(() => route.params.database_id, (newId) => {
} }
); );
// meta
watch(() => meta, () => {
filterQueryResults()
}, { deep: true })
// //
const queryExamples = ref([ const queryExamples = ref([
@ -875,9 +805,20 @@ const urlList = ref('');
const chunkData = () => { const chunkData = () => {
if (uploadMode.value === 'file') { if (uploadMode.value === 'file') {
handleFiles(); const files = fileList.value.filter(file => file.status === 'done').map(file => file.response.file_path);
console.log(files);
addFiles(files, 'file');
} else if (uploadMode.value === 'url') { } else if (uploadMode.value === 'url') {
handleUrls(); const urls = urlList.value.split('\n')
.map(url => url.trim())
.filter(url => url.length > 0 && (url.startsWith('http://') || url.startsWith('https://')));
if (urls.length === 0) {
message.error('请输入有效的网页链接必须以http://或https://开头)');
return;
}
addFiles(urls, 'url');
} }
} }
@ -1494,59 +1435,9 @@ const toggleAutoRefresh = (checked) => {
} }
} }
.chunk-preview {
margin-top: 20px;
.preview-header {
display: flex;
justify-content: space-between;
align-items: center;
margin-bottom: 16px;
h3 {
margin: 0;
color: var(--main-color);
font-size: 18px;
}
}
.result-cards {
display: grid;
grid-template-columns: repeat(auto-fill, minmax(600px, 1fr));
gap: 12px;
margin-top: 10px;
}
.chunk {
background-color: var(--main-light-5);
border: 1px solid var(--main-light-3);
border-radius: 8px;
padding: 16px;
word-wrap: break-word;
word-break: break-all;
transition: all 0.2s ease;
&:hover {
background-color: var(--main-light-4);
border-color: var(--main-light-2);
box-shadow: 0 2px 6px rgba(0, 0, 0, 0.06);
}
p {
margin: 0;
line-height: 1.6;
strong {
color: var(--main-color);
margin-right: 6px;
}
}
}
}
.url-input {
margin-bottom: 20px;
}