Files
Aether/_deprecated_py_src/services/gemini_files_mapping.py
fawney19 1d9c77522a refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
2026-04-03 16:26:16 +08:00

381 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Gemini Files API - 文件与 Key 绑定映射服务
用于在上传文件后记录 file_id -> provider_key_id
并在后续 generateContent 请求中优先使用同一 Key。
存储策略:
- 数据库(持久化):主存储,支持服务重启后恢复
- Redis缓存加速读取TTL=48小时
读取策略:
1. 先查 Redis 缓存
2. 缓存未命中时回查数据库
3. 从数据库读取后回填缓存
清理策略:
- 数据库中 expires_at 过期的记录由定时任务清理
- Redis 缓存由 TTL 自动过期
"""
from __future__ import annotations
import uuid
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import delete
from sqlalchemy.orm import Session
from src.core.cache_service import CacheService
from src.core.logger import logger
FILE_MAPPING_TTL_SECONDS = 60 * 60 * 48 # 48小时
FILE_MAPPING_CACHE_PREFIX = "gemini_files:key"
def _normalize_file_name(file_name: str) -> str:
"""规范化文件名,确保以 files/ 开头"""
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:
"""构建 Redis 缓存键"""
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,
user_id: str | None = None,
display_name: str | None = None,
mime_type: str | None = None,
source_hash: str | None = None,
) -> None:
"""
存储文件→Key 映射(同时写入 Redis 和数据库)
Args:
file_name: 文件名(如 files/abc123
key_id: Provider Key ID
user_id: 用户 ID可选用于权限验证
display_name: 文件显示名(可选)
mime_type: 文件 MIME 类型(可选)
source_hash: 源文件哈希(可选,用于关联相同源文件的不同上传)
"""
normalized_name = _normalize_file_name(file_name)
if not normalized_name or not key_id:
return
# 1. 写入 Redis 缓存
cache_key = build_file_mapping_key(normalized_name)
await CacheService.set(cache_key, str(key_id), ttl_seconds=FILE_MAPPING_TTL_SECONDS)
# 2. 写入数据库(异步执行,不阻塞主流程)
try:
await _store_to_database(
file_name=normalized_name,
key_id=key_id,
user_id=user_id,
display_name=display_name,
mime_type=mime_type,
source_hash=source_hash,
)
except Exception as e:
# 数据库写入失败只记录警告,不影响主流程
logger.warning(f"Failed to persist Gemini file mapping to database: {e}")
async def _store_to_database(
file_name: str,
key_id: str,
user_id: str | None = None,
display_name: str | None = None,
mime_type: str | None = None,
source_hash: str | None = None,
) -> None:
"""将映射写入数据库"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
now = datetime.now(timezone.utc)
expires_at = now + timedelta(hours=48)
with get_db_context() as db:
# 使用 upsert 逻辑:存在则更新,不存在则插入
existing = (
db.query(GeminiFileMapping).filter(GeminiFileMapping.file_name == file_name).first()
)
if existing:
# 更新现有记录
existing.key_id = key_id
existing.user_id = user_id
existing.display_name = display_name
existing.mime_type = mime_type
existing.source_hash = source_hash
existing.expires_at = expires_at
else:
# 插入新记录
mapping = GeminiFileMapping(
id=str(uuid.uuid4()),
file_name=file_name,
key_id=key_id,
user_id=user_id,
display_name=display_name,
mime_type=mime_type,
source_hash=source_hash,
created_at=now,
expires_at=expires_at,
)
db.add(mapping)
db.commit()
async def get_file_key_mapping(file_name: str) -> str | None:
"""
获取文件→Key 映射
读取策略:
1. 先查 Redis 缓存
2. 缓存未命中时回查数据库
3. 从数据库读取后回填缓存
Args:
file_name: 文件名(如 files/abc123
Returns:
Provider Key ID如果不存在或已过期则返回 None
"""
normalized_name = _normalize_file_name(file_name)
if not normalized_name:
return None
cache_key = build_file_mapping_key(normalized_name)
# 1. 先查 Redis 缓存
cached_value = await CacheService.get(cache_key)
if cached_value:
return str(cached_value)
# 2. 缓存未命中,回查数据库
key_id = await _get_from_database(normalized_name)
if key_id:
# 3. 回填缓存(使用剩余有效期或默认 TTL
await CacheService.set(cache_key, key_id, ttl_seconds=FILE_MAPPING_TTL_SECONDS)
logger.debug(f"Gemini file mapping cache refilled from database: {normalized_name}")
return key_id
async def _get_from_database(file_name: str) -> str | None:
"""从数据库查询映射"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
now = datetime.now(timezone.utc)
try:
with get_db_context() as db:
mapping = (
db.query(GeminiFileMapping)
.filter(
GeminiFileMapping.file_name == file_name,
GeminiFileMapping.expires_at > now, # 只返回未过期的
)
.first()
)
if mapping:
return str(mapping.key_id)
except Exception as e:
logger.warning(f"Failed to query Gemini file mapping from database: {e}")
return None
async def get_all_key_ids_for_file(file_name: str) -> list[str]:
"""
获取支持指定文件的所有 Key ID 列表
当同一个源文件被上传到多个 Key 时,返回所有可用的 Key ID。
这允许系统在首选 Key 不可用时选择其他 Key。
Args:
file_name: 文件名(如 files/abc123
Returns:
所有支持该文件的 Key ID 列表(包括原始映射和具有相同 source_hash 的映射)
"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
normalized_name = _normalize_file_name(file_name)
if not normalized_name:
return []
now = datetime.now(timezone.utc)
try:
with get_db_context() as db:
# 首先获取原始映射
original_mapping = (
db.query(GeminiFileMapping)
.filter(
GeminiFileMapping.file_name == normalized_name,
GeminiFileMapping.expires_at > now,
)
.first()
)
if not original_mapping:
return []
key_ids = [str(original_mapping.key_id)]
# 如果有 source_hash查找所有具有相同 source_hash 的映射
if original_mapping.source_hash:
related_mappings = (
db.query(GeminiFileMapping)
.filter(
GeminiFileMapping.source_hash == original_mapping.source_hash,
GeminiFileMapping.expires_at > now,
GeminiFileMapping.file_name != normalized_name, # 排除原始映射
)
.all()
)
for mapping in related_mappings:
kid = str(mapping.key_id)
if kid not in key_ids:
key_ids.append(kid)
return key_ids
except Exception as e:
logger.warning(f"Failed to query related Gemini file mappings: {e}")
return []
async def delete_file_key_mapping(file_name: str) -> None:
"""
删除文件→Key 映射(同时从 Redis 和数据库删除)
Args:
file_name: 文件名(如 files/abc123
"""
normalized_name = _normalize_file_name(file_name)
if not normalized_name:
return
# 1. 从 Redis 删除
cache_key = build_file_mapping_key(normalized_name)
await CacheService.delete(cache_key)
# 2. 从数据库删除
try:
await _delete_from_database(normalized_name)
except Exception as e:
logger.warning(f"Failed to delete Gemini file mapping from database: {e}")
async def _delete_from_database(file_name: str) -> None:
"""从数据库删除映射"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
with get_db_context() as db:
db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.file_name == file_name))
db.commit()
# =============================================================================
# 同步接口(用于定时任务等场景)
# =============================================================================
def cleanup_expired_mappings(db: Session) -> int:
"""
清理过期的文件映射记录(同步方法,供定时任务调用)
Args:
db: 数据库会话
Returns:
删除的记录数
"""
from src.models.database import GeminiFileMapping
now = datetime.now(timezone.utc)
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
db.commit()
deleted_count = result.rowcount
if deleted_count > 0:
logger.info(f"Cleaned up {deleted_count} expired Gemini file mappings")
return deleted_count
# =============================================================================
# 请求解析工具函数
# =============================================================================
def _extract_file_name_from_uri(file_uri: str) -> str | None:
"""
从 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: dict[str, Any] | None) -> 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