diff --git a/server/routers/data_router.py b/server/routers/data_router.py index c51b3397..2c397eb9 100644 --- a/server/routers/data_router.py +++ b/server/routers/data_router.py @@ -29,7 +29,7 @@ async def create_database( ): logger.debug(f"Create database {database_name}") try: - embed_info = config.embed_model_info[embed_model_name] + embed_info = config.embed_model_names[embed_model_name] database_info = knowledge_base.create_database( database_name, description, diff --git a/src/core/lightrag_based_kb.py b/src/core/lightrag_based_kb.py index e4046160..ec274351 100644 --- a/src/core/lightrag_based_kb.py +++ b/src/core/lightrag_based_kb.py @@ -76,6 +76,11 @@ class LightRagBasedKB: if db_id not in self.databases_meta: return None + llm_info = self.databases_meta[db_id].get("llm_info", {}) + embed_info = self.databases_meta[db_id].get("embed_info", {}) + logger.info(f"LLM info: {llm_info}") + logger.info(f"Embed info: {embed_info}") + # 创建 LightRAG 实例 working_dir = os.path.join(self.work_dir, db_id) os.makedirs(working_dir, exist_ok=True) @@ -84,8 +89,8 @@ class LightRagBasedKB: # 使用配置的 LLM 和 embedding 函数 rag = LightRAG( working_dir=working_dir, - llm_model_func=self._get_llm_func(), - embedding_func=self._get_embedding_func(), + llm_model_func=self._get_llm_func(llm_info), + embedding_func=self._get_embedding_func(embed_info), vector_storage="MilvusVectorDBStorage", kv_storage="JsonKVStorage", graph_storage="PGGraphStorage", @@ -110,30 +115,40 @@ class LightRagBasedKB: await rag.initialize_storages() await initialize_pipeline_status() - def _get_llm_func(self): + def _get_llm_func(self, llm_info: dict): """获取 LLM 函数""" + llm_info = llm_info | { + "model_name": "qwen3-1.7b", + "provider": "dashscope" + } + provider_info = config.model_names[llm_info.get("provider")] + api_key = os.getenv(provider_info.get("env")[0] or "OPENAI_API_KEY") or "no_api_key" + base_url = get_docker_safe_url(provider_info.get("base_url", "http://localhost:8081/v1")) async def llm_model_func(prompt, system_prompt=None, history_messages=[], **kwargs): return await openai_complete_if_cache( - "qwen3:32b", + llm_info.get("model_name"), prompt, system_prompt=system_prompt, history_messages=history_messages, - api_key="no_api_key", - base_url="http://172.19.13.7:8080/v1", + api_key=api_key, + base_url=base_url, + extra_body={"enable_thinking": False}, **kwargs, ) return llm_model_func - def _get_embedding_func(self): + def _get_embedding_func(self, embed_info: dict): """获取 embedding 函数""" + api_key = os.getenv(embed_info.get("api_key", "OPENAI_API_KEY")) or "no_api_key" + base_url = embed_info.get("base_url", "http://localhost:8081/v1").replace("/embeddings", "") return EmbeddingFunc( - embedding_dim=1024, + embedding_dim=embed_info.get("dimension") or 1024, max_token_size=4096, func=lambda texts: openai_embed( texts=texts, - model="Qwen3-Embedding-0.6B", - api_key="no_api_key", - base_url=get_docker_safe_url("http://localhost:8081/v1") + model=embed_info.get("model_name") or "Qwen3-Embedding-0.6B", + api_key=api_key, + base_url=get_docker_safe_url(base_url) ), ) @@ -248,6 +263,7 @@ class LightRagBasedKB: def delete_database(self, db_id): """删除数据库 - data_router.py 使用""" + # TODO 删除数据库时,需要删除文件记录,并删除 LightRAG 中的文件 if db_id in self.databases_meta: # 删除相关文件记录 files_to_delete = [fid for fid, finfo in self.files_meta.items() @@ -377,6 +393,7 @@ class LightRagBasedKB: async def delete_file(self, db_id, file_id): """删除文件 - data_router.py 使用""" + # TODO 删除文件时,需要删除文件记录,并删除 LightRAG 中的文件 rag = await self._get_lightrag_instance(db_id) if rag: try: diff --git a/src/models/embedding.py b/src/models/embedding.py index e40d9fff..f4823cc7 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -133,7 +133,7 @@ class OllamaEmbedding(BaseEmbeddingModel): def __init__(self) -> None: self.info = config.embed_model_names[config.embed_model] self.model = self.info["name"] - self.url = self.info.get("url", "http://localhost:11434/api/embed") + self.url = self.info.get("base_url", "http://localhost:11434/api/embed") self.url = get_docker_safe_url(self.url) self.dimension = self.info.get("dimension", None) self.embed_model_fullname = config.embed_model @@ -160,7 +160,7 @@ class OtherEmbedding(BaseEmbeddingModel): 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["url"]) + 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 = { "Authorization": f"Bearer {self.api_key}", diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py index 85109d91..01c4431d 100644 --- a/src/models/rerank_model.py +++ b/src/models/rerank_model.py @@ -25,7 +25,7 @@ def sigmoid(x): class OnlineRerank: def __init__(self, **kwargs): model_info = config.reranker_names[config.reranker] - self.url = get_docker_safe_url(model_info["url"]) + self.url = get_docker_safe_url(model_info["base_url"]) self.model = model_info["name"] api_key = os.getenv(model_info["api_key"], model_info["api_key"]) diff --git a/src/static/models.yaml b/src/static/models.yaml index cb8192c7..a1d9f4eb 100644 --- a/src/static/models.yaml +++ b/src/static/models.yaml @@ -111,44 +111,36 @@ MODEL_NAMES: EMBED_MODEL_INFO: - zhipu/zhipu-embedding-2: - name: embedding-2 - dimension: 1024 - - zhipu/zhipu-embedding-3: - name: embedding-3 - dimension: 2048 - siliconflow/BAAI/bge-m3: name: BAAI/bge-m3 dimension: 1024 - url: https://api.siliconflow.cn/v1/embeddings + base_url: https://api.siliconflow.cn/v1/embeddings api_key: SILICONFLOW_API_KEY vllm/Qwen/Qwen3-Embedding-0.6B: name: Qwen3-Embedding-0.6B dimension: 1024 - url: http://localhost:8081/v1/embeddings + base_url: http://localhost:8081/v1/embeddings api_key: no_api_key ollama/nomic-embed-text: name: nomic-embed-text - url: http://localhost:11434/api/embed + base_url: http://localhost:11434/api/embed dimension: 768 ollama/bge-m3: name: bge-m3 - url: http://localhost:11434/api/embed + base_url: http://localhost:11434/api/embed dimension: 1024 RERANKER_LIST: siliconflow/BAAI/bge-reranker-v2-m3: name: BAAI/bge-reranker-v2-m3 - url: https://api.siliconflow.cn/v1/rerank + base_url: https://api.siliconflow.cn/v1/rerank api_key: SILICONFLOW_API_KEY vllm/Qwen/Qwen3-Reranker-0.6B: name: Qwen/Qwen3-Reranker-0.6B - url: http://localhost:8081/v1/rerank + base_url: http://localhost:8081/v1/rerank api_key: no_api_key diff --git a/web/src/views/DataBaseView.vue b/web/src/views/DataBaseView.vue index 1a5b6326..a5dad062 100644 --- a/web/src/views/DataBaseView.vue +++ b/web/src/views/DataBaseView.vue @@ -8,11 +8,11 @@ - +

知识库名称*

嵌入模型

- +

知识库描述

在智能体流程中,这里的描述会作为工具的描述。智能体会根据知识库的标题和描述来选择合适的工具。所以这里描述的越详细,智能体越容易选择到合适的工具。

{ - return Object.keys(configStore.config?.embed_model_info || {}).map(key => ({ - label: configStore.config?.embed_model_info[key]?.name || key, + return Object.keys(configStore.config?.embed_model_names || {}).map(key => ({ + label: `${key} (${configStore.config?.embed_model_names[key]?.dimension})`, value: key, })) }) @@ -94,7 +94,7 @@ const embedModelOptions = computed(() => { const emptyEmbedInfo = { name: '', description: '', - embed_model_name: '', + embed_model_name: configStore.config?.embed_model, } const newDatabase = reactive({ @@ -270,4 +270,10 @@ onMounted(() => { .database-container { padding: 0; } + +.new-database-modal { + h3 { + margin-top: 10px; + } +}