diff --git a/app/exception/exceptions.py b/app/exception/exceptions.py index 0e9fb30..b951c03 100644 --- a/app/exception/exceptions.py +++ b/app/exception/exceptions.py @@ -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], } }, ) diff --git a/app/router/gemini_routes.py b/app/router/gemini_routes.py index b9945a7..3d126ca 100644 --- a/app/router/gemini_routes.py +++ b/app/router/gemini_routes.py @@ -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 - }) \ No newline at end of file + return JSONResponse( + { + "success": True, + "message": message, + "successful_keys": successful_keys, + "failed_keys": {}, + "valid_count": valid_count, + "invalid_count": 0, + } + ) diff --git a/app/router/openai_compatiable_routes.py b/app/router/openai_compatiable_routes.py index 958c8ea..5a7b38d 100644 --- a/app/router/openai_compatiable_routes.py +++ b/app/router/openai_compatiable_routes.py @@ -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 diff --git a/app/router/openai_routes.py b/app/router/openai_routes.py index 2c590bb..537e1f1 100644 --- a/app/router/openai_routes.py +++ b/app/router/openai_routes.py @@ -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") diff --git a/app/router/vertex_express_routes.py b/app/router/vertex_express_routes.py index 1809300..75bc455 100644 --- a/app/router/vertex_express_routes.py +++ b/app/router/vertex_express_routes.py @@ -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") \ No newline at end of file + 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") diff --git a/app/service/chat/gemini_chat_service.py b/app/service/chat/gemini_chat_service.py index 07551b7..45d0683 100644 --- a/app/service/chat/gemini_chat_service.py +++ b/app/service/chat/gemini_chat_service.py @@ -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) \ No newline at end of file diff --git a/app/service/chat/openai_chat_service.py b/app/service/chat/openai_chat_service.py index bb13dd9..1ab55d0 100644 --- a/app/service/chat/openai_chat_service.py +++ b/app/service/chat/openai_chat_service.py @@ -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, - ) \ No newline at end of file + ) diff --git a/app/service/chat/vertex_express_chat_service.py b/app/service/chat/vertex_express_chat_service.py index 362b10e..7a59c7d 100644 --- a/app/service/chat/vertex_express_chat_service.py +++ b/app/service/chat/vertex_express_chat_service.py @@ -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, - ) \ No newline at end of file + ) diff --git a/app/service/client/api_client.py b/app/service/client/api_client.py index 56d2a00..440949a 100644 --- a/app/service/client/api_client.py +++ b/app/service/client/api_client.py @@ -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() \ No newline at end of file + raise Exception(response.status_code, error_content) + return response.json() diff --git a/app/service/embedding/embedding_service.py b/app/service/embedding/embedding_service.py index 6954a5e..a8e7fc6 100644 --- a/app/service/embedding/embedding_service.py +++ b/app/service/embedding/embedding_service.py @@ -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() diff --git a/app/service/embedding/gemini_embedding_service.py b/app/service/embedding/gemini_embedding_service.py index 62c697e..4a0819a 100644 --- a/app/service/embedding/gemini_embedding_service.py +++ b/app/service/embedding/gemini_embedding_service.py @@ -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, diff --git a/app/service/openai_compatiable/openai_compatiable_service.py b/app/service/openai_compatiable/openai_compatiable_service.py index 6c09003..53af167 100644 --- a/app/service/openai_compatiable/openai_compatiable_service.py +++ b/app/service/openai_compatiable/openai_compatiable_service.py @@ -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"