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:
fawney19
2026-01-30 11:37:34 +08:00
parent 32b293ef3e
commit c33cd07ea9
12 changed files with 856 additions and 0 deletions

View File

@@ -439,6 +439,14 @@ class BaseMessageHandler:
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:
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
if provider_type:

View File

@@ -511,6 +511,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
capability_requirements = self._resolve_capability_requirements(
model_name=model,
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
@@ -529,6 +534,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
request_id=self.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref, # 传递容器引用
)
@@ -1126,6 +1132,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
capability_requirements = self._resolve_capability_requirements(
model_name=model,
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_id=self.request_id,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref, # 传递容器引用
)

View File

@@ -568,6 +568,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
capability_requirements = self._resolve_capability_requirements(
model_name=ctx.model,
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
@@ -586,6 +591,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
request_id=self.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref, # 传递容器引用
)
@@ -2324,6 +2330,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
capability_requirements = self._resolve_capability_requirements(
model_name=model,
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_id=self.request_id,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref, # 传递容器引用
)

View File

@@ -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.logger import logger
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
@@ -56,6 +57,16 @@ class GeminiChatAdapter(ChatAdapterBase):
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(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
) -> Dict[str, Any]:

View File

@@ -22,6 +22,58 @@ class GeminiChatHandler(ChatHandlerBase):
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(
self,
request_body: Dict[str, Any],

View File

@@ -6,6 +6,7 @@ from .capabilities import router as capabilities_router
from .catalog import router as catalog_router
from .claude import router as claude_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 .modules import router as modules_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(openai_router)
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(catalog_router)
router.include_router(capabilities_router)

View 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

View File

@@ -245,3 +245,12 @@ register_capability(
short_name="CLI 1M",
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="文件上传",
)

View File

@@ -389,6 +389,10 @@ openapi_tags = [
"name": "Gemini API",
"description": "Gemini API 代理接口,兼容 Google Gemini API 格式",
},
{
"name": "Gemini Files API",
"description": "Gemini Files API 代理接口,支持文件上传、查询、删除等操作",
},
{
"name": "System Catalog",
"description": "系统目录接口,用于获取可用模型列表等",

View 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

View File

@@ -52,6 +52,7 @@ class CandidateResolver:
request_id: Optional[str] = None,
is_stream: bool = False,
capability_requirements: Optional[Dict[str, bool]] = None,
preferred_key_ids: Optional[list[str]] = None,
) -> Tuple[List[ProviderCandidate], str]:
"""
获取所有可用候选
@@ -64,6 +65,7 @@ class CandidateResolver:
request_id: 请求 ID用于日志
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key
preferred_key_ids: 优先使用的 Provider Key ID 列表(匹配则置顶)
Returns:
(所有候选组合的列表, global_model_id)
@@ -107,6 +109,28 @@ class CandidateResolver:
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 作为后备
return all_candidates, global_model_id or model_name

View File

@@ -176,6 +176,7 @@ class FallbackOrchestrator:
request_id: Optional[str] = None,
is_stream: bool = False,
capability_requirements: Optional[Dict[str, bool]] = None,
preferred_key_ids: Optional[list[str]] = None,
) -> Tuple[List[ProviderCandidate], str]:
"""
收集所有可用的 Provider/Endpoint/Key 候选组合
@@ -190,6 +191,7 @@ class FallbackOrchestrator:
request_id: 请求 ID用于日志
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key
preferred_key_ids: 优先使用的 Provider Key ID 列表(匹配则置顶)
Returns:
(所有候选组合的列表, global_model_id)
@@ -206,6 +208,7 @@ class FallbackOrchestrator:
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
)
def _create_candidate_records(
@@ -988,6 +991,7 @@ class FallbackOrchestrator:
request_id: Optional[str] = None,
is_stream: bool = False,
capability_requirements: Optional[Dict[str, bool]] = None,
preferred_key_ids: Optional[list[str]] = None,
request_body_ref: Optional[Dict[str, Any]] = None,
) -> Tuple[Any, str, Optional[str], Optional[str], Optional[str], Optional[str]]:
"""
@@ -1001,6 +1005,7 @@ class FallbackOrchestrator:
request_id: 请求 ID用于日志
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(用于过滤不满足能力要求的 Key
preferred_key_ids: 优先使用的 Provider Key ID 列表(匹配则置顶)
request_body_ref: 请求体引用容器(用于 Thinking 签名错误重试)
Returns:
@@ -1036,6 +1041,7 @@ class FallbackOrchestrator:
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
)
# 2. 批量创建候选记录