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