Merge branch 'pr/chinrain/167'

This commit is contained in:
snaily
2025-07-03 17:28:58 +08:00
2 changed files with 128 additions and 169 deletions

View File

@@ -3,7 +3,6 @@ from starlette.middleware.base import BaseHTTPMiddleware
from app.config.config import settings
from app.log.logger import get_main_logger
import re
import json
logger = get_main_logger()
@@ -16,18 +15,18 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next):
if not settings.URL_NORMALIZATION_ENABLED:
return await call_next(request)
logger.debug(f"request: {request}")
original_path = str(request.url.path)
method = request.method
# 尝试修复URL
fixed_path, fix_info = self.fix_request_url(original_path, method, request)
if fixed_path != original_path:
logger.info(f"URL fixed: {method} {original_path}{fixed_path}")
if fix_info:
logger.debug(f"Fix details: {fix_info}")
# 重写请求路径
request.scope["path"] = fixed_path
request.scope["raw_path"] = fixed_path.encode()
@@ -35,47 +34,45 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
return await call_next(request)
def fix_request_url(self, path: str, method: str, request: Request) -> tuple:
"""修复错误的请求URL - 简化版本"""
"""简化的URL修复逻辑"""
# 首先检查是否已经是正确的格式,如果是则不处理
if self.is_already_correct_format(path):
return path, None
# 检测是否为流式请求
is_stream_request = self.detect_stream_request(path, request)
# 1. 最高优先级包含generateContent → Gemini格式
if "generatecontent" in path.lower() or "v1beta/models" in path.lower():
return self.fix_gemini_by_operation(path, method, request)
# 1. 优先检测OpenAI格式请求避免被v1beta误判
if self.is_openai_request(path, request):
return self.fix_openai_request(path, method, request)
# 2. 第二优先级:包含/openai/ → OpenAI格式
if "/openai/" in path.lower():
return self.fix_openai_by_operation(path, method)
# 2. 检测HF格式请求
if self.is_hf_request(path, request):
return self.fix_hf_request(path, method, request)
# 3. 第三优先级:包含/v1/ → v1格式
if "/v1/" in path.lower():
return self.fix_v1_by_operation(path, method)
# 3. 检测Vertex Express格式请求优先级高于Gemini
if self.is_vertex_express_request(path, request):
return self.fix_vertex_express_request(path, method, request, is_stream_request)
# 4. 第四优先级:包含/chat/completions → chat功能
if "/chat/completions" in path.lower():
return "/v1/chat/completions", {"type": "v1_chat"}
# 4. 检测Gemini请求
if self.is_gemini_request(path):
return self.fix_gemini_request(path, method, request, is_stream_request)
# 5. 默认处理其他请求转为最快的v1端点
return self.fix_default_request(path, method, request)
# 5. 默认:原样传递
return path, None
def is_already_correct_format(self, path: str) -> bool:
"""检查是否已经是正确的API格式"""
# 检查是否已经是正确的端点格式
correct_patterns = [
r'^/v1beta/models/[^/:]+:(generate|streamGenerate)Content$', # Gemini原生
r'^/gemini/v1beta/models/[^/:]+:(generate|streamGenerate)Content$', # Gemini带前缀
r'^/v1beta/models$', # Gemini模型列表
r'^/gemini/v1beta/models$', # Gemini带前缀的模型列表
r'^/v1/(chat/completions|models|embeddings|images/generations)$', # v1格式
r'^/openai/v1/(chat/completions|models|embeddings|images/generations)$', # OpenAI格式
r'^/hf/v1/(chat/completions|models|embeddings|images/generations)$', # HF格式
r'^/vertex-express/v1beta/models/[^/:]+:(generate|streamGenerate)Content$', # Vertex Express
r'^/vertex-express/v1beta/models$', # Vertex Express模型列表
r"^/v1beta/models/[^/:]+:(generate|streamGenerate)Content$", # Gemini原生
r"^/gemini/v1beta/models/[^/:]+:(generate|streamGenerate)Content$", # Gemini带前缀
r"^/v1beta/models$", # Gemini模型列表
r"^/gemini/v1beta/models$", # Gemini带前缀的模型列表
r"^/v1/(chat/completions|models|embeddings|images/generations)$", # v1格式
r"^/openai/v1/(chat/completions|models|embeddings|images/generations)$", # OpenAI格式
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$", # Vertex Express模型列表
r"^/vertex-express/v1/(chat/completions|models|embeddings|images/generations)$", # Vertex Express OpenAI格式
]
for pattern in correct_patterns:
@@ -84,54 +81,14 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
return False
def is_openai_request(self, path: str, request: Request) -> bool:
"""检测OpenAI格式请求"""
return '/openai/' in path.lower()
def is_hf_request(self, path: str, request: Request) -> bool:
"""检测HF格式请求"""
return '/hf/' in path.lower()
def is_vertex_express_request(self, path: str, request: Request) -> bool:
"""检测Vertex Express格式请求"""
return '/vertex-express/' in path.lower()
def fix_openai_request(self, path: str, method: str, request: Request) -> tuple:
"""修复OpenAI格式请求"""
if method == 'POST':
if 'chat' in path.lower() or 'completion' in path.lower():
return '/openai/v1/chat/completions', {'type': 'openai_chat'}
elif 'embedding' in path.lower():
return '/openai/v1/embeddings', {'type': 'openai_embeddings'}
elif 'image' in path.lower():
return '/openai/v1/images/generations', {'type': 'openai_images'}
elif method == 'GET':
if 'model' in path.lower():
return '/openai/v1/models', {'type': 'openai_models'}
return path, None
def fix_hf_request(self, path: str, method: str, request: Request) -> tuple:
"""修复HF格式请求"""
if method == 'POST':
if 'chat' in path.lower() or 'completion' in path.lower():
return '/hf/v1/chat/completions', {'type': 'hf_chat'}
elif 'embedding' in path.lower():
return '/hf/v1/embeddings', {'type': 'hf_embeddings'}
elif 'image' in path.lower():
return '/hf/v1/images/generations', {'type': 'hf_images'}
elif method == 'GET':
if 'model' in path.lower():
return '/hf/v1/models', {'type': 'hf_models'}
return path, None
def fix_vertex_express_request(self, path: str, method: str, request: Request, is_stream: bool) -> tuple:
"""修复Vertex Express请求"""
if method != 'POST':
if method == 'GET' and 'models' in path.lower():
return '/vertex-express/v1beta/models', {'rule': 'vertex_express_models', 'preference': 'vertex_express_format'}
return path, None
def fix_gemini_by_operation(
self, path: str, method: str, request: Request
) -> tuple:
"""根据Gemini操作修复考虑端点偏好"""
if method == "GET":
return "/v1beta/models", {
"role": "gemini_models",
}
# 提取模型名称
try:
@@ -140,98 +97,84 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
# 无法提取模型名称,返回原路径不做处理
return path, None
# 构建目标URL
if is_stream:
target_url = f'/vertex-express/v1beta/models/{model_name}:streamGenerateContent'
else:
target_url = f'/vertex-express/v1beta/models/{model_name}:generateContent'
# 检测是否为流式请求
is_stream = self.detect_stream_request(path, request)
fix_info = {
'rule': 'vertex_express_generate' if not is_stream else 'vertex_express_stream',
'preference': 'vertex_express_format',
'is_stream': is_stream,
'model': model_name
}
# 检查是否有vertex-express偏好
if "/vertex-express/" in path.lower():
if is_stream:
target_url = (
f"/vertex-express/v1beta/models/{model_name}:streamGenerateContent"
)
else:
target_url = (
f"/vertex-express/v1beta/models/{model_name}:generateContent"
)
fix_info = {
"rule": (
"vertex_express_generate"
if not is_stream
else "vertex_express_stream"
),
"preference": "vertex_express_format",
"is_stream": is_stream,
"model": model_name,
}
else:
# 标准Gemini端点
if is_stream:
target_url = f"/v1beta/models/{model_name}:streamGenerateContent"
else:
target_url = f"/v1beta/models/{model_name}:generateContent"
fix_info = {
"rule": "gemini_generate" if not is_stream else "gemini_stream",
"preference": "gemini_format",
"is_stream": is_stream,
"model": model_name,
}
return target_url, fix_info
def fix_default_request(self, path: str, method: str, request: Request) -> tuple:
"""修复默认请求转为最快的v1端点"""
if method == 'POST':
if 'chat' in path.lower() or 'completion' in path.lower():
return '/v1/chat/completions', {'type': 'default_chat'}
elif 'embedding' in path.lower():
return '/v1/embeddings', {'type': 'default_embeddings'}
elif 'image' in path.lower():
return '/v1/images/generations', {'type': 'default_images'}
elif method == 'GET':
if 'model' in path.lower():
return '/v1/models', {'type': 'default_models'}
def fix_openai_by_operation(self, path: str, method: str) -> tuple:
"""根据操作类型修复OpenAI格式"""
if method == "POST":
if "chat" in path.lower() or "completion" in path.lower():
return "/openai/v1/chat/completions", {"type": "openai_chat"}
elif "embedding" in path.lower():
return "/openai/v1/embeddings", {"type": "openai_embeddings"}
elif "image" in path.lower():
return "/openai/v1/images/generations", {"type": "openai_images"}
elif method == "GET":
if "model" in path.lower():
return "/openai/v1/models", {"type": "openai_models"}
return path, None
def fix_gemini_request(self, path: str, method: str, request: Request, is_stream: bool) -> tuple:
"""修复Gemini请求"""
if method != 'POST':
if method == 'GET' and 'models' in path.lower():
return '/v1beta/models', {'rule': 'gemini_models', 'preference': 'gemini_format'}
return path, None
def fix_v1_by_operation(self, path: str, method: str) -> tuple:
"""根据操作类型修复v1格式"""
if method == "POST":
if "chat" in path.lower() or "completion" in path.lower():
return "/v1/chat/completions", {"type": "v1_chat"}
elif "embedding" in path.lower():
return "/v1/embeddings", {"type": "v1_embeddings"}
elif "image" in path.lower():
return "/v1/images/generations", {"type": "v1_images"}
elif method == "GET":
if "model" in path.lower():
return "/v1/models", {"type": "v1_models"}
# 提取模型名称
try:
model_name = self.extract_model_name(path, request)
except ValueError:
# 无法提取模型名称,返回原路径不做处理
return path, None
# 构建目标URL
if is_stream:
target_url = f'/v1beta/models/{model_name}:streamGenerateContent'
else:
target_url = f'/v1beta/models/{model_name}:generateContent'
fix_info = {
'rule': 'gemini_generate' if not is_stream else 'gemini_stream',
'preference': 'gemini_format',
'is_stream': is_stream,
'model': model_name
}
return target_url, fix_info
return path, None
def detect_stream_request(self, path: str, request: Request) -> bool:
"""检测是否为流式请求"""
# 1. 路径中包含stream关键词
if 'stream' in path.lower():
if "stream" in path.lower():
return True
# 2. 查询参数
if request.query_params.get('stream') == 'true':
return True
return False
def is_gemini_request(self, path: str) -> bool:
"""判断是否为Gemini API请求"""
path_lower = path.lower()
# 如果已经是OpenAI、HF或Vertex Express格式不应该被识别为Gemini
if '/openai/' in path_lower or '/hf/' in path_lower or '/vertex-express/' in path_lower:
return False
# 1. 检查是否是明确的Gemini路径模式
gemini_path_patterns = [
r'/v1beta/models', # Gemini原生API路径
r'/gemini/v1beta', # 带gemini前缀的路径
]
for pattern in gemini_path_patterns:
if re.search(pattern, path_lower):
return True
# 2. 检查是否包含Gemini模型名称
if 'gemini' in path_lower and ('models' in path_lower or 'generatecontent' in path_lower):
if request.query_params.get("stream") == "true":
return True
return False
@@ -240,24 +183,24 @@ class SmartRoutingMiddleware(BaseHTTPMiddleware):
"""从请求中提取模型名称用于构建Gemini API URL"""
# 1. 从请求体中提取
try:
if hasattr(request, '_body') and request._body:
if hasattr(request, "_body") and request._body:
import json
body = json.loads(request._body.decode())
if 'model' in body and body['model']:
return body['model']
if "model" in body and body["model"]:
return body["model"]
except Exception:
pass
# 2. 从查询参数中提取
model_param = request.query_params.get('model')
model_param = request.query_params.get("model")
if model_param:
return model_param
# 3. 从路径中提取(用于已包含模型名称的路径)
match = re.search(r'/models/([^/:]+)', path, re.IGNORECASE)
match = re.search(r"/models/([^/:]+)", path, re.IGNORECASE)
if match:
return match.group(1)
# 4. 如果无法提取模型名称,抛出异常
raise ValueError("Unable to extract model name from request")

View File

@@ -937,13 +937,29 @@ endblock %} {% block head_extra_styles %}
</div>
<!-- 智能路由配置 -->
<div class="mb-6">
<label class="flex items-center">
<input type="checkbox" id="URL_NORMALIZATION_ENABLED" name="URL_NORMALIZATION_ENABLED"
class="mr-2 rounded" />
<span class="font-semibold text-gray-700">启用智能路由映射</span>
</label>
<div class="flex items-center justify-between">
<label
for="URL_NORMALIZATION_ENABLED"
class="font-semibold text-gray-700"
>启用智能路由映射</label
>
<div
class="relative inline-block w-10 mr-2 align-middle select-none transition duration-200 ease-in"
>
<input
type="checkbox"
name="URL_NORMALIZATION_ENABLED"
id="URL_NORMALIZATION_ENABLED"
class="toggle-checkbox absolute block w-6 h-6 rounded-full bg-white border-4 appearance-none cursor-pointer"
/>
<label
for="URL_NORMALIZATION_ENABLED"
class="toggle-label block overflow-hidden h-6 rounded-full bg-gray-300 cursor-pointer"
></label>
</div>
</div>
<small class="text-gray-500 mt-1 block">
自动客户端的各种URL格式映射到正确的API端点
自动客户端请求的url拼接为正确格式仅保证正常聊天出现问题请关闭
</small>
</div>
<!-- 最大失败次数 -->