refactor(error): 统一异常处理和响应格式

这次提交重构了整个应用的异常处理机制,保证了处理方式的一致性,还能提供更详细的错误信息。

主要改动包括:
- 修改了 `ApiClient`,现在抛出的异常会同时包含状态码和消息。这样上游服务就能传递准确的 HTTP 错误响应啦。
- 更新了所有服务层(`gemini`、`openai`、`vertex`、`embedding`),现在会捕获这些结构化的异常,不再从字符串里解析错误消息了。
- 增强了路由级别的错误处理,特别是针对流式端点,能正确捕获初始化错误,并返回结构化的 JSON 错误响应,而不是格式错误的 SSE 事件。
- 在所有 API 路由中添加了 `allowed_token` 的日志记录,方便追踪和调试授权问题。
- 还有一些常规的代码清理,比如调整了 import 顺序和格式化代码,提高了可读性和可维护性。
This commit is contained in:
snaily
2025-09-18 03:11:45 +08:00
parent e104a50cf4
commit 67dd1af583
12 changed files with 557 additions and 416 deletions
+3 -3
View File
@@ -130,11 +130,11 @@ def setup_exception_handlers(app: FastAPI) -> None:
"""处理通用异常"""
logger.exception(f"Unhandled Exception: {str(exc)}")
return JSONResponse(
status_code=500,
status_code=exc.args[0],
content={
"error": {
"code": "internal_server_error",
"message": "An unexpected error occurred",
"code": exc.args[0],
"message": exc.args[1],
}
},
)
+246 -155
View File
@@ -1,20 +1,28 @@
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse, JSONResponse
from copy import deepcopy
import json
import asyncio
from copy import deepcopy
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import JSONResponse, StreamingResponse
from app.config.config import settings
from app.log.logger import get_gemini_logger
from app.core.constants import API_VERSION
from app.core.security import SecurityService
from app.domain.gemini_models import GeminiContent, GeminiRequest, ResetSelectedKeysRequest, VerifySelectedKeysRequest, GeminiEmbedRequest, GeminiBatchEmbedRequest
from app.domain.gemini_models import (
GeminiBatchEmbedRequest,
GeminiContent,
GeminiEmbedRequest,
GeminiRequest,
ResetSelectedKeysRequest,
VerifySelectedKeysRequest,
)
from app.handler.error_handler import handle_route_errors
from app.handler.retry_handler import RetryHandler
from app.log.logger import get_gemini_logger
from app.service.chat.gemini_chat_service import GeminiChatService
from app.service.embedding.gemini_embedding_service import GeminiEmbeddingService
from app.service.key.key_manager import KeyManager, get_key_manager_instance
from app.service.tts.native.tts_routes import get_tts_chat_service
from app.service.model.model_service import ModelService
from app.handler.retry_handler import RetryHandler
from app.handler.error_handler import handle_route_errors
from app.core.constants import API_VERSION
from app.service.tts.native.tts_routes import get_tts_chat_service
from app.utils.helpers import redact_key_for_logging
router = APIRouter(prefix=f"/gemini/{API_VERSION}")
@@ -48,8 +56,8 @@ async def get_embedding_service(key_manager: KeyManager = Depends(get_key_manage
@router.get("/models")
@router_v1beta.get("/models")
async def list_models(
_=Depends(security_service.verify_key_or_goog_api_key),
key_manager: KeyManager = Depends(get_key_manager)
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
key_manager: KeyManager = Depends(get_key_manager),
):
"""获取可用的 Gemini 模型列表,并根据配置添加衍生模型(搜索、图像、非思考)。"""
operation_name = "list_gemini_models"
@@ -59,20 +67,30 @@ async def list_models(
try:
api_key = await key_manager.get_random_valid_key()
if not api_key:
raise HTTPException(status_code=503, detail="No valid API keys available to fetch models.")
raise HTTPException(
status_code=503, detail="No valid API keys available to fetch models."
)
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
models_data = await model_service.get_gemini_models(api_key)
if not models_data or "models" not in models_data:
raise HTTPException(status_code=500, detail="Failed to fetch base models list.")
raise HTTPException(
status_code=500, detail="Failed to fetch base models list."
)
models_json = deepcopy(models_data)
model_mapping = {x.get("name", "").split("/", maxsplit=1)[-1]: x for x in models_json.get("models", [])}
model_mapping = {
x.get("name", "").split("/", maxsplit=1)[-1]: x
for x in models_json.get("models", [])
}
def add_derived_model(base_name, suffix, display_suffix):
model = model_mapping.get(base_name)
if not model:
logger.warning(f"Base model '{base_name}' not found for derived model '{suffix}'.")
logger.warning(
f"Base model '{base_name}' not found for derived model '{suffix}'."
)
return
item = deepcopy(model)
item["name"] = f"models/{base_name}{suffix}"
@@ -86,7 +104,7 @@ async def list_models(
add_derived_model(name, "-search", " For Search")
if settings.IMAGE_MODELS:
for name in settings.IMAGE_MODELS:
add_derived_model(name, "-image", " For Image")
add_derived_model(name, "-image", " For Image")
if settings.THINKING_MODELS:
for name in settings.THINKING_MODELS:
add_derived_model(name, "-non-thinking", " Non Thinking")
@@ -98,7 +116,8 @@ async def list_models(
except Exception as e:
logger.error(f"Error getting Gemini models list: {str(e)}")
raise HTTPException(
status_code=500, detail="Internal server error while fetching Gemini models list"
status_code=500,
detail="Internal server error while fetching Gemini models list",
) from e
@@ -108,15 +127,19 @@ async def list_models(
async def generate_content(
model_name: str,
request: GeminiRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
chat_service: GeminiChatService = Depends(get_chat_service)
chat_service: GeminiChatService = Depends(get_chat_service),
):
"""处理 Gemini 非流式内容生成请求。"""
operation_name = "gemini_generate_content"
async with handle_route_errors(logger, operation_name, failure_message="Content generation failed"):
logger.info(f"Handling Gemini content generation request for model: {model_name}")
async with handle_route_errors(
logger, operation_name, failure_message="Content generation failed"
):
logger.info(
f"Handling Gemini content generation request for model: {model_name}"
)
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
# 检测是否为原生Gemini TTS请求
@@ -133,10 +156,13 @@ async def generate_content(
logger.info(f"TTS responseModalities: {response_modalities}")
logger.info(f"TTS speechConfig: {speech_config}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
# 所有原生TTS请求都使用TTS增强服务
if is_native_tts:
@@ -144,47 +170,51 @@ async def generate_content(
logger.info("Using native TTS enhanced service")
tts_service = await get_tts_chat_service(key_manager)
response = await tts_service.generate_content(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
return response
except Exception as e:
logger.warning(f"Native TTS processing failed, falling back to standard service: {e}")
logger.warning(
f"Native TTS processing failed, falling back to standard service: {e}"
)
# 使用标准服务处理所有其他请求(非TTS)
response = await chat_service.generate_content(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
return response
@router.post("/models/{model_name}:streamGenerateContent")
@router_v1beta.post("/models/{model_name}:streamGenerateContent")
@RetryHandler(key_arg="api_key")
async def stream_generate_content(
model_name: str,
request: GeminiRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
chat_service: GeminiChatService = Depends(get_chat_service)
chat_service: GeminiChatService = Depends(get_chat_service),
):
"""处理 Gemini 流式内容生成请求。"""
operation_name = "gemini_stream_generate_content"
async with handle_route_errors(logger, operation_name, failure_message="Streaming request initiation failed"):
logger.info(f"Handling Gemini streaming content generation for model: {model_name}")
async with handle_route_errors(
logger, operation_name, failure_message="Streaming request initiation failed"
):
logger.info(
f"Handling Gemini streaming content generation for model: {model_name}"
)
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
raw_stream = chat_service.stream_generate_content(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
try:
# 尝试获取第一条数据,判断是正常 SSE(data: 前缀)还是错误 JSON
@@ -195,12 +225,13 @@ async def stream_generate_content(
except Exception as e:
# 初始化流异常,直接返回 500 错误
return JSONResponse(
content={"error": {"code": 500, "message": str(e)}},
status_code=500
content={"error": {"code": e.args[0], "message": e.args[1]}},
status_code=e.args[0],
)
# 如果以 "data:" 开头,代表正常 SSE,将首块和后续块一起发送
if isinstance(first_chunk, str) and first_chunk.startswith("data:"):
async def combined():
yield first_chunk
async for chunk in raw_stream:
@@ -208,16 +239,6 @@ async def stream_generate_content(
return StreamingResponse(combined(), media_type="text/event-stream")
# 否则把首块当作错误 JSON 处理
try:
err = json.loads(first_chunk)
code = err.get("error", {}).get("code", 500)
except json.JSONDecodeError:
err = {"error": {"code": 500, "message": first_chunk}}
code = 500
return JSONResponse(content=err, status_code=code)
@router.post("/models/{model_name}:countTokens")
@router_v1beta.post("/models/{model_name}:countTokens")
@@ -225,53 +246,60 @@ async def stream_generate_content(
async def count_tokens(
model_name: str,
request: GeminiRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
chat_service: GeminiChatService = Depends(get_chat_service)
chat_service: GeminiChatService = Depends(get_chat_service),
):
"""处理 Gemini token 计数请求。"""
operation_name = "gemini_count_tokens"
async with handle_route_errors(logger, operation_name, failure_message="Token counting failed"):
async with handle_route_errors(
logger, operation_name, failure_message="Token counting failed"
):
logger.info(f"Handling Gemini token count request for model: {model_name}")
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
response = await chat_service.count_tokens(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
return response
@router.post("/models/{model_name}:embedContent")
@router_v1beta.post("/models/{model_name}:embedContent")
@RetryHandler(key_arg="api_key")
async def embed_content(
model_name: str,
request: GeminiEmbedRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
embedding_service: GeminiEmbeddingService = Depends(get_embedding_service)
embedding_service: GeminiEmbeddingService = Depends(get_embedding_service),
):
"""处理 Gemini 单一嵌入请求"""
operation_name = "gemini_embed_content"
async with handle_route_errors(logger, operation_name, failure_message="Embedding content generation failed"):
async with handle_route_errors(
logger, operation_name, failure_message="Embedding content generation failed"
):
logger.info(f"Handling Gemini embedding request for model: {model_name}")
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
response = await embedding_service.embed_content(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
return response
@@ -282,41 +310,48 @@ async def embed_content(
async def batch_embed_contents(
model_name: str,
request: GeminiBatchEmbedRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
embedding_service: GeminiEmbeddingService = Depends(get_embedding_service)
embedding_service: GeminiEmbeddingService = Depends(get_embedding_service),
):
"""处理 Gemini 批量嵌入请求"""
operation_name = "gemini_batch_embed_contents"
async with handle_route_errors(logger, operation_name, failure_message="Batch embedding content generation failed"):
async with handle_route_errors(
logger,
operation_name,
failure_message="Batch embedding content generation failed",
):
logger.info(f"Handling Gemini batch embedding request for model: {model_name}")
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
response = await embedding_service.batch_embed_contents(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
return response
@router.post("/reset-all-fail-counts")
async def reset_all_key_fail_counts(key_type: str = None, key_manager: KeyManager = Depends(get_key_manager)):
async def reset_all_key_fail_counts(
key_type: str = None, key_manager: KeyManager = Depends(get_key_manager)
):
"""批量重置Gemini API密钥的失败计数,可选择性地仅重置有效或无效密钥"""
logger.info("-" * 50 + "reset_all_gemini_key_fail_counts" + "-" * 50)
logger.info(f"Received reset request with key_type: {key_type}")
try:
# 获取分类后的密钥
keys_by_status = await key_manager.get_keys_by_status()
valid_keys = keys_by_status.get("valid_keys", {})
invalid_keys = keys_by_status.get("invalid_keys", {})
# 根据类型选择要重置的密钥
keys_to_reset = []
if key_type == "valid":
@@ -328,35 +363,45 @@ async def reset_all_key_fail_counts(key_type: str = None, key_manager: KeyManage
else:
# 重置所有密钥
await key_manager.reset_failure_counts()
return JSONResponse({"success": True, "message": "所有密钥的失败计数已重置"})
return JSONResponse(
{"success": True, "message": "所有密钥的失败计数已重置"}
)
# 批量重置指定类型的密钥
for key in keys_to_reset:
await key_manager.reset_key_failure_count(key)
return JSONResponse({
"success": True,
"message": f"{key_type}密钥的失败计数已重置",
"reset_count": len(keys_to_reset)
})
return JSONResponse(
{
"success": True,
"message": f"{key_type}密钥的失败计数已重置",
"reset_count": len(keys_to_reset),
}
)
except Exception as e:
logger.error(f"Failed to reset key failure counts: {str(e)}")
return JSONResponse({"success": False, "message": f"批量重置失败: {str(e)}"}, status_code=500)
return JSONResponse(
{"success": False, "message": f"批量重置失败: {str(e)}"}, status_code=500
)
@router.post("/reset-selected-fail-counts")
async def reset_selected_key_fail_counts(
request: ResetSelectedKeysRequest,
key_manager: KeyManager = Depends(get_key_manager)
key_manager: KeyManager = Depends(get_key_manager),
):
"""批量重置选定Gemini API密钥的失败计数"""
logger.info("-" * 50 + "reset_selected_gemini_key_fail_counts" + "-" * 50)
keys_to_reset = request.keys
key_type = request.key_type
logger.info(f"Received reset request for {len(keys_to_reset)} selected {key_type} keys.")
logger.info(
f"Received reset request for {len(keys_to_reset)} selected {key_type} keys."
)
if not keys_to_reset:
return JSONResponse({"success": False, "message": "没有提供需要重置的密钥"}, status_code=400)
return JSONResponse(
{"success": False, "message": "没有提供需要重置的密钥"}, status_code=400
)
reset_count = 0
errors = []
@@ -368,53 +413,79 @@ async def reset_selected_key_fail_counts(
if result:
reset_count += 1
else:
logger.warning(f"Key not found during selective reset: {redact_key_for_logging(key)}")
logger.warning(
f"Key not found during selective reset: {redact_key_for_logging(key)}"
)
except Exception as key_error:
logger.error(f"Error resetting key {redact_key_for_logging(key)}: {str(key_error)}")
logger.error(
f"Error resetting key {redact_key_for_logging(key)}: {str(key_error)}"
)
errors.append(f"Key {key}: {str(key_error)}")
if errors:
error_message = f"批量重置完成,但出现错误: {'; '.join(errors)}"
final_success = reset_count > 0
status_code = 207 if final_success and errors else 500
return JSONResponse({
"success": final_success,
"message": error_message,
"reset_count": reset_count
}, status_code=status_code)
error_message = f"批量重置完成,但出现错误: {'; '.join(errors)}"
final_success = reset_count > 0
status_code = 207 if final_success and errors else 500
return JSONResponse(
{
"success": final_success,
"message": error_message,
"reset_count": reset_count,
},
status_code=status_code,
)
return JSONResponse({
"success": True,
"message": f"成功重置 {reset_count} 个选定 {key_type} 密钥的失败计数",
"reset_count": reset_count
})
return JSONResponse(
{
"success": True,
"message": f"成功重置 {reset_count} 个选定 {key_type} 密钥的失败计数",
"reset_count": reset_count,
}
)
except Exception as e:
logger.error(f"Failed to process reset selected key failure counts request: {str(e)}")
return JSONResponse({"success": False, "message": f"批量重置处理失败: {str(e)}"}, status_code=500)
logger.error(
f"Failed to process reset selected key failure counts request: {str(e)}"
)
return JSONResponse(
{"success": False, "message": f"批量重置处理失败: {str(e)}"},
status_code=500,
)
@router.post("/reset-fail-count/{api_key}")
async def reset_key_fail_count(api_key: str, key_manager: KeyManager = Depends(get_key_manager)):
async def reset_key_fail_count(
api_key: str, key_manager: KeyManager = Depends(get_key_manager)
):
"""重置指定Gemini API密钥的失败计数"""
logger.info("-" * 50 + "reset_gemini_key_fail_count" + "-" * 50)
logger.info(f"Resetting failure count for API key: {redact_key_for_logging(api_key)}")
logger.info(
f"Resetting failure count for API key: {redact_key_for_logging(api_key)}"
)
try:
result = await key_manager.reset_key_failure_count(api_key)
if result:
return JSONResponse({"success": True, "message": "失败计数已重置"})
return JSONResponse({"success": False, "message": "未找到指定密钥"}, status_code=404)
return JSONResponse(
{"success": False, "message": "未找到指定密钥"}, status_code=404
)
except Exception as e:
logger.error(f"Failed to reset key failure count: {str(e)}")
return JSONResponse({"success": False, "message": f"重置失败: {str(e)}"}, status_code=500)
return JSONResponse(
{"success": False, "message": f"重置失败: {str(e)}"}, status_code=500
)
@router.post("/verify-key/{api_key}")
async def verify_key(api_key: str, chat_service: GeminiChatService = Depends(get_chat_service), key_manager: KeyManager = Depends(get_key_manager)):
async def verify_key(
api_key: str,
chat_service: GeminiChatService = Depends(get_chat_service),
key_manager: KeyManager = Depends(get_key_manager),
):
"""验证Gemini API密钥的有效性"""
logger.info("-" * 50 + "verify_gemini_key" + "-" * 50)
logger.info("Verifying API key validity")
try:
gemini_request = GeminiRequest(
contents=[
@@ -423,27 +494,27 @@ async def verify_key(api_key: str, chat_service: GeminiChatService = Depends(get
parts=[{"text": "hi"}],
)
],
generation_config={"temperature": 0.7, "topP": 1.0, "maxOutputTokens": 10}
generation_config={"temperature": 0.7, "topP": 1.0, "maxOutputTokens": 10},
)
response = await chat_service.generate_content(
settings.TEST_MODEL,
gemini_request,
api_key
settings.TEST_MODEL, gemini_request, api_key
)
if response:
# 如果密钥验证成功,则重置其失败计数
await key_manager.reset_key_failure_count(api_key)
return JSONResponse({"status": "valid"})
except Exception as e:
logger.error(f"Key verification failed: {str(e)}")
async with key_manager.failure_count_lock:
if api_key in key_manager.key_failure_counts:
key_manager.key_failure_counts[api_key] += 1
logger.warning(f"Verification exception for key: {redact_key_for_logging(api_key)}, incrementing failure count")
logger.warning(
f"Verification exception for key: {redact_key_for_logging(api_key)}, incrementing failure count"
)
return JSONResponse({"status": "invalid", "error": str(e)})
@@ -451,15 +522,19 @@ async def verify_key(api_key: str, chat_service: GeminiChatService = Depends(get
async def verify_selected_keys(
request: VerifySelectedKeysRequest,
chat_service: GeminiChatService = Depends(get_chat_service),
key_manager: KeyManager = Depends(get_key_manager)
key_manager: KeyManager = Depends(get_key_manager),
):
"""批量验证选定Gemini API密钥的有效性"""
logger.info("-" * 50 + "verify_selected_gemini_keys" + "-" * 50)
keys_to_verify = request.keys
logger.info(f"Received verification request for {len(keys_to_verify)} selected keys.")
logger.info(
f"Received verification request for {len(keys_to_verify)} selected keys."
)
if not keys_to_verify:
return JSONResponse({"success": False, "message": "没有提供需要验证的密钥"}, status_code=400)
return JSONResponse(
{"success": False, "message": "没有提供需要验证的密钥"}, status_code=400
)
successful_keys = []
failed_keys = {}
@@ -470,12 +545,14 @@ async def verify_selected_keys(
try:
gemini_request = GeminiRequest(
contents=[GeminiContent(role="user", parts=[{"text": "hi"}])],
generation_config={"temperature": 0.7, "topP": 1.0, "maxOutputTokens": 10}
generation_config={
"temperature": 0.7,
"topP": 1.0,
"maxOutputTokens": 10,
},
)
await chat_service.generate_content(
settings.TEST_MODEL,
gemini_request,
api_key
settings.TEST_MODEL, gemini_request, api_key
)
successful_keys.append(api_key)
# 如果密钥验证成功,则重置其失败计数
@@ -483,14 +560,20 @@ async def verify_selected_keys(
return api_key, "valid", None
except Exception as e:
error_message = str(e)
logger.warning(f"Key verification failed for {redact_key_for_logging(api_key)}: {error_message}")
logger.warning(
f"Key verification failed for {redact_key_for_logging(api_key)}: {error_message}"
)
async with key_manager.failure_count_lock:
if api_key in key_manager.key_failure_counts:
key_manager.key_failure_counts[api_key] += 1
logger.warning(f"Bulk verification exception for key: {redact_key_for_logging(api_key)}, incrementing failure count")
logger.warning(
f"Bulk verification exception for key: {redact_key_for_logging(api_key)}, incrementing failure count"
)
else:
key_manager.key_failure_counts[api_key] = 1
logger.warning(f"Bulk verification exception for key: {redact_key_for_logging(api_key)}, initializing failure count to 1")
key_manager.key_failure_counts[api_key] = 1
logger.warning(
f"Bulk verification exception for key: {redact_key_for_logging(api_key)}, initializing failure count to 1"
)
failed_keys[api_key] = error_message
return api_key, "invalid", error_message
@@ -499,34 +582,42 @@ async def verify_selected_keys(
for result in results:
if isinstance(result, Exception):
logger.error(f"An unexpected error occurred during bulk verification task: {result}")
logger.error(
f"An unexpected error occurred during bulk verification task: {result}"
)
elif result:
if not isinstance(result, Exception) and result:
key, status, error = result
elif isinstance(result, Exception):
logger.error(f"Task execution error during bulk verification: {result}")
if not isinstance(result, Exception) and result:
key, status, error = result
elif isinstance(result, Exception):
logger.error(f"Task execution error during bulk verification: {result}")
valid_count = len(successful_keys)
invalid_count = len(failed_keys)
logger.info(f"Bulk verification finished. Valid: {valid_count}, Invalid: {invalid_count}")
logger.info(
f"Bulk verification finished. Valid: {valid_count}, Invalid: {invalid_count}"
)
if failed_keys:
message = f"批量验证完成。成功: {valid_count}, 失败: {invalid_count}"
return JSONResponse({
"success": True,
"message": message,
"successful_keys": successful_keys,
"failed_keys": failed_keys,
"valid_count": valid_count,
"invalid_count": invalid_count
})
return JSONResponse(
{
"success": True,
"message": message,
"successful_keys": successful_keys,
"failed_keys": failed_keys,
"valid_count": valid_count,
"invalid_count": invalid_count,
}
)
else:
message = f"批量验证成功完成。所有 {valid_count} 个密钥均有效。"
return JSONResponse({
"success": True,
"message": message,
"successful_keys": successful_keys,
"failed_keys": {},
"valid_count": valid_count,
"invalid_count": 0
})
return JSONResponse(
{
"success": True,
"message": message,
"successful_keys": successful_keys,
"failed_keys": {},
"valid_count": valid_count,
"invalid_count": 0,
}
)
+46 -14
View File
@@ -1,5 +1,5 @@
from fastapi import APIRouter, Depends
from fastapi.responses import StreamingResponse
from fastapi.responses import JSONResponse, StreamingResponse
from app.config.config import settings
from app.core.security import SecurityService
@@ -8,19 +8,21 @@ from app.domain.openai_models import (
EmbeddingRequest,
ImageGenerationRequest,
)
from app.handler.retry_handler import RetryHandler
from app.handler.error_handler import handle_route_errors
from app.handler.retry_handler import RetryHandler
from app.log.logger import get_openai_compatible_logger
from app.service.key.key_manager import KeyManager, get_key_manager_instance
from app.service.openai_compatiable.openai_compatiable_service import OpenAICompatiableService
from app.service.openai_compatiable.openai_compatiable_service import (
OpenAICompatiableService,
)
from app.utils.helpers import redact_key_for_logging
router = APIRouter()
logger = get_openai_compatible_logger()
security_service = SecurityService()
async def get_key_manager():
return await get_key_manager_instance()
@@ -38,7 +40,7 @@ async def get_openai_service(key_manager: KeyManager = Depends(get_key_manager))
@router.get("/openai/v1/models")
async def list_models(
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
key_manager: KeyManager = Depends(get_key_manager),
openai_service: OpenAICompatiableService = Depends(get_openai_service),
):
@@ -47,6 +49,7 @@ async def list_models(
async with handle_route_errors(logger, operation_name):
logger.info("Handling models list request")
api_key = await key_manager.get_random_valid_key()
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
return await openai_service.get_models(api_key)
@@ -55,7 +58,7 @@ async def list_models(
@RetryHandler(key_arg="api_key")
async def chat_completion(
request: ChatRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
api_key: str = Depends(get_next_working_key_wrapper),
key_manager: KeyManager = Depends(get_key_manager),
openai_service: OpenAICompatiableService = Depends(get_openai_service),
@@ -70,28 +73,56 @@ async def chat_completion(
async with handle_route_errors(logger, operation_name):
logger.info(f"Handling chat completion request for model: {request.model}")
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(current_api_key)}")
raw_response = None
if is_image_chat:
response = await openai_service.create_image_chat_completion(request, current_api_key)
return response
raw_response = await openai_service.create_image_chat_completion(
request, current_api_key
)
else:
response = await openai_service.create_chat_completion(request, current_api_key)
if request.stream:
return StreamingResponse(response, media_type="text/event-stream")
return response
raw_response = await openai_service.create_chat_completion(
request, current_api_key
)
if request.stream:
try:
# 尝试获取第一条数据,判断是正常 SSE(data: 前缀)还是错误 JSON
first_chunk = await raw_response.__anext__()
except StopAsyncIteration:
# 如果流直接结束,退回标准 SSE 输出
return StreamingResponse(raw_response, media_type="text/event-stream")
except Exception as e:
# 初始化流异常,直接返回 500 错误
return JSONResponse(
content={"error": {"code": e.args[0], "message": e.args[1]}},
status_code=e.args[0],
)
# 如果以 "data:" 开头,代表正常 SSE,将首块和后续块一起发送
if isinstance(first_chunk, str) and first_chunk.startswith("data:"):
async def combined():
yield first_chunk
async for chunk in raw_response:
yield chunk
return StreamingResponse(combined(), media_type="text/event-stream")
else:
return raw_response
@router.post("/openai/v1/images/generations")
async def generate_image(
request: ImageGenerationRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
openai_service: OpenAICompatiableService = Depends(get_openai_service),
):
"""处理图像生成请求。"""
operation_name = "generate_image"
async with handle_route_errors(logger, operation_name):
logger.info(f"Handling image generation request for prompt: {request.prompt}")
logger.info(f"Using allowed token: {allowed_token}")
request.model = settings.CREATE_IMAGE_MODEL
return await openai_service.generate_images(request)
@@ -99,7 +130,7 @@ async def generate_image(
@router.post("/openai/v1/embeddings")
async def embedding(
request: EmbeddingRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
key_manager: KeyManager = Depends(get_key_manager),
openai_service: OpenAICompatiableService = Depends(get_openai_service),
):
@@ -108,6 +139,7 @@ async def embedding(
async with handle_route_errors(logger, operation_name):
logger.info(f"Handling embedding request for model: {request.model}")
api_key = await key_manager.get_next_working_key()
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
return await openai_service.create_embeddings(
input_text=request.input, model=request.model, api_key=api_key
+44 -16
View File
@@ -1,5 +1,5 @@
from fastapi import APIRouter, Depends, HTTPException, Response
from fastapi.responses import StreamingResponse
from fastapi.responses import JSONResponse, StreamingResponse
from app.config.config import settings
from app.core.security import SecurityService
@@ -9,15 +9,15 @@ from app.domain.openai_models import (
ImageGenerationRequest,
TTSRequest,
)
from app.handler.retry_handler import RetryHandler
from app.handler.error_handler import handle_route_errors
from app.handler.retry_handler import RetryHandler
from app.log.logger import get_openai_logger
from app.service.chat.openai_chat_service import OpenAIChatService
from app.service.embedding.embedding_service import EmbeddingService
from app.service.image.image_create_service import ImageCreateService
from app.service.tts.tts_service import TTSService
from app.service.key.key_manager import KeyManager, get_key_manager_instance
from app.service.model.model_service import ModelService
from app.service.tts.tts_service import TTSService
from app.utils.helpers import redact_key_for_logging
router = APIRouter()
@@ -53,7 +53,7 @@ async def get_tts_service():
@router.get("/v1/models")
@router.get("/hf/v1/models")
async def list_models(
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
key_manager: KeyManager = Depends(get_key_manager),
):
"""获取可用的 OpenAI 模型列表 (兼容 Gemini 和 OpenAI)。"""
@@ -61,6 +61,7 @@ async def list_models(
async with handle_route_errors(logger, operation_name):
logger.info("Handling models list request")
api_key = await key_manager.get_random_valid_key()
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
return await model_service.get_gemini_openai_models(api_key)
@@ -70,7 +71,7 @@ async def list_models(
@RetryHandler(key_arg="api_key")
async def chat_completion(
request: ChatRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
api_key: str = Depends(get_next_working_key_wrapper),
key_manager: KeyManager = Depends(get_key_manager),
chat_service: OpenAIChatService = Depends(get_openai_chat_service),
@@ -92,23 +93,48 @@ async def chat_completion(
status_code=400, detail=f"Model {request.model} is not supported"
)
raw_response = None
if is_image_chat:
response = await chat_service.create_image_chat_completion(request, current_api_key)
if request.stream:
return StreamingResponse(response, media_type="text/event-stream")
return response
raw_response = await chat_service.create_image_chat_completion(
request, current_api_key
)
else:
response = await chat_service.create_chat_completion(request, current_api_key)
if request.stream:
return StreamingResponse(response, media_type="text/event-stream")
return response
raw_response = await chat_service.create_chat_completion(
request, current_api_key
)
if request.stream:
try:
# 尝试获取第一条数据,判断是正常 SSE(data: 前缀)还是错误 JSON
first_chunk = await raw_response.__anext__()
except StopAsyncIteration:
# 如果流直接结束,退回标准 SSE 输出
return StreamingResponse(raw_response, media_type="text/event-stream")
except Exception as e:
# 初始化流异常,直接返回 500 错误
return JSONResponse(
content={"error": {"code": e.args[0], "message": e.args[1]}},
status_code=e.args[0],
)
# 如果以 "data:" 开头,代表正常 SSE,将首块和后续块一起发送
if isinstance(first_chunk, str) and first_chunk.startswith("data:"):
async def combined():
yield first_chunk
async for chunk in raw_response:
yield chunk
return StreamingResponse(combined(), media_type="text/event-stream")
else:
return raw_response
@router.post("/v1/images/generations")
@router.post("/hf/v1/images/generations")
async def generate_image(
request: ImageGenerationRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
):
"""处理 OpenAI 图像生成请求。"""
operation_name = "generate_image"
@@ -122,7 +148,7 @@ async def generate_image(
@router.post("/hf/v1/embeddings")
async def embedding(
request: EmbeddingRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
key_manager: KeyManager = Depends(get_key_manager),
):
"""处理 OpenAI 文本嵌入请求。"""
@@ -130,6 +156,7 @@ async def embedding(
async with handle_route_errors(logger, operation_name):
logger.info(f"Handling embedding request for model: {request.model}")
api_key = await key_manager.get_next_working_key()
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
response = await embedding_service.create_embedding(
input_text=request.input, model=request.model, api_key=api_key
@@ -162,7 +189,7 @@ async def get_keys_list(
@router.post("/hf/v1/audio/speech")
async def text_to_speech(
request: TTSRequest,
_=Depends(security_service.verify_authorization),
allowed_token=Depends(security_service.verify_authorization),
api_key: str = Depends(get_next_working_key_wrapper),
tts_service: TTSService = Depends(get_tts_service),
):
@@ -171,6 +198,7 @@ async def text_to_speech(
async with handle_route_errors(logger, operation_name):
logger.info(f"Handling TTS request for model: {request.model}")
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
audio_data = await tts_service.create_tts(request, api_key)
return Response(content=audio_data, media_type="audio/wav")
+76 -32
View File
@@ -1,16 +1,18 @@
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse
from copy import deepcopy
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import JSONResponse, StreamingResponse
from app.config.config import settings
from app.log.logger import get_vertex_express_logger
from app.core.constants import API_VERSION
from app.core.security import SecurityService
from app.domain.gemini_models import GeminiRequest
from app.handler.error_handler import handle_route_errors
from app.handler.retry_handler import RetryHandler
from app.log.logger import get_vertex_express_logger
from app.service.chat.vertex_express_chat_service import GeminiChatService
from app.service.key.key_manager import KeyManager, get_key_manager_instance
from app.service.model.model_service import ModelService
from app.handler.retry_handler import RetryHandler
from app.handler.error_handler import handle_route_errors
from app.core.constants import API_VERSION
from app.utils.helpers import redact_key_for_logging
router = APIRouter(prefix=f"/vertex-express/{API_VERSION}")
@@ -37,8 +39,8 @@ async def get_chat_service(key_manager: KeyManager = Depends(get_key_manager)):
@router.get("/models")
async def list_models(
_=Depends(security_service.verify_key_or_goog_api_key),
key_manager: KeyManager = Depends(get_key_manager)
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
key_manager: KeyManager = Depends(get_key_manager),
):
"""获取可用的 Gemini 模型列表,并根据配置添加衍生模型(搜索、图像、非思考)。"""
operation_name = "list_gemini_models"
@@ -48,20 +50,30 @@ async def list_models(
try:
api_key = await key_manager.get_random_valid_key()
if not api_key:
raise HTTPException(status_code=503, detail="No valid API keys available to fetch models.")
raise HTTPException(
status_code=503, detail="No valid API keys available to fetch models."
)
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
models_data = await model_service.get_gemini_models(api_key)
if not models_data or "models" not in models_data:
raise HTTPException(status_code=500, detail="Failed to fetch base models list.")
raise HTTPException(
status_code=500, detail="Failed to fetch base models list."
)
models_json = deepcopy(models_data)
model_mapping = {x.get("name", "").split("/", maxsplit=1)[-1]: x for x in models_json.get("models", [])}
model_mapping = {
x.get("name", "").split("/", maxsplit=1)[-1]: x
for x in models_json.get("models", [])
}
def add_derived_model(base_name, suffix, display_suffix):
model = model_mapping.get(base_name)
if not model:
logger.warning(f"Base model '{base_name}' not found for derived model '{suffix}'.")
logger.warning(
f"Base model '{base_name}' not found for derived model '{suffix}'."
)
return
item = deepcopy(model)
item["name"] = f"models/{base_name}{suffix}"
@@ -75,7 +87,7 @@ async def list_models(
add_derived_model(name, "-search", " For Search")
if settings.IMAGE_MODELS:
for name in settings.IMAGE_MODELS:
add_derived_model(name, "-image", " For Image")
add_derived_model(name, "-image", " For Image")
if settings.THINKING_MODELS:
for name in settings.THINKING_MODELS:
add_derived_model(name, "-non-thinking", " Non Thinking")
@@ -87,7 +99,8 @@ async def list_models(
except Exception as e:
logger.error(f"Error getting Gemini models list: {str(e)}")
raise HTTPException(
status_code=500, detail="Internal server error while fetching Gemini models list"
status_code=500,
detail="Internal server error while fetching Gemini models list",
) from e
@@ -96,25 +109,30 @@ async def list_models(
async def generate_content(
model_name: str,
request: GeminiRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
chat_service: GeminiChatService = Depends(get_chat_service)
chat_service: GeminiChatService = Depends(get_chat_service),
):
"""处理 Gemini 非流式内容生成请求。"""
operation_name = "gemini_generate_content"
async with handle_route_errors(logger, operation_name, failure_message="Content generation failed"):
logger.info(f"Handling Gemini content generation request for model: {model_name}")
async with handle_route_errors(
logger, operation_name, failure_message="Content generation failed"
):
logger.info(
f"Handling Gemini content generation request for model: {model_name}"
)
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
response = await chat_service.generate_content(
model=model_name,
request=request,
api_key=api_key
model=model_name, request=request, api_key=api_key
)
return response
@@ -124,24 +142,50 @@ async def generate_content(
async def stream_generate_content(
model_name: str,
request: GeminiRequest,
_=Depends(security_service.verify_key_or_goog_api_key),
allowed_token=Depends(security_service.verify_key_or_goog_api_key),
api_key: str = Depends(get_next_working_key),
key_manager: KeyManager = Depends(get_key_manager),
chat_service: GeminiChatService = Depends(get_chat_service)
chat_service: GeminiChatService = Depends(get_chat_service),
):
"""处理 Gemini 流式内容生成请求。"""
operation_name = "gemini_stream_generate_content"
async with handle_route_errors(logger, operation_name, failure_message="Streaming request initiation failed"):
logger.info(f"Handling Gemini streaming content generation for model: {model_name}")
async with handle_route_errors(
logger, operation_name, failure_message="Streaming request initiation failed"
):
logger.info(
f"Handling Gemini streaming content generation for model: {model_name}"
)
logger.debug(f"Request: \n{request.model_dump_json(indent=2)}")
logger.info(f"Using allowed token: {allowed_token}")
logger.info(f"Using API key: {redact_key_for_logging(api_key)}")
if not await model_service.check_model_support(model_name):
raise HTTPException(status_code=400, detail=f"Model {model_name} is not supported")
raise HTTPException(
status_code=400, detail=f"Model {model_name} is not supported"
)
response_stream = chat_service.stream_generate_content(
model=model_name,
request=request,
api_key=api_key
raw_stream = chat_service.stream_generate_content(
model=model_name, request=request, api_key=api_key
)
return StreamingResponse(response_stream, media_type="text/event-stream")
try:
# 尝试获取第一条数据,判断是正常 SSE(data: 前缀)还是错误 JSON
first_chunk = await raw_stream.__anext__()
except StopAsyncIteration:
# 如果流直接结束,退回标准 SSE 输出
return StreamingResponse(raw_stream, media_type="text/event-stream")
except Exception as e:
# 初始化流异常,直接返回 500 错误
return JSONResponse(
content={"error": {"code": e.args[0], "message": e.args[1]}},
status_code=e.args[0],
)
# 如果以 "data:" 开头,代表正常 SSE,将首块和后续块一起发送
if isinstance(first_chunk, str) and first_chunk.startswith("data:"):
async def combined():
yield first_chunk
async for chunk in raw_stream:
yield chunk
return StreamingResponse(combined(), media_type="text/event-stream")
+8 -45
View File
@@ -365,13 +365,9 @@ class GeminiChatService:
return self.response_handler.handle_response(response, model, stream=False)
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"Normal API call failed with error: {error_log_msg}")
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=api_key,
@@ -416,13 +412,9 @@ class GeminiChatService:
return response
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"Count tokens API call failed with error: {error_log_msg}")
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=api_key,
@@ -470,7 +462,6 @@ class GeminiChatService:
is_success = False
status_code = None
final_api_key = api_key
last_error_msg = None
while retries < max_retries:
request_datetime = datetime.datetime.now()
@@ -509,16 +500,11 @@ class GeminiChatService:
except Exception as e:
retries += 1
is_success = False
error_log_msg = str(e)
last_error_msg = error_log_msg
status_code = e.args[0]
error_log_msg = e.args[1]
logger.warning(
f"Streaming API call failed with error: {error_log_msg}. Attempt {retries} of {max_retries}"
)
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=current_attempt_key,
@@ -539,11 +525,11 @@ class GeminiChatService:
)
else:
logger.error(f"No valid API key available after {retries} retries.")
break
raise
if retries >= max_retries:
logger.error(f"Max retries ({max_retries}) reached for streaming.")
break
raise
finally:
end_time = time.perf_counter()
latency_ms = int((end_time - start_time) * 1000)
@@ -555,26 +541,3 @@ class GeminiChatService:
latency_ms=latency_ms,
request_time=request_datetime,
)
# Emit final error SSE event if all retries failed
if not is_success:
# 从错误消息中提取嵌套JSON
parsed_error = None
if last_error_msg:
try:
# 查找JSON起始位置
json_start = last_error_msg.find('{')
if json_start != -1:
json_str = last_error_msg[json_start:]
parsed_error = json.loads(json_str)
except json.JSONDecodeError:
pass
error_data = {
"error": {
"code": parsed_error['error']['code'] if (parsed_error and 'error' in parsed_error and 'code' in parsed_error['error']) else (status_code or 500),
"message": parsed_error['error']['message'] if (parsed_error and 'error' in parsed_error and 'message' in parsed_error['error']) else (last_error_msg or "Streaming failed"),
"status": parsed_error['error']['status'] if (parsed_error and 'error' in parsed_error and 'status' in parsed_error['error']) else "INTERNAL"
}
}
yield json.dumps(error_data, ensure_ascii=False)
+13 -31
View File
@@ -3,7 +3,6 @@
import asyncio
import datetime
import json
import re
import time
from copy import deepcopy
from typing import Any, AsyncGenerator, Dict, List, Optional, Union
@@ -339,7 +338,8 @@ class OpenAIChatService:
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"API call failed for model {model}: {error_log_msg}")
# 特别记录 max_tokens 相关的错误
@@ -353,9 +353,6 @@ class OpenAIChatService:
if "parts" in error_log_msg:
logger.error("This is likely a response processing error")
match = re.search(r"status code (\d+)", error_log_msg)
status_code = int(match.group(1)) if match else 500
await add_error_log(
gemini_key=api_key,
model_name=model,
@@ -540,20 +537,12 @@ class OpenAIChatService:
except Exception as e:
retries += 1
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.warning(
f"Streaming API call failed with error: {error_log_msg}. Attempt {retries} of {max_retries} with key {current_attempt_key}"
)
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
if isinstance(e, asyncio.TimeoutError):
status_code = 408
else:
status_code = 500
await add_error_log(
gemini_key=current_attempt_key,
model_name=model,
@@ -577,7 +566,7 @@ class OpenAIChatService:
logger.error(
f"No valid API key available after {retries} retries, ceasing attempts for this request."
)
break
raise
else:
logger.error(
"KeyManager not available, cannot switch API key. Ceasing attempts for this request."
@@ -588,6 +577,7 @@ class OpenAIChatService:
logger.error(
f"Max retries ({max_retries}) reached for streaming model {model}."
)
raise
finally:
end_time = time.perf_counter()
latency_ms = int((end_time - start_time) * 1000)
@@ -600,13 +590,6 @@ class OpenAIChatService:
request_time=request_datetime,
)
if not is_success:
logger.error(
f"Streaming failed permanently for model {model} after {retries} attempts."
)
yield f"data: {json.dumps({'error': f'Streaming failed after {retries} retries.'})}\n\n"
yield "data: [DONE]\n\n"
async def create_image_chat_completion(
self, request: ChatRequest, api_key: str
) -> Union[Dict[str, Any], AsyncGenerator[str, None]]:
@@ -665,9 +648,9 @@ class OpenAIChatService:
yield "data: [DONE]\n\n"
except Exception as e:
is_success = False
error_log_msg = f"Stream image completion failed for model {model}: {e}"
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(error_log_msg)
status_code = 500
await add_error_log(
gemini_key=api_key,
model_name=model,
@@ -677,8 +660,7 @@ class OpenAIChatService:
request_msg={"image_data_truncated": image_data[:1000]},
request_datetime=request_datetime,
)
yield f"data: {json.dumps({'error': error_log_msg})}\n\n"
yield "data: [DONE]\n\n"
raise
finally:
end_time = time.perf_counter()
latency_ms = int((end_time - start_time) * 1000)
@@ -716,9 +698,9 @@ class OpenAIChatService:
return result
except Exception as e:
is_success = False
error_log_msg = f"Normal image completion failed for model {model}: {e}"
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(error_log_msg)
status_code = 500
await add_error_log(
gemini_key=api_key,
model_name=model,
@@ -728,7 +710,7 @@ class OpenAIChatService:
request_msg={"image_data_truncated": image_data[:1000]},
request_datetime=request_datetime,
)
raise e
raise
finally:
end_time = time.perf_counter()
latency_ms = int((end_time - start_time) * 1000)
@@ -742,4 +724,4 @@ class OpenAIChatService:
status_code=status_code,
latency_ms=latency_ms,
request_time=request_datetime,
)
)
@@ -2,7 +2,6 @@
import datetime
import json
import re
import time
from typing import Any, AsyncGenerator, Dict, List
@@ -278,13 +277,9 @@ class GeminiChatService:
return self.response_handler.handle_response(response, model, stream=False)
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"Normal API call failed with error: {error_log_msg}")
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=api_key,
@@ -356,15 +351,11 @@ class GeminiChatService:
except Exception as e:
retries += 1
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.warning(
f"Streaming API call failed with error: {error_log_msg}. Attempt {retries} of {max_retries}"
)
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=current_attempt_key,
@@ -385,11 +376,11 @@ class GeminiChatService:
)
else:
logger.error(f"No valid API key available after {retries} retries.")
break
raise
if retries >= max_retries:
logger.error(f"Max retries ({max_retries}) reached for streaming.")
break
raise
finally:
end_time = time.perf_counter()
latency_ms = int((end_time - start_time) * 1000)
@@ -400,4 +391,4 @@ class GeminiChatService:
status_code=status_code,
latency_ms=latency_ms,
request_time=request_datetime,
)
)
+103 -66
View File
@@ -1,24 +1,31 @@
# app/services/chat/api_client.py
from typing import Dict, Any, AsyncGenerator, Optional
import httpx
import random
from abc import ABC, abstractmethod
from typing import Any, AsyncGenerator, Dict, Optional
import httpx
from app.config.config import settings
from app.log.logger import get_api_client_logger
from app.core.constants import DEFAULT_TIMEOUT
from app.log.logger import get_api_client_logger
logger = get_api_client_logger()
class ApiClient(ABC):
"""API客户端基类"""
@abstractmethod
async def generate_content(self, payload: Dict[str, Any], model: str, api_key: str) -> Dict[str, Any]:
async def generate_content(
self, payload: Dict[str, Any], model: str, api_key: str
) -> Dict[str, Any]:
pass
@abstractmethod
async def stream_generate_content(self, payload: Dict[str, Any], model: str, api_key: str) -> AsyncGenerator[str, None]:
async def stream_generate_content(
self, payload: Dict[str, Any], model: str, api_key: str
) -> AsyncGenerator[str, None]:
pass
@@ -50,7 +57,7 @@ class GeminiApiClient(ApiClient):
async def get_models(self, api_key: str) -> Optional[Dict[str, Any]]:
"""获取可用的 Gemini 模型列表"""
timeout = httpx.Timeout(timeout=5)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -73,11 +80,13 @@ class GeminiApiClient(ApiClient):
except httpx.RequestError as e:
logger.error(f"请求模型列表失败: {e}")
return None
async def generate_content(self, payload: Dict[str, Any], model: str, api_key: str) -> Dict[str, Any]:
async def generate_content(
self, payload: Dict[str, Any], model: str, api_key: str
) -> Dict[str, Any]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
model = self._get_real_model(model)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -85,42 +94,46 @@ class GeminiApiClient(ApiClient):
else:
proxy_to_use = random.choice(settings.PROXIES)
logger.info(f"Using proxy for getting models: {proxy_to_use}")
headers = self._prepare_headers()
async with httpx.AsyncClient(timeout=timeout, proxy=proxy_to_use) as client:
url = f"{self.base_url}/models/{model}:generateContent?key={api_key}"
try:
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
logger.error(f"API call failed - Status: {response.status_code}, Content: {error_content}")
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
logger.error(
f"API call failed - Status: {response.status_code}, Content: {error_content}"
)
raise Exception(response.status_code, error_content)
response_data = response.json()
# 检查响应结构的基本信息
if not response_data.get("candidates"):
logger.warning("No candidates found in API response")
return response_data
except httpx.TimeoutException as e:
logger.error(f"Request timeout: {e}")
raise Exception(f"Request timeout: {e}")
raise Exception(500, f"Request timeout: {e}")
except httpx.RequestError as e:
logger.error(f"Request error: {e}")
raise Exception(f"Request error: {e}")
raise Exception(500, f"Request error: {e}")
except Exception as e:
logger.error(f"Unexpected error: {e}")
raise
raise Exception(500, f"Unexpected error: {e}")
async def stream_generate_content(self, payload: Dict[str, Any], model: str, api_key: str) -> AsyncGenerator[str, None]:
async def stream_generate_content(
self, payload: Dict[str, Any], model: str, api_key: str
) -> AsyncGenerator[str, None]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
model = self._get_real_model(model)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -132,15 +145,19 @@ class GeminiApiClient(ApiClient):
headers = self._prepare_headers()
async with httpx.AsyncClient(timeout=timeout, proxy=proxy_to_use) as client:
url = f"{self.base_url}/models/{model}:streamGenerateContent?alt=sse&key={api_key}"
async with client.stream(method="POST", url=url, json=payload, headers=headers) as response:
async with client.stream(
method="POST", url=url, json=payload, headers=headers
) as response:
if response.status_code != 200:
error_content = await response.aread()
error_msg = error_content.decode("utf-8")
raise Exception(f"API call failed with status code {response.status_code}, {error_msg}")
raise Exception(response.status_code, error_msg)
async for line in response.aiter_lines():
yield line
async def count_tokens(self, payload: Dict[str, Any], model: str, api_key: str) -> Dict[str, Any]:
async def count_tokens(
self, payload: Dict[str, Any], model: str, api_key: str
) -> Dict[str, Any]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
model = self._get_real_model(model)
@@ -158,14 +175,16 @@ class GeminiApiClient(ApiClient):
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
raise Exception(response.status_code, error_content)
return response.json()
async def embed_content(self, payload: Dict[str, Any], model: str, api_key: str) -> Dict[str, Any]:
async def embed_content(
self, payload: Dict[str, Any], model: str, api_key: str
) -> Dict[str, Any]:
"""单一嵌入内容生成"""
timeout = httpx.Timeout(self.timeout, read=self.timeout)
model = self._get_real_model(model)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -177,32 +196,36 @@ class GeminiApiClient(ApiClient):
headers = self._prepare_headers()
async with httpx.AsyncClient(timeout=timeout, proxy=proxy_to_use) as client:
url = f"{self.base_url}/models/{model}:embedContent?key={api_key}"
try:
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
logger.error(f"Embedding API call failed - Status: {response.status_code}, Content: {error_content}")
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
logger.error(
f"Embedding API call failed - Status: {response.status_code}, Content: {error_content}"
)
raise Exception(response.status_code, error_content)
return response.json()
except httpx.TimeoutException as e:
logger.error(f"Embedding request timeout: {e}")
raise Exception(f"Request timeout: {e}")
raise Exception(500, f"Request timeout: {e}")
except httpx.RequestError as e:
logger.error(f"Embedding request error: {e}")
raise Exception(f"Request error: {e}")
raise Exception(500, f"Request error: {e}")
except Exception as e:
logger.error(f"Unexpected embedding error: {e}")
raise
raise Exception(500, f"Unexpected embedding error: {e}")
async def batch_embed_contents(self, payload: Dict[str, Any], model: str, api_key: str) -> Dict[str, Any]:
async def batch_embed_contents(
self, payload: Dict[str, Any], model: str, api_key: str
) -> Dict[str, Any]:
"""批量嵌入内容生成"""
timeout = httpx.Timeout(self.timeout, read=self.timeout)
model = self._get_real_model(model)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -214,26 +237,28 @@ class GeminiApiClient(ApiClient):
headers = self._prepare_headers()
async with httpx.AsyncClient(timeout=timeout, proxy=proxy_to_use) as client:
url = f"{self.base_url}/models/{model}:batchEmbedContents?key={api_key}"
try:
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
logger.error(f"Batch embedding API call failed - Status: {response.status_code}, Content: {error_content}")
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
logger.error(
f"Batch embedding API call failed - Status: {response.status_code}, Content: {error_content}"
)
raise Exception(response.status_code, error_content)
return response.json()
except httpx.TimeoutException as e:
logger.error(f"Batch embedding request timeout: {e}")
raise Exception(f"Request timeout: {e}")
raise Exception(500, f"Request timeout: {e}")
except httpx.RequestError as e:
logger.error(f"Batch embedding request error: {e}")
raise Exception(f"Request error: {e}")
raise Exception(500, f"Request error: {e}")
except Exception as e:
logger.error(f"Unexpected batch embedding error: {e}")
raise
raise Exception(500, f"Unexpected batch embedding error: {e}")
class OpenaiApiClient(ApiClient):
@@ -242,7 +267,7 @@ class OpenaiApiClient(ApiClient):
def __init__(self, base_url: str, timeout: int = DEFAULT_TIMEOUT):
self.base_url = base_url
self.timeout = timeout
def _prepare_headers(self, api_key: str) -> Dict[str, str]:
headers = {"Authorization": f"Bearer {api_key}"}
if settings.CUSTOM_HEADERS:
@@ -267,12 +292,16 @@ class OpenaiApiClient(ApiClient):
response = await client.get(url, headers=headers)
if response.status_code != 200:
error_content = response.text
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
raise Exception(response.status_code, error_content)
return response.json()
async def generate_content(self, payload: Dict[str, Any], api_key: str) -> Dict[str, Any]:
async def generate_content(
self, payload: Dict[str, Any], api_key: str
) -> Dict[str, Any]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
logger.info(f"settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY: {settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY}")
logger.info(
f"settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY: {settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY}"
)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -287,10 +316,12 @@ class OpenaiApiClient(ApiClient):
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
raise Exception(response.status_code, error_content)
return response.json()
async def stream_generate_content(self, payload: Dict[str, Any], api_key: str) -> AsyncGenerator[str, None]:
async def stream_generate_content(
self, payload: Dict[str, Any], api_key: str
) -> AsyncGenerator[str, None]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
proxy_to_use = None
if settings.PROXIES:
@@ -303,17 +334,21 @@ class OpenaiApiClient(ApiClient):
headers = self._prepare_headers(api_key)
async with httpx.AsyncClient(timeout=timeout, proxy=proxy_to_use) as client:
url = f"{self.base_url}/openai/chat/completions"
async with client.stream(method="POST", url=url, json=payload, headers=headers) as response:
async with client.stream(
method="POST", url=url, json=payload, headers=headers
) as response:
if response.status_code != 200:
error_content = await response.aread()
error_msg = error_content.decode("utf-8")
raise Exception(f"API call failed with status code {response.status_code}, {error_msg}")
raise Exception(response.status_code, error_msg)
async for line in response.aiter_lines():
yield line
async def create_embeddings(self, input: str, model: str, api_key: str) -> Dict[str, Any]:
async def create_embeddings(
self, input: str, model: str, api_key: str
) -> Dict[str, Any]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
proxy_to_use = None
if settings.PROXIES:
if settings.PROXIES_USE_CONSISTENCY_HASH_BY_API_KEY:
@@ -332,10 +367,12 @@ class OpenaiApiClient(ApiClient):
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
raise Exception(response.status_code, error_content)
return response.json()
async def generate_images(self, payload: Dict[str, Any], api_key: str) -> Dict[str, Any]:
async def generate_images(
self, payload: Dict[str, Any], api_key: str
) -> Dict[str, Any]:
timeout = httpx.Timeout(self.timeout, read=self.timeout)
proxy_to_use = None
@@ -352,5 +389,5 @@ class OpenaiApiClient(ApiClient):
response = await client.post(url, json=payload, headers=headers)
if response.status_code != 200:
error_content = response.text
raise Exception(f"API call failed with status code {response.status_code}, {error_content}")
return response.json()
raise Exception(response.status_code, error_content)
return response.json()
+1 -6
View File
@@ -1,5 +1,4 @@
import datetime
import re
import time
from typing import List, Union
@@ -56,13 +55,9 @@ class EmbeddingService:
raise e
except Exception as e:
is_success = False
status_code = 500
error_log_msg = f"Generic error: {e}"
logger.error(f"Error creating embedding (Exception): {error_log_msg}")
match = re.search(r"status code (\d+)", str(e))
if match:
status_code = int(match.group(1))
else:
status_code = 500
raise e
finally:
end_time = time.perf_counter()
@@ -1,7 +1,6 @@
# app/service/embedding/gemini_embedding_service.py
import datetime
import re
import time
from typing import Any, Dict
@@ -69,13 +68,9 @@ class GeminiEmbeddingService:
return response
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"Single embedding API call failed: {error_log_msg}")
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=api_key,
@@ -119,13 +114,9 @@ class GeminiEmbeddingService:
return response
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"Batch embedding API call failed: {error_log_msg}")
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=api_key,
@@ -1,6 +1,4 @@
import datetime
import json
import re
import time
from typing import Any, AsyncGenerator, Dict, Union
@@ -80,13 +78,9 @@ class OpenAICompatiableService:
return response
except Exception as e:
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.error(f"Normal API call failed with error: {error_log_msg}")
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=api_key,
@@ -138,15 +132,11 @@ class OpenAICompatiableService:
except Exception as e:
retries += 1
is_success = False
error_log_msg = str(e)
status_code = e.args[0]
error_log_msg = e.args[1]
logger.warning(
f"Streaming API call failed with error: {error_log_msg}. Attempt {retries} of {max_retries}"
)
match = re.search(r"status code (\d+)", error_log_msg)
if match:
status_code = int(match.group(1))
else:
status_code = 500
await add_error_log(
gemini_key=current_attempt_key,
@@ -170,14 +160,14 @@ class OpenAICompatiableService:
logger.error(
f"No valid API key available after {retries} retries."
)
break
raise
else:
logger.error("KeyManager not available for retry logic.")
break
if retries >= max_retries:
logger.error(f"Max retries ({max_retries}) reached for streaming.")
break
raise
finally:
end_time = time.perf_counter()
latency_ms = int((end_time - start_time) * 1000)
@@ -189,6 +179,3 @@ class OpenAICompatiableService:
latency_ms=latency_ms,
request_time=request_datetime,
)
if not is_success and retries >= max_retries:
yield f"data: {json.dumps({'error': 'Streaming failed after retries'})}\n\n"
yield "data: [DONE]\n\n"