fix:优化智能路由中间件,增强URL处理逻辑

- 增加对新路径模式的支持,包括对`v1beta/models`的处理
- 统一日志记录格式,提升调试信息的可读性
- 规范化代码风格,确保一致性和可维护性
- 修复了请求体和查询参数的模型名称提取逻辑
This commit is contained in:
snaily
2025-07-03 17:25:50 +08:00
parent 94d1041961
commit f79a52f839
+76 -63
View File
@@ -15,18 +15,18 @@ 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
# 尝试修复URL # 尝试修复URL
fixed_path, fix_info = self.fix_request_url(original_path, method, request) fixed_path, fix_info = self.fix_request_url(original_path, method, request)
if fixed_path != original_path: if fixed_path != original_path:
logger.info(f"URL fixed: {method} {original_path}{fixed_path}") logger.info(f"URL fixed: {method} {original_path}{fixed_path}")
if fix_info: if fix_info:
logger.debug(f"Fix details: {fix_info}") logger.debug(f"Fix details: {fix_info}")
# 重写请求路径 # 重写请求路径
request.scope["path"] = fixed_path request.scope["path"] = fixed_path
request.scope["raw_path"] = fixed_path.encode() request.scope["raw_path"] = fixed_path.encode()
@@ -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,23 +183,24 @@ 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)
# 4. 如果无法提取模型名称,抛出异常 # 4. 如果无法提取模型名称,抛出异常
raise ValueError("Unable to extract model name from request") raise ValueError("Unable to extract model name from request")