mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
feat: add Gemini Files API proxy with key capability support
- Add gemini_files_api capability for key-level Files API support - Implement Gemini Files API proxy endpoints (upload/list/get/delete) - Add file-to-key mapping cache for correct key routing - Support preferred_key_ids to prioritize keys that uploaded files - Enhance logging for file mapping diagnostics
This commit is contained in:
@@ -439,6 +439,14 @@ class BaseMessageHandler:
|
|||||||
adapter_detector=self.adapter_detector,
|
adapter_detector=self.adapter_detector,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def _resolve_preferred_key_ids(
|
||||||
|
self,
|
||||||
|
model_name: str,
|
||||||
|
request_body: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> Optional[list[str]]:
|
||||||
|
"""可选的 Key 优先级解析钩子(默认不启用)。"""
|
||||||
|
return None
|
||||||
|
|
||||||
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
|
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
|
||||||
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
|
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
|
||||||
if provider_type:
|
if provider_type:
|
||||||
|
|||||||
@@ -511,6 +511,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
capability_requirements = self._resolve_capability_requirements(
|
capability_requirements = self._resolve_capability_requirements(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
)
|
||||||
|
preferred_key_ids = await self._resolve_preferred_key_ids(
|
||||||
|
model_name=model,
|
||||||
|
request_body=original_request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 执行请求(通过 FallbackOrchestrator)
|
# 执行请求(通过 FallbackOrchestrator)
|
||||||
@@ -529,6 +534,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
preferred_key_ids=preferred_key_ids or None,
|
||||||
request_body_ref=request_body_ref, # 传递容器引用
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1126,6 +1132,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
capability_requirements = self._resolve_capability_requirements(
|
capability_requirements = self._resolve_capability_requirements(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
)
|
||||||
|
preferred_key_ids = await self._resolve_preferred_key_ids(
|
||||||
|
model_name=model,
|
||||||
|
request_body=original_request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
(
|
(
|
||||||
@@ -1142,6 +1153,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
request_func=sync_request_func,
|
request_func=sync_request_func,
|
||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
preferred_key_ids=preferred_key_ids or None,
|
||||||
request_body_ref=request_body_ref, # 传递容器引用
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -568,6 +568,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
capability_requirements = self._resolve_capability_requirements(
|
capability_requirements = self._resolve_capability_requirements(
|
||||||
model_name=ctx.model,
|
model_name=ctx.model,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
)
|
||||||
|
preferred_key_ids = await self._resolve_preferred_key_ids(
|
||||||
|
model_name=ctx.model,
|
||||||
|
request_body=original_request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 执行请求(通过 FallbackOrchestrator)
|
# 执行请求(通过 FallbackOrchestrator)
|
||||||
@@ -586,6 +591,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
preferred_key_ids=preferred_key_ids or None,
|
||||||
request_body_ref=request_body_ref, # 传递容器引用
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2324,6 +2330,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
capability_requirements = self._resolve_capability_requirements(
|
capability_requirements = self._resolve_capability_requirements(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
)
|
||||||
|
preferred_key_ids = await self._resolve_preferred_key_ids(
|
||||||
|
model_name=model,
|
||||||
|
request_body=original_request_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
(
|
(
|
||||||
@@ -2340,6 +2351,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
request_func=sync_request_func,
|
request_func=sync_request_func,
|
||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
capability_requirements=capability_requirements or None,
|
capability_requirements=capability_requirements or None,
|
||||||
|
preferred_key_ids=preferred_key_ids or None,
|
||||||
request_body_ref=request_body_ref, # 传递容器引用
|
request_body_ref=request_body_ref, # 传递容器引用
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
|||||||
from src.core.api_format import extract_client_api_key_with_query
|
from src.core.api_format import extract_client_api_key_with_query
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.gemini import GeminiRequest
|
from src.models.gemini import GeminiRequest
|
||||||
|
from src.services.gemini_files_mapping import extract_file_names_from_request
|
||||||
from src.services.provider.transport import redact_url_for_log
|
from src.services.provider.transport import redact_url_for_log
|
||||||
|
|
||||||
|
|
||||||
@@ -56,6 +57,16 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
self._get_api_format(),
|
self._get_api_format(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def detect_capability_requirements(
|
||||||
|
self,
|
||||||
|
headers: Dict[str, str], # noqa: ARG002 - 预留
|
||||||
|
request_body: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> Dict[str, bool]:
|
||||||
|
"""检测是否需要 Gemini Files API 能力"""
|
||||||
|
if request_body and extract_file_names_from_request(request_body):
|
||||||
|
return {"gemini_files_api": True}
|
||||||
|
return {}
|
||||||
|
|
||||||
def _merge_path_params(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
|
|||||||
@@ -22,6 +22,58 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
FORMAT_ID = "GEMINI"
|
FORMAT_ID = "GEMINI"
|
||||||
|
|
||||||
|
async def _resolve_preferred_key_ids(
|
||||||
|
self,
|
||||||
|
model_name: str, # noqa: ARG002 - 仅做文件绑定
|
||||||
|
request_body: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> Optional[list[str]]:
|
||||||
|
"""
|
||||||
|
从 files/xxx 绑定关系中解析优先 Key ID 列表。
|
||||||
|
|
||||||
|
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
|
||||||
|
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
|
||||||
|
|
||||||
|
注意事项:
|
||||||
|
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
|
||||||
|
- 如果多个文件属于不同 Key,只能使用其中一个,其他文件可能无法访问
|
||||||
|
"""
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.services.gemini_files_mapping import (
|
||||||
|
extract_file_names_from_request,
|
||||||
|
get_file_key_mapping,
|
||||||
|
)
|
||||||
|
|
||||||
|
file_names = extract_file_names_from_request(request_body or {})
|
||||||
|
if not file_names:
|
||||||
|
return None
|
||||||
|
|
||||||
|
preferred_key_ids: list[str] = []
|
||||||
|
unmapped_files: list[str] = [] # 记录找不到映射的文件
|
||||||
|
|
||||||
|
for file_name in file_names:
|
||||||
|
key_id = await get_file_key_mapping(file_name)
|
||||||
|
if key_id:
|
||||||
|
if key_id not in preferred_key_ids:
|
||||||
|
preferred_key_ids.append(key_id)
|
||||||
|
else:
|
||||||
|
unmapped_files.append(file_name)
|
||||||
|
|
||||||
|
# 警告:映射缺失
|
||||||
|
if unmapped_files:
|
||||||
|
logger.warning(
|
||||||
|
f"[{self.request_id}] Gemini 文件→Key 映射缺失: {unmapped_files},"
|
||||||
|
"请求可能失败(文件属于其他 Key 或映射已过期)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 警告:多个文件属于不同 Key
|
||||||
|
if len(preferred_key_ids) > 1:
|
||||||
|
logger.warning(
|
||||||
|
f"[{self.request_id}] 请求使用了多个文件,但它们属于不同的 Key: "
|
||||||
|
f"{preferred_key_ids},只能使用第一个 Key,其他文件可能无法访问"
|
||||||
|
)
|
||||||
|
|
||||||
|
return preferred_key_ids or None
|
||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: Dict[str, Any],
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from .capabilities import router as capabilities_router
|
|||||||
from .catalog import router as catalog_router
|
from .catalog import router as catalog_router
|
||||||
from .claude import router as claude_router
|
from .claude import router as claude_router
|
||||||
from .gemini import router as gemini_router
|
from .gemini import router as gemini_router
|
||||||
|
from .gemini_files import router as gemini_files_router
|
||||||
from .models import router as models_router
|
from .models import router as models_router
|
||||||
from .modules import router as modules_router
|
from .modules import router as modules_router
|
||||||
from .openai import router as openai_router
|
from .openai import router as openai_router
|
||||||
@@ -17,6 +18,7 @@ router.include_router(models_router)
|
|||||||
router.include_router(claude_router, tags=["Claude API"])
|
router.include_router(claude_router, tags=["Claude API"])
|
||||||
router.include_router(openai_router)
|
router.include_router(openai_router)
|
||||||
router.include_router(gemini_router, tags=["Gemini API"])
|
router.include_router(gemini_router, tags=["Gemini API"])
|
||||||
|
router.include_router(gemini_files_router, tags=["Gemini Files API"])
|
||||||
router.include_router(system_catalog_router, tags=["System Catalog"])
|
router.include_router(system_catalog_router, tags=["System Catalog"])
|
||||||
router.include_router(catalog_router)
|
router.include_router(catalog_router)
|
||||||
router.include_router(capabilities_router)
|
router.include_router(capabilities_router)
|
||||||
|
|||||||
621
src/api/public/gemini_files.py
Normal file
621
src/api/public/gemini_files.py
Normal file
@@ -0,0 +1,621 @@
|
|||||||
|
"""
|
||||||
|
Gemini Files API 代理端点
|
||||||
|
|
||||||
|
代理 Google Gemini Files API,支持文件的上传、查询、删除等操作。
|
||||||
|
|
||||||
|
端点列表:
|
||||||
|
- POST /upload/v1beta/files - 上传文件(可恢复上传)
|
||||||
|
- GET /v1beta/files - 列出文件
|
||||||
|
- GET /v1beta/files/{name} - 获取文件元数据
|
||||||
|
- DELETE /v1beta/files/{name} - 删除文件
|
||||||
|
|
||||||
|
认证方式:
|
||||||
|
- x-goog-api-key 请求头
|
||||||
|
- ?key= URL 参数
|
||||||
|
|
||||||
|
参考文档:
|
||||||
|
https://ai.google.dev/api/files
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Dict, Optional, Tuple
|
||||||
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from src.clients.http_client import HTTPClientPool
|
||||||
|
from src.core.api_format import APIFormat, extract_client_api_key_with_query
|
||||||
|
from src.core.api_format.metadata import get_api_format_definition
|
||||||
|
from src.core.crypto import crypto_service
|
||||||
|
from src.core.logger import logger
|
||||||
|
from src.database import get_db
|
||||||
|
from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User
|
||||||
|
from src.services.auth.service import AuthService
|
||||||
|
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||||
|
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
|
||||||
|
from src.services.provider.transport import redact_url_for_log
|
||||||
|
|
||||||
|
# 从配置获取路径前缀
|
||||||
|
_gemini_def = get_api_format_definition(APIFormat.GEMINI)
|
||||||
|
|
||||||
|
router = APIRouter(tags=["Gemini Files API"], prefix=_gemini_def.path_prefix)
|
||||||
|
|
||||||
|
# Gemini Files API 基础 URL
|
||||||
|
GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
|
||||||
|
|
||||||
|
# Gemini Files API 能力标签
|
||||||
|
REQUIRED_CAPABILITIES = {"gemini_files_api": True}
|
||||||
|
|
||||||
|
# 需要从客户端请求中移除的头部(这些会由代理重新设置或不应转发)
|
||||||
|
HEADERS_TO_REMOVE = frozenset({
|
||||||
|
"host",
|
||||||
|
"content-length",
|
||||||
|
"transfer-encoding",
|
||||||
|
"connection",
|
||||||
|
"x-goog-api-key",
|
||||||
|
"authorization",
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_gemini_api_key(request: Request) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
从请求中提取 Gemini API Key
|
||||||
|
|
||||||
|
优先级(与 Google SDK 行为一致):
|
||||||
|
1. URL 参数 ?key=
|
||||||
|
2. x-goog-api-key 请求头
|
||||||
|
"""
|
||||||
|
return extract_client_api_key_with_query(
|
||||||
|
dict(request.headers),
|
||||||
|
dict(request.query_params),
|
||||||
|
APIFormat.GEMINI,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_upstream_headers(
|
||||||
|
original_headers: Dict[str, str],
|
||||||
|
upstream_api_key: str,
|
||||||
|
) -> Dict[str, str]:
|
||||||
|
"""
|
||||||
|
构建上游请求头
|
||||||
|
|
||||||
|
Args:
|
||||||
|
original_headers: 原始请求头
|
||||||
|
upstream_api_key: 上游 API Key
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
处理后的请求头字典
|
||||||
|
"""
|
||||||
|
headers = {}
|
||||||
|
|
||||||
|
# 透传非敏感头部
|
||||||
|
for name, value in original_headers.items():
|
||||||
|
if name.lower() not in HEADERS_TO_REMOVE:
|
||||||
|
headers[name] = value
|
||||||
|
|
||||||
|
# 设置认证头
|
||||||
|
headers["x-goog-api-key"] = upstream_api_key
|
||||||
|
|
||||||
|
return headers
|
||||||
|
|
||||||
|
|
||||||
|
def _build_upstream_url(
|
||||||
|
base_url: str,
|
||||||
|
path: str,
|
||||||
|
query_params: Optional[Dict[str, Any]] = None,
|
||||||
|
is_upload: bool = False,
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
构建上游 URL
|
||||||
|
|
||||||
|
Args:
|
||||||
|
base_url: 上游基础 URL
|
||||||
|
path: API 路径
|
||||||
|
query_params: 查询参数
|
||||||
|
is_upload: 是否为上传端点
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
完整的上游 URL
|
||||||
|
"""
|
||||||
|
# 移除 key 参数(认证通过 header)
|
||||||
|
effective_params = dict(query_params) if query_params else {}
|
||||||
|
effective_params.pop("key", None)
|
||||||
|
|
||||||
|
# 处理 base_url 可能包含 /v1beta 的情况,避免重复
|
||||||
|
normalized_base_url = base_url.rstrip("/")
|
||||||
|
if normalized_base_url.endswith("/v1beta"):
|
||||||
|
normalized_base_url = normalized_base_url[:-len("/v1beta")]
|
||||||
|
|
||||||
|
# 上传端点使用不同的路径前缀
|
||||||
|
if is_upload:
|
||||||
|
url = f"{normalized_base_url}/upload{path}"
|
||||||
|
else:
|
||||||
|
url = f"{normalized_base_url}{path}"
|
||||||
|
|
||||||
|
if effective_params:
|
||||||
|
query_string = urlencode(effective_params, doseq=True)
|
||||||
|
url = f"{url}?{query_string}"
|
||||||
|
|
||||||
|
return url
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_files_model_name(
|
||||||
|
db: Session,
|
||||||
|
user_api_key: ApiKey,
|
||||||
|
user: Optional[User],
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
为 Files API 选择一个可用的模型名(用于 Key 选择与权限过滤)
|
||||||
|
|
||||||
|
选择顺序:
|
||||||
|
1. 用户/Key 的 allowed_models(取交集后选第一个)
|
||||||
|
2. 任意支持 Gemini 格式的 GlobalModel
|
||||||
|
"""
|
||||||
|
from src.core.model_permissions import merge_allowed_models
|
||||||
|
|
||||||
|
allowed_models = merge_allowed_models(
|
||||||
|
user_api_key.allowed_models,
|
||||||
|
user.allowed_models if user else None,
|
||||||
|
)
|
||||||
|
if allowed_models is not None:
|
||||||
|
if not allowed_models:
|
||||||
|
return None
|
||||||
|
return sorted(allowed_models)[0]
|
||||||
|
|
||||||
|
row = (
|
||||||
|
db.query(GlobalModel.name)
|
||||||
|
.join(Model, Model.global_model_id == GlobalModel.id)
|
||||||
|
.join(Provider, Provider.id == Model.provider_id)
|
||||||
|
.join(ProviderEndpoint, ProviderEndpoint.provider_id == Provider.id)
|
||||||
|
.filter(
|
||||||
|
GlobalModel.is_active == True,
|
||||||
|
Model.is_active == True,
|
||||||
|
Provider.is_active == True,
|
||||||
|
ProviderEndpoint.is_active == True,
|
||||||
|
ProviderEndpoint.api_format == APIFormat.GEMINI.value,
|
||||||
|
)
|
||||||
|
.distinct()
|
||||||
|
.order_by(GlobalModel.name.asc())
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
return row[0] if row else None
|
||||||
|
|
||||||
|
|
||||||
|
async def _select_provider_candidate(
|
||||||
|
db: Session,
|
||||||
|
user_api_key: ApiKey,
|
||||||
|
model_name: str,
|
||||||
|
) -> Optional[ProviderCandidate]:
|
||||||
|
"""选择支持 Files API 的 Provider/Endpoint/Key 组合"""
|
||||||
|
scheduler = CacheAwareScheduler()
|
||||||
|
candidates, _global_model_id = await scheduler.list_all_candidates(
|
||||||
|
db=db,
|
||||||
|
api_format=APIFormat.GEMINI,
|
||||||
|
model_name=model_name,
|
||||||
|
affinity_key=str(user_api_key.id),
|
||||||
|
user_api_key=user_api_key,
|
||||||
|
max_candidates=10,
|
||||||
|
capability_requirements=REQUIRED_CAPABILITIES,
|
||||||
|
)
|
||||||
|
for candidate in candidates:
|
||||||
|
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
|
||||||
|
if auth_type == "api_key":
|
||||||
|
return candidate
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def _resolve_upstream_context(
|
||||||
|
request: Request,
|
||||||
|
db: Session,
|
||||||
|
) -> Tuple[str, str, str]:
|
||||||
|
"""
|
||||||
|
解析上游 Key 与 Base URL
|
||||||
|
|
||||||
|
仅允许系统 API Key,通过能力标签选择支持 Files API 的 Provider Key。
|
||||||
|
"""
|
||||||
|
client_key = _extract_gemini_api_key(request)
|
||||||
|
if not client_key:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=401,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 401,
|
||||||
|
"message": "API key required. Provide via x-goog-api-key header or ?key= parameter.",
|
||||||
|
"status": "UNAUTHENTICATED",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
auth_result = AuthService.authenticate_api_key(db, client_key)
|
||||||
|
if not auth_result:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=401,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 401,
|
||||||
|
"message": "API key not valid. Please pass a valid API key.",
|
||||||
|
"status": "UNAUTHENTICATED",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
user, user_api_key = auth_result
|
||||||
|
model_name = _resolve_files_model_name(db, user_api_key, user)
|
||||||
|
if not model_name:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=503,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 503,
|
||||||
|
"message": "No available model for Gemini Files API routing",
|
||||||
|
"status": "UNAVAILABLE",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
candidate = await _select_provider_candidate(db, user_api_key, model_name)
|
||||||
|
if not candidate:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=503,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 503,
|
||||||
|
"message": "No available key with gemini_files_api capability",
|
||||||
|
"status": "UNAVAILABLE",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
upstream_key = crypto_service.decrypt(candidate.key.api_key)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Failed to decrypt provider key for Gemini Files API: {exc}")
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=500,
|
||||||
|
detail={
|
||||||
|
"error": {
|
||||||
|
"code": 500,
|
||||||
|
"message": "Failed to decrypt provider key",
|
||||||
|
"status": "INTERNAL",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
base_url = candidate.endpoint.base_url or GEMINI_FILES_BASE_URL
|
||||||
|
return upstream_key, base_url, str(candidate.key.id)
|
||||||
|
|
||||||
|
|
||||||
|
async def _proxy_request(
|
||||||
|
method: str,
|
||||||
|
upstream_url: str,
|
||||||
|
headers: Dict[str, str],
|
||||||
|
content: Optional[bytes] = None,
|
||||||
|
json_body: Optional[Dict[str, Any]] = None,
|
||||||
|
file_key_id: Optional[str] = None,
|
||||||
|
) -> Response:
|
||||||
|
"""
|
||||||
|
代理请求到上游 Gemini API
|
||||||
|
|
||||||
|
Args:
|
||||||
|
method: HTTP 方法
|
||||||
|
upstream_url: 上游 URL
|
||||||
|
headers: 请求头
|
||||||
|
content: 原始请求体(二进制)
|
||||||
|
json_body: JSON 请求体
|
||||||
|
file_key_id: 上游 Provider Key ID,用于成功响应时存储 file→key 映射
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
FastAPI Response 对象
|
||||||
|
"""
|
||||||
|
client = await HTTPClientPool.get_default_client_async()
|
||||||
|
|
||||||
|
try:
|
||||||
|
if method.upper() == "GET":
|
||||||
|
response = await client.get(upstream_url, headers=headers)
|
||||||
|
elif method.upper() == "DELETE":
|
||||||
|
response = await client.delete(upstream_url, headers=headers)
|
||||||
|
elif method.upper() == "POST":
|
||||||
|
if content is not None:
|
||||||
|
response = await client.post(
|
||||||
|
upstream_url, headers=headers, content=content
|
||||||
|
)
|
||||||
|
elif json_body is not None:
|
||||||
|
response = await client.post(
|
||||||
|
upstream_url, headers=headers, json=json_body
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
response = await client.post(upstream_url, headers=headers)
|
||||||
|
else:
|
||||||
|
raise HTTPException(status_code=405, detail="Method not allowed")
|
||||||
|
|
||||||
|
# 构建响应头(排除 hop-by-hop 头部)
|
||||||
|
response_headers = {}
|
||||||
|
hop_by_hop = {"connection", "keep-alive", "transfer-encoding", "upgrade"}
|
||||||
|
for name, value in response.headers.items():
|
||||||
|
if name.lower() not in hop_by_hop:
|
||||||
|
response_headers[name] = value
|
||||||
|
|
||||||
|
if (
|
||||||
|
file_key_id
|
||||||
|
and response.status_code < 300
|
||||||
|
and response.headers.get("content-type", "").startswith("application/json")
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
payload = response.json()
|
||||||
|
file_name = None
|
||||||
|
if isinstance(payload, dict):
|
||||||
|
file_name = payload.get("name")
|
||||||
|
if not file_name and isinstance(payload.get("file"), dict):
|
||||||
|
file_name = payload["file"].get("name")
|
||||||
|
if file_name:
|
||||||
|
await store_file_key_mapping(file_name, file_key_id)
|
||||||
|
logger.debug(
|
||||||
|
f"Gemini file→key 映射已存储: {file_name} → key_id={file_key_id}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 为 list_files 响应中的所有文件建立映射
|
||||||
|
# 这是正确的:Gemini API 按 Key 隔离文件,返回的文件必然属于当前 Key
|
||||||
|
files_list = payload.get("files")
|
||||||
|
if isinstance(files_list, list):
|
||||||
|
mapped_count = 0
|
||||||
|
for item in files_list:
|
||||||
|
if isinstance(item, dict) and item.get("name"):
|
||||||
|
await store_file_key_mapping(item["name"], file_key_id)
|
||||||
|
mapped_count += 1
|
||||||
|
if mapped_count > 0:
|
||||||
|
logger.debug(
|
||||||
|
f"Gemini list_files 批量映射已存储: {mapped_count} 个文件 → key_id={file_key_id}"
|
||||||
|
)
|
||||||
|
except (ValueError, KeyError) as e:
|
||||||
|
logger.debug(f"Failed to store Gemini file mapping: {e}")
|
||||||
|
|
||||||
|
return Response(
|
||||||
|
content=response.content,
|
||||||
|
status_code=response.status_code,
|
||||||
|
headers=response_headers,
|
||||||
|
media_type=response.headers.get("content-type", "application/json"),
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
sanitized_error = redact_url_for_log(str(e))
|
||||||
|
logger.error(f"Gemini Files API proxy error: {sanitized_error}")
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=502,
|
||||||
|
content={
|
||||||
|
"error": {
|
||||||
|
"code": 502,
|
||||||
|
"message": "Upstream request failed",
|
||||||
|
"status": "BAD_GATEWAY",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# 文件上传端点
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/upload/v1beta/files")
|
||||||
|
async def upload_file(
|
||||||
|
request: Request,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
上传文件到 Gemini Files API
|
||||||
|
|
||||||
|
支持可恢复上传协议(Resumable Upload Protocol):
|
||||||
|
1. 初始请求:设置元数据,获取上传 URL
|
||||||
|
2. 上传请求:上传实际文件内容
|
||||||
|
|
||||||
|
**认证方式**:
|
||||||
|
- `x-goog-api-key` 请求头,或
|
||||||
|
- `?key=` URL 参数
|
||||||
|
|
||||||
|
**请求头(可恢复上传)**:
|
||||||
|
- `X-Goog-Upload-Protocol: resumable`
|
||||||
|
- `X-Goog-Upload-Command: start` | `upload, finalize`
|
||||||
|
- `X-Goog-Upload-Header-Content-Length`: 文件大小
|
||||||
|
- `X-Goog-Upload-Header-Content-Type`: 文件 MIME 类型
|
||||||
|
|
||||||
|
**请求体(初始请求)**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"file": {
|
||||||
|
"display_name": "文件名"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
|
||||||
|
|
||||||
|
# 读取请求体
|
||||||
|
body = await request.body()
|
||||||
|
|
||||||
|
# 构建上游请求
|
||||||
|
upstream_url = _build_upstream_url(
|
||||||
|
base_url,
|
||||||
|
"/v1beta/files",
|
||||||
|
dict(request.query_params),
|
||||||
|
is_upload=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
||||||
|
|
||||||
|
logger.debug(f"Gemini Files upload proxy: POST {redact_url_for_log(upstream_url)}")
|
||||||
|
|
||||||
|
return await _proxy_request(
|
||||||
|
"POST", upstream_url, headers, content=body, file_key_id=file_key_id
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# 文件列表端点
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/v1beta/files")
|
||||||
|
async def list_files(
|
||||||
|
request: Request,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
pageSize: Optional[int] = None,
|
||||||
|
pageToken: Optional[str] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
列出已上传的文件
|
||||||
|
|
||||||
|
**认证方式**:
|
||||||
|
- `x-goog-api-key` 请求头,或
|
||||||
|
- `?key=` URL 参数
|
||||||
|
|
||||||
|
**查询参数**:
|
||||||
|
- `pageSize`: 每页返回的文件数量(默认 10,最大 100)
|
||||||
|
- `pageToken`: 分页令牌
|
||||||
|
|
||||||
|
**响应格式**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"files": [
|
||||||
|
{
|
||||||
|
"name": "files/abc-123",
|
||||||
|
"displayName": "文件名",
|
||||||
|
"mimeType": "image/jpeg",
|
||||||
|
"sizeBytes": "12345",
|
||||||
|
"createTime": "2024-01-01T00:00:00Z",
|
||||||
|
"updateTime": "2024-01-01T00:00:00Z",
|
||||||
|
"expirationTime": "2024-01-03T00:00:00Z",
|
||||||
|
"sha256Hash": "...",
|
||||||
|
"uri": "https://...",
|
||||||
|
"state": "ACTIVE"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"nextPageToken": "..."
|
||||||
|
}
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
|
||||||
|
|
||||||
|
# 构建查询参数
|
||||||
|
query_params = dict(request.query_params)
|
||||||
|
if pageSize is not None:
|
||||||
|
query_params["pageSize"] = pageSize
|
||||||
|
if pageToken is not None:
|
||||||
|
query_params["pageToken"] = pageToken
|
||||||
|
|
||||||
|
# 构建上游请求
|
||||||
|
upstream_url = _build_upstream_url(base_url, "/v1beta/files", query_params)
|
||||||
|
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
||||||
|
|
||||||
|
logger.debug(f"Gemini Files list proxy: GET {redact_url_for_log(upstream_url)}")
|
||||||
|
|
||||||
|
return await _proxy_request("GET", upstream_url, headers, file_key_id=file_key_id)
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# 获取文件元数据端点
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/v1beta/files/{file_name:path}")
|
||||||
|
async def get_file(
|
||||||
|
file_name: str,
|
||||||
|
request: Request,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
获取指定文件的元数据
|
||||||
|
|
||||||
|
**认证方式**:
|
||||||
|
- `x-goog-api-key` 请求头,或
|
||||||
|
- `?key=` URL 参数
|
||||||
|
|
||||||
|
**路径参数**:
|
||||||
|
- `file_name`: 文件名(格式:files/xxx 或 xxx)
|
||||||
|
|
||||||
|
**响应格式**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"name": "files/abc-123",
|
||||||
|
"displayName": "文件名",
|
||||||
|
"mimeType": "image/jpeg",
|
||||||
|
"sizeBytes": "12345",
|
||||||
|
"createTime": "2024-01-01T00:00:00Z",
|
||||||
|
"updateTime": "2024-01-01T00:00:00Z",
|
||||||
|
"expirationTime": "2024-01-03T00:00:00Z",
|
||||||
|
"sha256Hash": "...",
|
||||||
|
"uri": "https://...",
|
||||||
|
"state": "ACTIVE"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
|
||||||
|
|
||||||
|
# 规范化文件名(确保以 files/ 开头)
|
||||||
|
if not file_name.startswith("files/"):
|
||||||
|
file_name = f"files/{file_name}"
|
||||||
|
|
||||||
|
# 构建上游请求
|
||||||
|
upstream_url = _build_upstream_url(
|
||||||
|
base_url,
|
||||||
|
f"/v1beta/{file_name}",
|
||||||
|
dict(request.query_params),
|
||||||
|
)
|
||||||
|
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
||||||
|
|
||||||
|
logger.debug(f"Gemini Files get proxy: GET {redact_url_for_log(upstream_url)}")
|
||||||
|
|
||||||
|
return await _proxy_request("GET", upstream_url, headers, file_key_id=file_key_id)
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# 删除文件端点
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/v1beta/files/{file_name:path}")
|
||||||
|
async def delete_file(
|
||||||
|
file_name: str,
|
||||||
|
request: Request,
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
删除指定文件
|
||||||
|
|
||||||
|
**认证方式**:
|
||||||
|
- `x-goog-api-key` 请求头,或
|
||||||
|
- `?key=` URL 参数
|
||||||
|
|
||||||
|
**路径参数**:
|
||||||
|
- `file_name`: 文件名(格式:files/xxx 或 xxx)
|
||||||
|
|
||||||
|
**响应格式**:
|
||||||
|
成功时返回空 JSON 对象:`{}`
|
||||||
|
"""
|
||||||
|
upstream_key, base_url, _file_key_id = await _resolve_upstream_context(request, db)
|
||||||
|
|
||||||
|
# 规范化文件名(确保以 files/ 开头)
|
||||||
|
if not file_name.startswith("files/"):
|
||||||
|
file_name = f"files/{file_name}"
|
||||||
|
|
||||||
|
# 构建上游请求
|
||||||
|
upstream_url = _build_upstream_url(
|
||||||
|
base_url,
|
||||||
|
f"/v1beta/{file_name}",
|
||||||
|
dict(request.query_params),
|
||||||
|
)
|
||||||
|
headers = _build_upstream_headers(dict(request.headers), upstream_key)
|
||||||
|
|
||||||
|
logger.debug(f"Gemini Files delete proxy: DELETE {redact_url_for_log(upstream_url)}")
|
||||||
|
|
||||||
|
del _file_key_id # 显式标记:delete 端点不需要存储映射
|
||||||
|
response = await _proxy_request("DELETE", upstream_url, headers)
|
||||||
|
if response.status_code < 300:
|
||||||
|
await delete_file_key_mapping(file_name)
|
||||||
|
else:
|
||||||
|
logger.debug(
|
||||||
|
f"Gemini Files delete failed, skip mapping cleanup: status={response.status_code}"
|
||||||
|
)
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
@@ -245,3 +245,12 @@ register_capability(
|
|||||||
short_name="CLI 1M",
|
short_name="CLI 1M",
|
||||||
error_patterns=["context", "token", "length", "exceed"], # 上下文超限错误
|
error_patterns=["context", "token", "length", "exceed"], # 上下文超限错误
|
||||||
)
|
)
|
||||||
|
|
||||||
|
register_capability(
|
||||||
|
name="gemini_files_api",
|
||||||
|
display_name="Gemini文件上传",
|
||||||
|
description="支持 Gemini Files API(上传、查询、删除),第三方 Key 通常不支持",
|
||||||
|
match_mode=CapabilityMatchMode.COMPATIBLE, # 需要时选有的,不需要时都可选
|
||||||
|
config_mode=CapabilityConfigMode.REQUEST_PARAM, # 从请求路径检测
|
||||||
|
short_name="文件上传",
|
||||||
|
)
|
||||||
|
|||||||
@@ -389,6 +389,10 @@ openapi_tags = [
|
|||||||
"name": "Gemini API",
|
"name": "Gemini API",
|
||||||
"description": "Gemini API 代理接口,兼容 Google Gemini API 格式",
|
"description": "Gemini API 代理接口,兼容 Google Gemini API 格式",
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"name": "Gemini Files API",
|
||||||
|
"description": "Gemini Files API 代理接口,支持文件上传、查询、删除等操作",
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"name": "System Catalog",
|
"name": "System Catalog",
|
||||||
"description": "系统目录接口,用于获取可用模型列表等",
|
"description": "系统目录接口,用于获取可用模型列表等",
|
||||||
|
|||||||
95
src/services/gemini_files_mapping.py
Normal file
95
src/services/gemini_files_mapping.py
Normal file
@@ -0,0 +1,95 @@
|
|||||||
|
"""
|
||||||
|
Gemini Files API - 文件与 Key 绑定缓存
|
||||||
|
|
||||||
|
用于在上传文件后记录 file_id -> provider_key_id,
|
||||||
|
并在后续 generateContent 请求中优先使用同一 Key。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Dict, Optional, Set
|
||||||
|
|
||||||
|
from src.core.cache_service import CacheService
|
||||||
|
|
||||||
|
FILE_MAPPING_TTL_SECONDS = 60 * 60 * 48 # 48小时
|
||||||
|
FILE_MAPPING_CACHE_PREFIX = "gemini_files:key"
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_file_name(file_name: str) -> str:
|
||||||
|
name = (file_name or "").strip()
|
||||||
|
if not name:
|
||||||
|
return ""
|
||||||
|
return name if name.startswith("files/") else f"files/{name}"
|
||||||
|
|
||||||
|
|
||||||
|
def build_file_mapping_key(file_name: str) -> str:
|
||||||
|
normalized = _normalize_file_name(file_name)
|
||||||
|
return f"{FILE_MAPPING_CACHE_PREFIX}:{normalized}" if normalized else ""
|
||||||
|
|
||||||
|
|
||||||
|
async def store_file_key_mapping(file_name: str, key_id: str) -> None:
|
||||||
|
cache_key = build_file_mapping_key(file_name)
|
||||||
|
if not cache_key or not key_id:
|
||||||
|
return
|
||||||
|
await CacheService.set(cache_key, str(key_id), ttl_seconds=FILE_MAPPING_TTL_SECONDS)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_file_key_mapping(file_name: str) -> Optional[str]:
|
||||||
|
cache_key = build_file_mapping_key(file_name)
|
||||||
|
if not cache_key:
|
||||||
|
return None
|
||||||
|
value = await CacheService.get(cache_key)
|
||||||
|
if value:
|
||||||
|
return str(value)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def delete_file_key_mapping(file_name: str) -> None:
|
||||||
|
cache_key = build_file_mapping_key(file_name)
|
||||||
|
if cache_key:
|
||||||
|
await CacheService.delete(cache_key)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_file_name_from_uri(file_uri: str) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
从 fileUri 提取 files/xxx 名称。
|
||||||
|
|
||||||
|
支持两种格式:
|
||||||
|
- 完整 URL: https://generativelanguage.googleapis.com/v1beta/files/abc123
|
||||||
|
- 短格式: files/abc123
|
||||||
|
"""
|
||||||
|
if not file_uri:
|
||||||
|
return None
|
||||||
|
# 完整 URL 格式
|
||||||
|
if "/files/" in file_uri:
|
||||||
|
idx = file_uri.rfind("/files/")
|
||||||
|
return file_uri[idx + 1:] # 提取 files/xxx 部分
|
||||||
|
# 短格式
|
||||||
|
if file_uri.startswith("files/"):
|
||||||
|
return file_uri
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def extract_file_names_from_request(payload: Optional[Dict[str, Any]]) -> Set[str]:
|
||||||
|
"""
|
||||||
|
从 Gemini 请求体中提取 fileUri 使用到的 files/xxx 名称集合。
|
||||||
|
"""
|
||||||
|
results: Set[str] = set()
|
||||||
|
|
||||||
|
def walk(node: Any) -> None:
|
||||||
|
if isinstance(node, dict):
|
||||||
|
file_data = node.get("fileData") or node.get("file_data")
|
||||||
|
if isinstance(file_data, dict):
|
||||||
|
file_uri = file_data.get("fileUri") or file_data.get("file_uri")
|
||||||
|
if isinstance(file_uri, str):
|
||||||
|
file_name = _extract_file_name_from_uri(file_uri)
|
||||||
|
if file_name:
|
||||||
|
results.add(file_name)
|
||||||
|
for value in node.values():
|
||||||
|
walk(value)
|
||||||
|
elif isinstance(node, list):
|
||||||
|
for item in node:
|
||||||
|
walk(item)
|
||||||
|
|
||||||
|
if payload:
|
||||||
|
walk(payload)
|
||||||
|
|
||||||
|
return results
|
||||||
@@ -52,6 +52,7 @@ class CandidateResolver:
|
|||||||
request_id: Optional[str] = None,
|
request_id: Optional[str] = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
|
preferred_key_ids: Optional[list[str]] = None,
|
||||||
) -> Tuple[List[ProviderCandidate], str]:
|
) -> Tuple[List[ProviderCandidate], str]:
|
||||||
"""
|
"""
|
||||||
获取所有可用候选
|
获取所有可用候选
|
||||||
@@ -64,6 +65,7 @@ class CandidateResolver:
|
|||||||
request_id: 请求 ID(用于日志)
|
request_id: 请求 ID(用于日志)
|
||||||
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
||||||
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
||||||
|
preferred_key_ids: 优先使用的 Provider Key ID 列表(匹配则置顶)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(所有候选组合的列表, global_model_id)
|
(所有候选组合的列表, global_model_id)
|
||||||
@@ -107,6 +109,28 @@ class CandidateResolver:
|
|||||||
|
|
||||||
logger.debug(f" [{request_id}] 获取到 {len(all_candidates)} 个候选组合")
|
logger.debug(f" [{request_id}] 获取到 {len(all_candidates)} 个候选组合")
|
||||||
|
|
||||||
|
if preferred_key_ids:
|
||||||
|
preferred_set = {str(kid) for kid in preferred_key_ids if kid}
|
||||||
|
if preferred_set:
|
||||||
|
preferred_candidates = [
|
||||||
|
c for c in all_candidates if c.key and str(c.key.id) in preferred_set
|
||||||
|
]
|
||||||
|
other_candidates = [
|
||||||
|
c for c in all_candidates if not (c.key and str(c.key.id) in preferred_set)
|
||||||
|
]
|
||||||
|
if preferred_candidates:
|
||||||
|
matched_key_ids = [str(c.key.id) for c in preferred_candidates if c.key]
|
||||||
|
logger.debug(
|
||||||
|
f" [{request_id}] 优先候选命中: {len(preferred_candidates)} 个 "
|
||||||
|
f"(key_ids={matched_key_ids[:3]}{'...' if len(matched_key_ids) > 3 else ''})"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.debug(
|
||||||
|
f" [{request_id}] 优先候选未命中: 请求的 key_ids={list(preferred_set)[:3]} "
|
||||||
|
"不在可用候选中,将使用普通优先级"
|
||||||
|
)
|
||||||
|
all_candidates = preferred_candidates + other_candidates
|
||||||
|
|
||||||
# 如果没有解析到 global_model_id,使用原始 model_name 作为后备
|
# 如果没有解析到 global_model_id,使用原始 model_name 作为后备
|
||||||
return all_candidates, global_model_id or model_name
|
return all_candidates, global_model_id or model_name
|
||||||
|
|
||||||
|
|||||||
@@ -176,6 +176,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: Optional[str] = None,
|
request_id: Optional[str] = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
|
preferred_key_ids: Optional[list[str]] = None,
|
||||||
) -> Tuple[List[ProviderCandidate], str]:
|
) -> Tuple[List[ProviderCandidate], str]:
|
||||||
"""
|
"""
|
||||||
收集所有可用的 Provider/Endpoint/Key 候选组合
|
收集所有可用的 Provider/Endpoint/Key 候选组合
|
||||||
@@ -190,6 +191,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: 请求 ID(用于日志)
|
request_id: 请求 ID(用于日志)
|
||||||
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
||||||
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
||||||
|
preferred_key_ids: 优先使用的 Provider Key ID 列表(匹配则置顶)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(所有候选组合的列表, global_model_id)
|
(所有候选组合的列表, global_model_id)
|
||||||
@@ -206,6 +208,7 @@ class FallbackOrchestrator:
|
|||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
capability_requirements=capability_requirements,
|
capability_requirements=capability_requirements,
|
||||||
|
preferred_key_ids=preferred_key_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _create_candidate_records(
|
def _create_candidate_records(
|
||||||
@@ -988,6 +991,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: Optional[str] = None,
|
request_id: Optional[str] = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
|
preferred_key_ids: Optional[list[str]] = None,
|
||||||
request_body_ref: Optional[Dict[str, Any]] = None,
|
request_body_ref: Optional[Dict[str, Any]] = None,
|
||||||
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
|
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
|
||||||
"""
|
"""
|
||||||
@@ -1001,6 +1005,7 @@ class FallbackOrchestrator:
|
|||||||
request_id: 请求 ID(用于日志)
|
request_id: 请求 ID(用于日志)
|
||||||
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
|
||||||
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key)
|
||||||
|
preferred_key_ids: 优先使用的 Provider Key ID 列表(匹配则置顶)
|
||||||
request_body_ref: 请求体引用容器(用于 Thinking 签名错误重试)
|
request_body_ref: 请求体引用容器(用于 Thinking 签名错误重试)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -1036,6 +1041,7 @@ class FallbackOrchestrator:
|
|||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
capability_requirements=capability_requirements,
|
capability_requirements=capability_requirements,
|
||||||
|
preferred_key_ids=preferred_key_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. 批量创建候选记录
|
# 2. 批量创建候选记录
|
||||||
|
|||||||
Reference in New Issue
Block a user