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