mirror of
https://github.com/snailyp/gemini-balance.git
synced 2026-09-06 16:16:37 +08:00
fix:优化智能路由中间件,增强URL处理逻辑
- 增加对新路径模式的支持,包括对`v1beta/models`的处理 - 统一日志记录格式,提升调试信息的可读性 - 规范化代码风格,确保一致性和可维护性 - 修复了请求体和查询参数的模型名称提取逻辑
This commit is contained in:
@@ -15,7 +15,7 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
|
|||||||
async def dispatch(self, request: Request, call_next):
|
async def dispatch(self, request: Request, call_next):
|
||||||
if not settings.URL_NORMALIZATION_ENABLED:
|
if not settings.URL_NORMALIZATION_ENABLED:
|
||||||
return await call_next(request)
|
return await call_next(request)
|
||||||
|
logger.debug(f"request: {request}")
|
||||||
original_path = str(request.url.path)
|
original_path = str(request.url.path)
|
||||||
method = request.method
|
method = request.method
|
||||||
|
|
||||||
@@ -41,20 +41,20 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
|
|||||||
return path, None
|
return path, None
|
||||||
|
|
||||||
# 1. 最高优先级:包含generateContent → Gemini格式
|
# 1. 最高优先级:包含generateContent → Gemini格式
|
||||||
if 'generatecontent' in path.lower():
|
if "generatecontent" in path.lower() or "v1beta/models" in path.lower():
|
||||||
return self.fix_gemini_by_operation(path, method, request)
|
return self.fix_gemini_by_operation(path, method, request)
|
||||||
|
|
||||||
# 2. 第二优先级:包含/openai/ → OpenAI格式
|
# 2. 第二优先级:包含/openai/ → OpenAI格式
|
||||||
if '/openai/' in path.lower():
|
if "/openai/" in path.lower():
|
||||||
return self.fix_openai_by_operation(path, method)
|
return self.fix_openai_by_operation(path, method)
|
||||||
|
|
||||||
# 3. 第三优先级:包含/v1/ → v1格式
|
# 3. 第三优先级:包含/v1/ → v1格式
|
||||||
if '/v1/' in path.lower():
|
if "/v1/" in path.lower():
|
||||||
return self.fix_v1_by_operation(path, method)
|
return self.fix_v1_by_operation(path, method)
|
||||||
|
|
||||||
# 4. 第四优先级:包含/chat/completions → chat功能
|
# 4. 第四优先级:包含/chat/completions → chat功能
|
||||||
if '/chat/completions' in path.lower():
|
if "/chat/completions" in path.lower():
|
||||||
return '/v1/chat/completions', {'type': 'v1_chat'}
|
return "/v1/chat/completions", {"type": "v1_chat"}
|
||||||
|
|
||||||
# 5. 默认:原样传递
|
# 5. 默认:原样传递
|
||||||
return path, None
|
return path, None
|
||||||
@@ -63,16 +63,16 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
|
|||||||
"""检查是否已经是正确的API格式"""
|
"""检查是否已经是正确的API格式"""
|
||||||
# 检查是否已经是正确的端点格式
|
# 检查是否已经是正确的端点格式
|
||||||
correct_patterns = [
|
correct_patterns = [
|
||||||
r'^/v1beta/models/[^/:]+:(generate|streamGenerate)Content$', # Gemini原生
|
r"^/v1beta/models/[^/:]+:(generate|streamGenerate)Content$", # Gemini原生
|
||||||
r'^/gemini/v1beta/models/[^/:]+:(generate|streamGenerate)Content$', # Gemini带前缀
|
r"^/gemini/v1beta/models/[^/:]+:(generate|streamGenerate)Content$", # Gemini带前缀
|
||||||
r'^/v1beta/models$', # Gemini模型列表
|
r"^/v1beta/models$", # Gemini模型列表
|
||||||
r'^/gemini/v1beta/models$', # Gemini带前缀的模型列表
|
r"^/gemini/v1beta/models$", # Gemini带前缀的模型列表
|
||||||
r'^/v1/(chat/completions|models|embeddings|images/generations)$', # v1格式
|
r"^/v1/(chat/completions|models|embeddings|images/generations)$", # v1格式
|
||||||
r'^/openai/v1/(chat/completions|models|embeddings|images/generations)$', # OpenAI格式
|
r"^/openai/v1/(chat/completions|models|embeddings|images/generations)$", # OpenAI格式
|
||||||
r'^/hf/v1/(chat/completions|models|embeddings|images/generations)$', # HF格式
|
r"^/hf/v1/(chat/completions|models|embeddings|images/generations)$", # HF格式
|
||||||
r'^/vertex-express/v1beta/models/[^/:]+:(generate|streamGenerate)Content$', # Vertex Express Gemini格式
|
r"^/vertex-express/v1beta/models/[^/:]+:(generate|streamGenerate)Content$", # Vertex Express Gemini格式
|
||||||
r'^/vertex-express/v1beta/models$', # Vertex Express模型列表
|
r"^/vertex-express/v1beta/models$", # Vertex Express模型列表
|
||||||
r'^/vertex-express/v1/(chat/completions|models|embeddings|images/generations)$', # Vertex Express OpenAI格式
|
r"^/vertex-express/v1/(chat/completions|models|embeddings|images/generations)$", # Vertex Express OpenAI格式
|
||||||
]
|
]
|
||||||
|
|
||||||
for pattern in correct_patterns:
|
for pattern in correct_patterns:
|
||||||
@@ -81,10 +81,14 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def fix_gemini_by_operation(self, path: str, method: str, request: Request) -> tuple:
|
def fix_gemini_by_operation(
|
||||||
|
self, path: str, method: str, request: Request
|
||||||
|
) -> tuple:
|
||||||
"""根据Gemini操作修复,考虑端点偏好"""
|
"""根据Gemini操作修复,考虑端点偏好"""
|
||||||
if method != 'POST':
|
if method == "GET":
|
||||||
return path, None
|
return "/v1beta/models", {
|
||||||
|
"role": "gemini_models",
|
||||||
|
}
|
||||||
|
|
||||||
# 提取模型名称
|
# 提取模型名称
|
||||||
try:
|
try:
|
||||||
@@ -97,72 +101,80 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
|
|||||||
is_stream = self.detect_stream_request(path, request)
|
is_stream = self.detect_stream_request(path, request)
|
||||||
|
|
||||||
# 检查是否有vertex-express偏好
|
# 检查是否有vertex-express偏好
|
||||||
if '/vertex-express/' in path.lower():
|
if "/vertex-express/" in path.lower():
|
||||||
if is_stream:
|
if is_stream:
|
||||||
target_url = f'/vertex-express/v1beta/models/{model_name}:streamGenerateContent'
|
target_url = (
|
||||||
|
f"/vertex-express/v1beta/models/{model_name}:streamGenerateContent"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
target_url = f'/vertex-express/v1beta/models/{model_name}:generateContent'
|
target_url = (
|
||||||
|
f"/vertex-express/v1beta/models/{model_name}:generateContent"
|
||||||
|
)
|
||||||
|
|
||||||
fix_info = {
|
fix_info = {
|
||||||
'rule': 'vertex_express_generate' if not is_stream else 'vertex_express_stream',
|
"rule": (
|
||||||
'preference': 'vertex_express_format',
|
"vertex_express_generate"
|
||||||
'is_stream': is_stream,
|
if not is_stream
|
||||||
'model': model_name
|
else "vertex_express_stream"
|
||||||
|
),
|
||||||
|
"preference": "vertex_express_format",
|
||||||
|
"is_stream": is_stream,
|
||||||
|
"model": model_name,
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
# 标准Gemini端点
|
# 标准Gemini端点
|
||||||
if is_stream:
|
if is_stream:
|
||||||
target_url = f'/v1beta/models/{model_name}:streamGenerateContent'
|
target_url = f"/v1beta/models/{model_name}:streamGenerateContent"
|
||||||
else:
|
else:
|
||||||
target_url = f'/v1beta/models/{model_name}:generateContent'
|
target_url = f"/v1beta/models/{model_name}:generateContent"
|
||||||
|
|
||||||
fix_info = {
|
fix_info = {
|
||||||
'rule': 'gemini_generate' if not is_stream else 'gemini_stream',
|
"rule": "gemini_generate" if not is_stream else "gemini_stream",
|
||||||
'preference': 'gemini_format',
|
"preference": "gemini_format",
|
||||||
'is_stream': is_stream,
|
"is_stream": is_stream,
|
||||||
'model': model_name
|
"model": model_name,
|
||||||
}
|
}
|
||||||
|
|
||||||
return target_url, fix_info
|
return target_url, fix_info
|
||||||
|
|
||||||
def fix_openai_by_operation(self, path: str, method: str) -> tuple:
|
def fix_openai_by_operation(self, path: str, method: str) -> tuple:
|
||||||
"""根据操作类型修复OpenAI格式"""
|
"""根据操作类型修复OpenAI格式"""
|
||||||
if method == 'POST':
|
if method == "POST":
|
||||||
if 'chat' in path.lower() or 'completion' in path.lower():
|
if "chat" in path.lower() or "completion" in path.lower():
|
||||||
return '/openai/v1/chat/completions', {'type': 'openai_chat'}
|
return "/openai/v1/chat/completions", {"type": "openai_chat"}
|
||||||
elif 'embedding' in path.lower():
|
elif "embedding" in path.lower():
|
||||||
return '/openai/v1/embeddings', {'type': 'openai_embeddings'}
|
return "/openai/v1/embeddings", {"type": "openai_embeddings"}
|
||||||
elif 'image' in path.lower():
|
elif "image" in path.lower():
|
||||||
return '/openai/v1/images/generations', {'type': 'openai_images'}
|
return "/openai/v1/images/generations", {"type": "openai_images"}
|
||||||
elif method == 'GET':
|
elif method == "GET":
|
||||||
if 'model' in path.lower():
|
if "model" in path.lower():
|
||||||
return '/openai/v1/models', {'type': 'openai_models'}
|
return "/openai/v1/models", {"type": "openai_models"}
|
||||||
|
|
||||||
return path, None
|
return path, None
|
||||||
|
|
||||||
def fix_v1_by_operation(self, path: str, method: str) -> tuple:
|
def fix_v1_by_operation(self, path: str, method: str) -> tuple:
|
||||||
"""根据操作类型修复v1格式"""
|
"""根据操作类型修复v1格式"""
|
||||||
if method == 'POST':
|
if method == "POST":
|
||||||
if 'chat' in path.lower() or 'completion' in path.lower():
|
if "chat" in path.lower() or "completion" in path.lower():
|
||||||
return '/v1/chat/completions', {'type': 'v1_chat'}
|
return "/v1/chat/completions", {"type": "v1_chat"}
|
||||||
elif 'embedding' in path.lower():
|
elif "embedding" in path.lower():
|
||||||
return '/v1/embeddings', {'type': 'v1_embeddings'}
|
return "/v1/embeddings", {"type": "v1_embeddings"}
|
||||||
elif 'image' in path.lower():
|
elif "image" in path.lower():
|
||||||
return '/v1/images/generations', {'type': 'v1_images'}
|
return "/v1/images/generations", {"type": "v1_images"}
|
||||||
elif method == 'GET':
|
elif method == "GET":
|
||||||
if 'model' in path.lower():
|
if "model" in path.lower():
|
||||||
return '/v1/models', {'type': 'v1_models'}
|
return "/v1/models", {"type": "v1_models"}
|
||||||
|
|
||||||
return path, None
|
return path, None
|
||||||
|
|
||||||
def detect_stream_request(self, path: str, request: Request) -> bool:
|
def detect_stream_request(self, path: str, request: Request) -> bool:
|
||||||
"""检测是否为流式请求"""
|
"""检测是否为流式请求"""
|
||||||
# 1. 路径中包含stream关键词
|
# 1. 路径中包含stream关键词
|
||||||
if 'stream' in path.lower():
|
if "stream" in path.lower():
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# 2. 查询参数
|
# 2. 查询参数
|
||||||
if request.query_params.get('stream') == 'true':
|
if request.query_params.get("stream") == "true":
|
||||||
return True
|
return True
|
||||||
|
|
||||||
return False
|
return False
|
||||||
@@ -171,21 +183,22 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
|
|||||||
"""从请求中提取模型名称,用于构建Gemini API URL"""
|
"""从请求中提取模型名称,用于构建Gemini API URL"""
|
||||||
# 1. 从请求体中提取
|
# 1. 从请求体中提取
|
||||||
try:
|
try:
|
||||||
if hasattr(request, '_body') and request._body:
|
if hasattr(request, "_body") and request._body:
|
||||||
import json
|
import json
|
||||||
|
|
||||||
body = json.loads(request._body.decode())
|
body = json.loads(request._body.decode())
|
||||||
if 'model' in body and body['model']:
|
if "model" in body and body["model"]:
|
||||||
return body['model']
|
return body["model"]
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# 2. 从查询参数中提取
|
# 2. 从查询参数中提取
|
||||||
model_param = request.query_params.get('model')
|
model_param = request.query_params.get("model")
|
||||||
if model_param:
|
if model_param:
|
||||||
return model_param
|
return model_param
|
||||||
|
|
||||||
# 3. 从路径中提取(用于已包含模型名称的路径)
|
# 3. 从路径中提取(用于已包含模型名称的路径)
|
||||||
match = re.search(r'/models/([^/:]+)', path, re.IGNORECASE)
|
match = re.search(r"/models/([^/:]+)", path, re.IGNORECASE)
|
||||||
if match:
|
if match:
|
||||||
return match.group(1)
|
return match.group(1)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user