From 794c1f4b251bf2536300afb2da1b0a9c3f524410 Mon Sep 17 00:00:00 2001 From: Lynn Date: Fri, 5 Jun 2026 09:45:44 +0800 Subject: [PATCH] Fix: volc engine and other json key factories (#15653) ### What problem does this PR solve? Fix: - VolcEngine adapt to new api_key format - Save dict api_key as json ### Type of change - [x] Bug Fix (non-breaking change which fixes an issue) --- api/apps/services/provider_api_service.py | 17 ++++++++++------- rag/llm/chat_model.py | 10 +++++++--- rag/llm/cv_model.py | 18 +++++++++++++----- rag/llm/embedding_model.py | 8 ++++++-- 4 files changed, 36 insertions(+), 17 deletions(-) diff --git a/api/apps/services/provider_api_service.py b/api/apps/services/provider_api_service.py index 595f3928ee..a8898f0915 100644 --- a/api/apps/services/provider_api_service.py +++ b/api/apps/services/provider_api_service.py @@ -224,7 +224,7 @@ def show_provider_model(provider_name: str, model_name: str): } -async def create_provider_instance(tenant_id: str, provider_name: str, instance_name: str, api_key: str, base_url: str, region: str, model_info: dict=None): +async def create_provider_instance(tenant_id: str, provider_name: str, instance_name: str, api_key: str|dict, base_url: str, region: str, model_info: dict=None): """ Create a provider instance. @@ -263,8 +263,10 @@ async def create_provider_instance(tenant_id: str, provider_name: str, instance_ if not provider_obj: return False, f"Provider '{provider_name}' does not exist" + api_key_str = "" if api_key: - same_key_instance = TenantModelInstanceService.get_by_provider_id_and_api_key(provider_obj.id, api_key) + api_key_str = api_key if isinstance(api_key, str) else json.dumps(api_key) + same_key_instance = TenantModelInstanceService.get_by_provider_id_and_api_key(provider_obj.id, api_key_str) if same_key_instance: return False, f"Already exist instance: {same_key_instance.instance_name} with api_key {api_key}" success, msg = await verify_api_key(provider_name, api_key, base_url, region, model_info) @@ -276,7 +278,7 @@ async def create_provider_instance(tenant_id: str, provider_name: str, instance_ extra_fields["base_url"] = base_url if region: extra_fields["region"] = region - TenantModelInstanceService.create_instance(provider_id=provider_obj.id,instance_name=instance_name,api_key=api_key, extra=json.dumps(extra_fields)) + TenantModelInstanceService.create_instance(provider_id=provider_obj.id,instance_name=instance_name,api_key=api_key_str, extra=json.dumps(extra_fields)) if model_info: success, msg = add_model_to_instance(tenant_id, provider_name, instance_name, **model_info) if not success: @@ -319,7 +321,7 @@ def list_provider_instances(tenant_id: str, provider_name: str): return True, active_instances + inactive_instances -async def verify_api_key(provider_name: str, api_key: str, base_url: str=None, region: str=None, model_info: dict=None): +async def verify_api_key(provider_name: str, api_key: str|dict, base_url: str=None, region: str=None, model_info: dict=None): """ Verify API key for a provider. @@ -364,10 +366,11 @@ async def verify_api_key(provider_name: str, api_key: str, base_url: str=None, r timeout_seconds = int(os.environ.get("LLM_TIMEOUT_SECONDS", 10)) extra = {"provider": provider_name} msg = "" + api_key_str = api_key if isinstance(api_key, str) else json.dumps(api_key) for llm in factory_llms: if not embd_passed and llm["model_type"] == LLMType.EMBEDDING.value: assert provider_name in EmbeddingModel, f"Embedding model from {provider_name} is not supported yet." - mdl = EmbeddingModel[provider_name](api_key, llm["llm_name"], base_url=base_url) + mdl = EmbeddingModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url) try: arr, tc = asyncio.wait_for( asyncio.to_thread(mdl.encode, ["Test if the api key is available"]), @@ -380,7 +383,7 @@ async def verify_api_key(provider_name: str, api_key: str, base_url: str=None, r msg += f"\nFail to access embedding model({llm['llm_name']}) using this api key." + str(e) elif not chat_passed and llm["model_type"] == LLMType.CHAT.value: assert provider_name in ChatModel, f"Chat model from {provider_name} is not supported yet." - mdl = ChatModel[provider_name](api_key, llm["llm_name"], base_url=base_url, **extra) + mdl = ChatModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url, **extra) try: async def check_streamly(): async for chunk in mdl.async_chat_streamly( @@ -401,7 +404,7 @@ async def verify_api_key(provider_name: str, api_key: str, base_url: str=None, r msg += f"\nFail to access model({provider_name}/{llm['llm_name']}) using this api key." + str(e) elif not rerank_passed and llm["model_type"] == LLMType.RERANK.value: assert provider_name in RerankModel, f"Rerank model from {provider_name} is not supported yet." - mdl = RerankModel[provider_name](api_key, llm["llm_name"], base_url=base_url) + mdl = RerankModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url) try: arr, tc = await asyncio.wait_for( asyncio.to_thread(mdl.similarity, "What's the weather?", ["Is it sunny today?"]), diff --git a/rag/llm/chat_model.py b/rag/llm/chat_model.py index 9e20514d30..1200fe3c5e 100644 --- a/rag/llm/chat_model.py +++ b/rag/llm/chat_model.py @@ -25,6 +25,7 @@ from copy import deepcopy from urllib.parse import urljoin import json_repair +from json.decoder import JSONDecodeError import litellm import openai from openai import AsyncOpenAI, OpenAI @@ -831,9 +832,12 @@ class VolcEngineChat(Base): model_name is for display only """ base_url = base_url if base_url else "https://ark.cn-beijing.volces.com/api/v3" - ark_api_key = json.loads(key).get("ark_api_key", "") - model_name = json.loads(key).get("ep_id", "") + json.loads(key).get("endpoint_id", "") - super().__init__(ark_api_key, model_name, base_url, **kwargs) + try: + ark_api_key = json.loads(key).get("ark_api_key", "") + model_name = json.loads(key).get("ep_id", "") + json.loads(key).get("endpoint_id", "") + super().__init__(ark_api_key, model_name, base_url, **kwargs) + except JSONDecodeError: + super().__init__(key, model_name, base_url, **kwargs) class MistralChat(Base): diff --git a/rag/llm/cv_model.py b/rag/llm/cv_model.py index d247f2c44d..bcb93d5b42 100644 --- a/rag/llm/cv_model.py +++ b/rag/llm/cv_model.py @@ -25,6 +25,7 @@ from copy import deepcopy from io import BytesIO from pathlib import Path from urllib.parse import urljoin +from json.decoder import JSONDecodeError import requests from openai import OpenAI, AsyncOpenAI @@ -558,14 +559,21 @@ class VolcEngineCV(GptV4): def __init__(self, key, model_name, lang="Chinese", base_url="https://ark.cn-beijing.volces.com/api/v3", **kwargs): if not base_url: base_url = "https://ark.cn-beijing.volces.com/api/v3" - ark_api_key = json.loads(key).get("ark_api_key", "") - self.client = OpenAI(api_key=ark_api_key, base_url=base_url) - self.async_client = AsyncOpenAI(api_key=ark_api_key, base_url=base_url) - self.model_name = json.loads(key).get("ep_id", "") + json.loads(key).get("endpoint_id", "") + + try: + api_key = json.loads(key).get("ark_api_key", "") + llm_name = json.loads(key).get("ep_id", "") + json.loads(key).get("endpoint_id", "") + + except JSONDecodeError: + api_key = key + llm_name = model_name + + self.client = OpenAI(api_key=api_key, base_url=base_url) + self.async_client = AsyncOpenAI(api_key=api_key, base_url=base_url) + self.model_name = llm_name self.lang = lang Base.__init__(self, **kwargs) - class LmStudioCV(GptV4): _FACTORY_NAME = "LM-Studio" diff --git a/rag/llm/embedding_model.py b/rag/llm/embedding_model.py index 79fa69eef0..516f3dad5a 100644 --- a/rag/llm/embedding_model.py +++ b/rag/llm/embedding_model.py @@ -19,6 +19,7 @@ import threading from abc import ABC from contextlib import contextmanager from urllib.parse import urljoin +from json.decoder import JSONDecodeError import dashscope import numpy as np @@ -1084,8 +1085,11 @@ class VolcEngineEmbed(Base): base_url = "https://ark.cn-beijing.volces.com/api/v3" self.base_url = base_url - cfg = json.loads(key) - self.ark_api_key = cfg.get("ark_api_key", "") + try: + cfg = json.loads(key) + self.ark_api_key = cfg.get("ark_api_key", "") + except JSONDecodeError: + self.ark_api_key = key self.model_name = model_name @staticmethod