mirror of
https://github.com/snailyp/gemini-balance.git
synced 2026-09-07 00:36:36 +08:00
refactor(config): 将服务配置改为从 settings 获取
将 SecurityService, ModelService, EmbeddingService 的配置依赖从构造函数注入改为直接从 app.config.config.settings 获取。 这简化了服务类的实例化过程,并实现了配置的集中管理。
This commit is contained in:
+7
-10
@@ -13,12 +13,9 @@ def verify_auth_token(token: str) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
class SecurityService:
|
class SecurityService:
|
||||||
def __init__(self, allowed_tokens: list, auth_token: str):
|
|
||||||
self.allowed_tokens = allowed_tokens
|
|
||||||
self.auth_token = auth_token
|
|
||||||
|
|
||||||
async def verify_key(self, key: str):
|
async def verify_key(self, key: str):
|
||||||
if key not in self.allowed_tokens and key != self.auth_token:
|
if key not in settings.ALLOWED_TOKENS and key != settings.AUTH_TOKEN:
|
||||||
logger.error("Invalid key")
|
logger.error("Invalid key")
|
||||||
raise HTTPException(status_code=401, detail="Invalid key")
|
raise HTTPException(status_code=401, detail="Invalid key")
|
||||||
return key
|
return key
|
||||||
@@ -37,7 +34,7 @@ class SecurityService:
|
|||||||
)
|
)
|
||||||
|
|
||||||
token = authorization.replace("Bearer ", "")
|
token = authorization.replace("Bearer ", "")
|
||||||
if token not in self.allowed_tokens and token != self.auth_token:
|
if token not in settings.ALLOWED_TOKENS and token != settings.AUTH_TOKEN:
|
||||||
logger.error("Invalid token")
|
logger.error("Invalid token")
|
||||||
raise HTTPException(status_code=401, detail="Invalid token")
|
raise HTTPException(status_code=401, detail="Invalid token")
|
||||||
|
|
||||||
@@ -52,8 +49,8 @@ class SecurityService:
|
|||||||
raise HTTPException(status_code=401, detail="Missing x-goog-api-key header")
|
raise HTTPException(status_code=401, detail="Missing x-goog-api-key header")
|
||||||
|
|
||||||
if (
|
if (
|
||||||
x_goog_api_key not in self.allowed_tokens
|
x_goog_api_key not in settings.ALLOWED_TOKENS
|
||||||
and x_goog_api_key != self.auth_token
|
and x_goog_api_key != settings.AUTH_TOKEN
|
||||||
):
|
):
|
||||||
logger.error("Invalid x-goog-api-key")
|
logger.error("Invalid x-goog-api-key")
|
||||||
raise HTTPException(status_code=401, detail="Invalid x-goog-api-key")
|
raise HTTPException(status_code=401, detail="Invalid x-goog-api-key")
|
||||||
@@ -67,7 +64,7 @@ class SecurityService:
|
|||||||
logger.error("Missing auth_token header")
|
logger.error("Missing auth_token header")
|
||||||
raise HTTPException(status_code=401, detail="Missing auth_token header")
|
raise HTTPException(status_code=401, detail="Missing auth_token header")
|
||||||
token = authorization.replace("Bearer ", "")
|
token = authorization.replace("Bearer ", "")
|
||||||
if token != self.auth_token:
|
if token != settings.AUTH_TOKEN:
|
||||||
logger.error("Invalid auth_token")
|
logger.error("Invalid auth_token")
|
||||||
raise HTTPException(status_code=401, detail="Invalid auth_token")
|
raise HTTPException(status_code=401, detail="Invalid auth_token")
|
||||||
|
|
||||||
@@ -78,7 +75,7 @@ class SecurityService:
|
|||||||
) -> str:
|
) -> str:
|
||||||
"""验证URL中的key或请求头中的x-goog-api-key"""
|
"""验证URL中的key或请求头中的x-goog-api-key"""
|
||||||
# 如果URL中的key有效,直接返回
|
# 如果URL中的key有效,直接返回
|
||||||
if key in self.allowed_tokens or key == self.auth_token:
|
if key in settings.ALLOWED_TOKENS or key == settings.AUTH_TOKEN:
|
||||||
return key
|
return key
|
||||||
|
|
||||||
# 否则检查请求头中的x-goog-api-key
|
# 否则检查请求头中的x-goog-api-key
|
||||||
@@ -86,7 +83,7 @@ class SecurityService:
|
|||||||
logger.error("Invalid key and missing x-goog-api-key header")
|
logger.error("Invalid key and missing x-goog-api-key header")
|
||||||
raise HTTPException(status_code=401, detail="Invalid key and missing x-goog-api-key header")
|
raise HTTPException(status_code=401, detail="Invalid key and missing x-goog-api-key header")
|
||||||
|
|
||||||
if x_goog_api_key not in self.allowed_tokens and x_goog_api_key != self.auth_token:
|
if x_goog_api_key not in settings.ALLOWED_TOKENS and x_goog_api_key != settings.AUTH_TOKEN:
|
||||||
logger.error("Invalid key and invalid x-goog-api-key")
|
logger.error("Invalid key and invalid x-goog-api-key")
|
||||||
raise HTTPException(status_code=401, detail="Invalid key and invalid x-goog-api-key")
|
raise HTTPException(status_code=401, detail="Invalid key and invalid x-goog-api-key")
|
||||||
|
|
||||||
|
|||||||
@@ -17,8 +17,8 @@ router_v1beta = APIRouter(prefix=f"/{API_VERSION}")
|
|||||||
logger = get_gemini_logger()
|
logger = get_gemini_logger()
|
||||||
|
|
||||||
# 初始化服务
|
# 初始化服务
|
||||||
security_service = SecurityService(settings.ALLOWED_TOKENS, settings.AUTH_TOKEN)
|
security_service = SecurityService()
|
||||||
model_service = ModelService(settings.SEARCH_MODELS, settings.IMAGE_MODELS)
|
model_service = ModelService()
|
||||||
|
|
||||||
|
|
||||||
async def get_key_manager():
|
async def get_key_manager():
|
||||||
|
|||||||
@@ -20,9 +20,9 @@ router = APIRouter()
|
|||||||
logger = get_openai_logger()
|
logger = get_openai_logger()
|
||||||
|
|
||||||
# 初始化服务
|
# 初始化服务
|
||||||
security_service = SecurityService(settings.ALLOWED_TOKENS, settings.AUTH_TOKEN)
|
security_service = SecurityService()
|
||||||
model_service = ModelService(settings.SEARCH_MODELS, settings.IMAGE_MODELS)
|
model_service = ModelService()
|
||||||
embedding_service = EmbeddingService(settings.BASE_URL)
|
embedding_service = EmbeddingService()
|
||||||
image_create_service = ImageCreateService()
|
image_create_service = ImageCreateService()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,22 +2,20 @@ from typing import List, Union
|
|||||||
|
|
||||||
import openai
|
import openai
|
||||||
from openai.types import CreateEmbeddingResponse
|
from openai.types import CreateEmbeddingResponse
|
||||||
|
from app.config.config import settings
|
||||||
from app.log.logger import get_embeddings_logger
|
from app.log.logger import get_embeddings_logger
|
||||||
|
|
||||||
logger = get_embeddings_logger()
|
logger = get_embeddings_logger()
|
||||||
|
|
||||||
|
|
||||||
class EmbeddingService:
|
class EmbeddingService:
|
||||||
def __init__(self, base_url: str):
|
|
||||||
self.base_url = base_url
|
|
||||||
|
|
||||||
async def create_embedding(
|
async def create_embedding(
|
||||||
self, input_text: Union[str, List[str]], model: str, api_key: str
|
self, input_text: Union[str, List[str]], model: str, api_key: str
|
||||||
) -> CreateEmbeddingResponse:
|
) -> CreateEmbeddingResponse:
|
||||||
"""Create embeddings using OpenAI API"""
|
"""Create embeddings using OpenAI API"""
|
||||||
try:
|
try:
|
||||||
client = openai.OpenAI(api_key=api_key, base_url=self.base_url)
|
client = openai.OpenAI(api_key=api_key, base_url=settings.BASE_URL)
|
||||||
response = client.embeddings.create(input=input_text, model=model)
|
response = client.embeddings.create(input=input_text, model=model)
|
||||||
return response
|
return response
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -10,14 +10,8 @@ logger = get_model_logger()
|
|||||||
|
|
||||||
|
|
||||||
class ModelService:
|
class ModelService:
|
||||||
def __init__(self, search_models: list, image_models: list):
|
|
||||||
self.search_models = search_models
|
|
||||||
self.image_models = image_models
|
|
||||||
self.base_url = settings.BASE_URL
|
|
||||||
self.filtered_models = settings.FILTERED_MODELS
|
|
||||||
|
|
||||||
def get_gemini_models(self, api_key: str) -> Optional[Dict[str, Any]]:
|
def get_gemini_models(self, api_key: str) -> Optional[Dict[str, Any]]:
|
||||||
url = f"{self.base_url}/models?key={api_key}"
|
url = f"{settings.BASE_URL}/models?key={api_key}"
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response = requests.get(url)
|
response = requests.get(url)
|
||||||
@@ -27,7 +21,7 @@ class ModelService:
|
|||||||
filtered_models_list = []
|
filtered_models_list = []
|
||||||
for model in gemini_models.get("models", []):
|
for model in gemini_models.get("models", []):
|
||||||
model_id = model["name"].split("/")[-1]
|
model_id = model["name"].split("/")[-1]
|
||||||
if model_id not in self.filtered_models:
|
if model_id not in settings.FILTERED_MODELS:
|
||||||
filtered_models_list.append(model)
|
filtered_models_list.append(model)
|
||||||
else:
|
else:
|
||||||
logger.info(f"Filtered out model: {model_id}")
|
logger.info(f"Filtered out model: {model_id}")
|
||||||
@@ -68,11 +62,11 @@ class ModelService:
|
|||||||
}
|
}
|
||||||
openai_format["data"].append(openai_model)
|
openai_format["data"].append(openai_model)
|
||||||
|
|
||||||
if model_id in self.search_models:
|
if model_id in settings.SEARCH_MODELS:
|
||||||
search_model = openai_model.copy()
|
search_model = openai_model.copy()
|
||||||
search_model["id"] = f"{model_id}-search"
|
search_model["id"] = f"{model_id}-search"
|
||||||
openai_format["data"].append(search_model)
|
openai_format["data"].append(search_model)
|
||||||
if model_id in self.image_models:
|
if model_id in settings.IMAGE_MODELS:
|
||||||
image_model = openai_model.copy()
|
image_model = openai_model.copy()
|
||||||
image_model["id"] = f"{model_id}-image"
|
image_model["id"] = f"{model_id}-image"
|
||||||
openai_format["data"].append(image_model)
|
openai_format["data"].append(image_model)
|
||||||
@@ -90,9 +84,9 @@ class ModelService:
|
|||||||
model = model.strip()
|
model = model.strip()
|
||||||
if model.endswith("-search"):
|
if model.endswith("-search"):
|
||||||
model = model[:-7]
|
model = model[:-7]
|
||||||
return model in self.search_models
|
return model in settings.SEARCH_MODELS
|
||||||
if model.endswith("-image"):
|
if model.endswith("-image"):
|
||||||
model = model[:-6]
|
model = model[:-6]
|
||||||
return model in self.image_models
|
return model in settings.IMAGE_MODELS
|
||||||
|
|
||||||
return model not in self.filtered_models
|
return model not in settings.FILTERED_MODELS
|
||||||
|
|||||||
Reference in New Issue
Block a user