From b7ef7fa0c2421d866d17ccb5570fdf4fbcb181fa Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 14 Mar 2025 03:28:53 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E8=8E=B7=E5=8F=96=E5=AE=89?= =?UTF-8?q?=E5=85=A8API=E7=9A=84=E6=96=B9=E6=B3=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/config/__init__.py | 32 ++++++++++++++++++++++++++++++++ src/models/chat_model.py | 2 +- src/routers/base_router.py | 15 +++++++-------- 3 files changed, 40 insertions(+), 9 deletions(-) diff --git a/src/config/__init__.py b/src/config/__init__.py index d0b516c5..0682393d 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -11,6 +11,7 @@ MODEL_NAMES = _models["MODEL_NAMES"] EMBED_MODEL_INFO = _models["EMBED_MODEL_INFO"] RERANKER_LIST = _models["RERANKER_LIST"] +DEFAULT_MOCK_API = 'this_is_mock_api_key_in_frontend' class SimpleConfig(dict): @@ -174,3 +175,34 @@ class Config(SimpleConfig): json.dump(self, f, indent=4) logger.info(f"Config file {self.filename} saved") + + def get_safe_config(self): + """ + 获取安全的配置,即过滤掉 api_key + """ + + config = json.loads(str(self)) + + # 过滤掉 api_key + for model in config.get("custom_models", []): + model["api_key"] = DEFAULT_MOCK_API if model.get("api_key") else "" + + return config + + def compare_custom_models(self, value): + """ + 比较 custom_models 中的 api_key,如果输入的 api_key 与当前的 api_key 相同,则不修改 + 如果输入的 api_key 为 DEFAULT_MOCK_API,则使用当前的 api_key + """ + current_models_dict = {model["custom_id"]: model.get("api_key") for model in self.custom_models} + + for i, model in enumerate(value): + input_custom_id = model.get("custom_id") + input_api_key = model.get("api_key") + + if input_custom_id in current_models_dict: + current_api_key = current_models_dict[input_custom_id] + if input_api_key == DEFAULT_MOCK_API or input_api_key == current_api_key: + value[i]["api_key"] = current_api_key + + return value diff --git a/src/models/chat_model.py b/src/models/chat_model.py index d62da34c..fe5c9c82 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -58,7 +58,7 @@ class OpenModel(OpenAIBase): class CustomModel(OpenAIBase): def __init__(self, model_info): model_name = model_info["name"] - api_key = model_info["api_key"] or "custom_model" + api_key = model_info.get("api_key", "custom_model") base_url = get_docker_safe_url(model_info["api_base"]) logger.info(f"> Custom model: {model_name}, base_url: {base_url}") diff --git a/src/routers/base_router.py b/src/routers/base_router.py index 0dcbbde7..16986210 100644 --- a/src/routers/base_router.py +++ b/src/routers/base_router.py @@ -1,14 +1,10 @@ +from fastapi import Request, Body from fastapi import APIRouter base = APIRouter() -from fastapi import FastAPI, HTTPException -from fastapi.responses import JSONResponse -from fastapi import Request, Body - -from src.core import HistoryManager from src.core.startup import startup -from src.utils import logger +from src.utils import hashstr, logger @base.get("/") @@ -17,13 +13,16 @@ async def route_index(): @base.get("/config") def get_config(): - return startup.config + return startup.config.get_safe_config() @base.post("/config") async def update_config(key = Body(...), value = Body(...)): + if key == "custom_models": + value = startup.config.compare_custom_models(value) + startup.config[key] = value startup.config.save() - return startup.config + return startup.config.get_safe_config() @base.post("/restart") async def restart():