refactor: 调度器迁移至独立模块,消除 services->api 反向依赖

- 将调度器相关模块从 src/services/cache/ 迁移到 src/services/scheduling/
- 下沉类型定义到 core 层: AccessRestrictions, ProviderAuthInfo, ParsedChunk/StreamStats, 视频工具函数
- 提取 thinking_cache 签名缓存到 core/api_format/conversion/
- 提取 provider 认证逻辑到 services/provider/auth
- 提取遥测记录到 services/usage/telemetry
- 提取 models 列表缓存到 services/cache/model_list_cache
- 更新所有引用方的 import 路径及相关测试
This commit is contained in:
fawney19
2026-02-16 11:00:48 +08:00
parent 4dc401677d
commit 63870931af
85 changed files with 1907 additions and 1463 deletions

View File

@@ -32,7 +32,7 @@ from src.models.database import (
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache.aware_scheduler import CacheAwareScheduler
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
from src.services.system.config import SystemConfigService
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])

View File

@@ -22,8 +22,8 @@ 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, User
from src.services.cache.affinity_manager import get_affinity_manager
from src.services.cache.aware_scheduler import CacheAwareScheduler, get_cache_aware_scheduler
from src.services.scheduling.affinity_manager import get_affinity_manager
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, get_cache_aware_scheduler
from src.services.system.config import SystemConfigService
router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitoring: Cache"])
@@ -1103,8 +1103,8 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
class AdminCacheConfigAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.config.constants import ConcurrencyDefaults
from src.services.cache.affinity_manager import CacheAffinityManager
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
from src.services.scheduling.affinity_manager import CacheAffinityManager
# 获取动态预留管理器的配置
reservation_manager = get_adaptive_reservation_manager()

View File

@@ -643,7 +643,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
if self.key in ("scheduling_mode", "provider_priority_mode"):
try:
from src.clients.redis_client import get_redis_client_sync
from src.services.cache.aware_scheduler import get_cache_aware_scheduler
from src.services.scheduling.aware_scheduler import get_cache_aware_scheduler
redis_client = get_redis_client_sync()
# 从数据库读取两个调度配置的最新值,确保一致性

View File

@@ -19,15 +19,18 @@ from sqlalchemy import tuple_
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL
from src.core.access_restrictions import AccessRestrictions
from src.core.api_format.conversion.compatibility import is_format_compatible
from src.core.cache_service import CacheService
from src.core.logger import logger
from src.models.database import ApiKey, Model, Provider, ProviderEndpoint, User
from src.models.database import Model, Provider, ProviderEndpoint
from src.services.cache.model_list_cache import MODELS_LIST_CACHE_PREFIX as _CACHE_KEY_PREFIX
from src.services.cache.model_list_cache import (
invalidate_models_list_cache,
)
from src.services.model.availability import ModelAvailabilityQuery
from src.services.provider.format import normalize_endpoint_signature
# 缓存 key 前缀
_CACHE_KEY_PREFIX = "models:list"
_CACHE_TTL = CacheTTL.MODEL # 300 秒
@@ -70,21 +73,7 @@ async def _set_cached_models(
logger.warning(f"[ModelsService] 缓存写入失败: {e}")
async def invalidate_models_list_cache() -> None:
"""
清除所有 /v1/models 列表缓存
在模型创建、更新、删除时调用,确保模型列表实时更新
"""
try:
# 使用通配符删除所有 models:list:* 缓存(包括多格式组合的 key
deleted = await CacheService.delete_pattern(f"{_CACHE_KEY_PREFIX}:*")
if deleted > 0:
logger.info(f"[ModelsService] 已清除 {deleted}{_CACHE_KEY_PREFIX} 缓存")
else:
logger.debug(f"[ModelsService] 无 {_CACHE_KEY_PREFIX} 缓存需要清除")
except Exception as e:
logger.warning(f"[ModelsService] 清除缓存失败: {e}")
__all__ = ["AccessRestrictions", "invalidate_models_list_cache", "ModelInfo"]
@dataclass
@@ -115,91 +104,7 @@ class ModelInfo:
output_modalities: list[str] | None = None
@dataclass
class AccessRestrictions:
"""API Key 或 User 的访问限制"""
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
allowed_models: list[str] | None = None # 允许的模型名称列表
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
@classmethod
def from_api_key_and_user(cls, api_key: ApiKey | None, user: User | None) -> AccessRestrictions:
"""
从 API Key 和 User 合并访问限制
限制逻辑:
- API Key 的限制优先于 User 的限制
- 如果 API Key 有限制,使用 API Key 的限制
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
- 两者都无限制则返回空限制
"""
allowed_providers: list[str] | None = None
allowed_models: list[str] | None = None
allowed_api_formats: list[str] | None = None
# 优先使用 API Key 的限制
if api_key:
if api_key.allowed_providers is not None:
allowed_providers = api_key.allowed_providers
if api_key.allowed_models is not None:
allowed_models = api_key.allowed_models
if api_key.allowed_api_formats is not None:
allowed_api_formats = api_key.allowed_api_formats
# 如果 API Key 没有限制,检查 User 的限制
if user:
if allowed_providers is None and user.allowed_providers is not None:
allowed_providers = user.allowed_providers
if allowed_models is None and user.allowed_models is not None:
allowed_models = user.allowed_models
if allowed_api_formats is None and user.allowed_api_formats is not None:
allowed_api_formats = user.allowed_api_formats
return cls(
allowed_providers=allowed_providers,
allowed_models=allowed_models,
allowed_api_formats=allowed_api_formats,
)
def is_api_format_allowed(self, api_format: str) -> bool:
"""
检查 API 格式是否被允许
Args:
api_format: endpoint signature"openai:chat"
Returns:
True 如果格式被允许False 否则
"""
if self.allowed_api_formats is None:
return True
target = normalize_endpoint_signature(api_format)
allowed = {normalize_endpoint_signature(f) for f in self.allowed_api_formats if f}
return target in allowed
def is_model_allowed(self, model_id: str, provider_id: str) -> bool:
"""
检查模型是否被允许访问
Args:
model_id: 模型 ID
provider_id: Provider ID
Returns:
True 如果模型被允许False 否则
"""
# 检查 Provider 限制
if self.allowed_providers is not None:
if provider_id not in self.allowed_providers:
return False
# 检查模型限制
if self.allowed_models is not None:
if model_id not in self.allowed_models:
return False
return True
# AccessRestrictions -- re-export from src.core.access_restrictions (see __all__)
def _normalize_api_formats(

View File

@@ -45,8 +45,8 @@ from sqlalchemy.orm import Session
from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger
from src.services.provider.format import normalize_endpoint_signature
from src.services.system.audit import audit_service
from src.services.usage.service import UsageService
from src.services.usage.telemetry import MessageTelemetry # re-export
if TYPE_CHECKING:
from src.api.handlers.base.stream_context import StreamContext
@@ -55,286 +55,8 @@ if TYPE_CHECKING:
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
class MessageTelemetry:
"""
负责记录 Usage/Audit避免处理器里重复代码。
"""
def __init__(
self, db: Session, user: Any, api_key: Any, request_id: str, client_ip: str
) -> None:
self.db = db
self.user = user
self.api_key = api_key
self.request_id = request_id
self.client_ip = client_ip
async def calculate_cost(
self,
provider: str,
model: str,
*,
input_tokens: int,
output_tokens: int,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
) -> float:
input_price, output_price = await UsageService.get_model_price_async(
self.db, provider, model
)
_, _, _, _, _, _, total_cost = UsageService.calculate_cost(
input_tokens,
output_tokens,
input_price,
output_price,
cache_creation_tokens,
cache_read_tokens,
*await UsageService.get_cache_prices_async(self.db, provider, model, input_price),
)
return total_cost
async def record_success(
self,
*,
provider: str,
model: str,
input_tokens: int,
output_tokens: int,
response_time_ms: int,
status_code: int,
request_body: dict[str, Any],
request_headers: dict[str, Any],
response_body: Any,
response_headers: dict[str, Any],
client_response_headers: dict[str, Any] | None = None,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
is_stream: bool = False,
provider_request_headers: dict[str, Any] | None = None,
# 时间指标
first_byte_time_ms: int | None = None, # 首字时间/TTFB
# Provider 侧追踪信息(用于记录真实成本)
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
api_format: str | None = None,
# 格式转换追踪
endpoint_api_format: str | None = None, # 端点原生 API 格式
has_format_conversion: bool = False, # 是否发生了格式转换
# 模型映射信息
target_model: str | None = None,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata: dict[str, Any] | None = None,
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> float:
metadata = response_metadata
if request_metadata:
merged = dict(request_metadata)
if response_metadata:
merged.setdefault("response", response_metadata)
metadata = merged
usage = await UsageService.record_usage(
db=self.db,
user=self.user,
api_key=self.api_key,
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=cache_read_tokens,
request_type="chat",
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
is_stream=is_stream,
response_time_ms=response_time_ms,
first_byte_time_ms=first_byte_time_ms, # 传递首字时间
status_code=status_code,
request_headers=request_headers,
request_body=request_body,
provider_request_headers=provider_request_headers or {},
response_headers=response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
request_id=self.request_id,
# Provider 侧追踪信息(用于记录真实成本)
provider_id=provider_id,
provider_endpoint_id=provider_endpoint_id,
provider_api_key_id=provider_api_key_id,
# 模型映射信息
target_model=target_model,
# Provider 响应元数据/请求元数据
metadata=metadata,
)
total_cost = float(getattr(usage, "total_cost_usd", 0.0) or 0.0)
if self.user and self.api_key:
audit_service.log_api_request(
db=self.db,
user_id=self.user.id,
api_key_id=self.api_key.id,
request_id=self.request_id,
model=model,
provider=provider,
success=True,
ip_address=self.client_ip,
status_code=status_code,
input_tokens=getattr(usage, "input_tokens", input_tokens),
output_tokens=getattr(usage, "output_tokens", output_tokens),
cost_usd=total_cost,
)
return total_cost
async def record_failure(
self,
*,
provider: str,
model: str,
response_time_ms: int,
status_code: int,
error_message: str,
request_body: dict[str, Any],
request_headers: dict[str, Any],
is_stream: bool,
api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
# 格式转换追踪
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
# 模型映射信息
target_model: str | None = None,
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> None:
"""
记录失败请求
注意Provider 链路信息provider_id, endpoint_id, key_id不在此处记录
因为 RequestCandidate 表已经记录了完整的请求链路追踪信息。
Args:
input_tokens: 预估输入 tokens来自 message_start用于中断请求的成本估算
output_tokens: 预估输出 tokens来自已收到的内容
cache_creation_tokens: 缓存创建 tokens
cache_read_tokens: 缓存读取 tokens
response_body: 响应体(如果有部分响应)
response_headers: 响应头Provider 返回的原始响应头)
client_response_headers: 返回给客户端的响应头
target_model: 映射后的目标模型名(如果发生了映射)
"""
provider_name = provider or "unknown"
if provider_name == "unknown":
logger.warning(
f"[Telemetry] Recording failure with unknown provider (request_id={self.request_id})"
)
await UsageService.record_usage(
db=self.db,
user=self.user,
api_key=self.api_key,
provider=provider_name,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=cache_read_tokens,
request_type="chat",
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
is_stream=is_stream,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=error_message,
request_headers=request_headers,
request_body=request_body,
provider_request_headers=provider_request_headers or {},
response_headers=response_headers or {},
client_response_headers=client_response_headers,
response_body=response_body or {"error": error_message},
request_id=self.request_id,
# 模型映射信息
target_model=target_model,
# 请求元数据
metadata=request_metadata,
)
async def record_cancelled(
self,
*,
provider: str,
model: str,
response_time_ms: int,
first_byte_time_ms: int | None,
status_code: int,
request_body: dict[str, Any],
request_headers: dict[str, Any],
is_stream: bool,
api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
# 格式转换追踪
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
target_model: str | None = None,
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> None:
"""
记录客户端取消的请求
客户端主动断开连接不算系统失败,使用 cancelled 状态。
"""
provider_name = provider or "unknown"
await UsageService.record_usage(
db=self.db,
user=self.user,
api_key=self.api_key,
provider=provider_name,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=cache_read_tokens,
request_type="chat",
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
is_stream=is_stream,
response_time_ms=response_time_ms,
first_byte_time_ms=first_byte_time_ms,
status_code=status_code,
status="cancelled",
request_headers=request_headers,
request_body=request_body,
provider_request_headers=provider_request_headers or {},
response_headers=response_headers or {},
client_response_headers=client_response_headers,
response_body=response_body or {},
request_id=self.request_id,
target_model=target_model,
metadata=request_metadata,
)
# MessageTelemetry -- re-export from src.services.usage.telemetry (see import above)
__all__ = ["MessageTelemetry", "MessageHandlerProtocol", "AdapterDetectorType"]
@runtime_checkable

View File

@@ -13,8 +13,8 @@ from src.api.handlers.base.utils import get_format_converter_registry
from src.core.exceptions import ThinkingSignatureException, UpstreamClientException
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.transport import get_vertex_ai_effective_format
from src.services.scheduling.aware_scheduler import ProviderCandidate
def _get_error_status_code(e: Exception, default: int = 400) -> int:

View File

@@ -75,7 +75,6 @@ from src.models.database import (
ProviderEndpoint,
User,
)
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.behavior import get_provider_behavior
from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream,
@@ -85,6 +84,7 @@ from src.services.provider.stream_policy import (
from src.services.provider.transport import (
build_provider_url,
)
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.system.config import SystemConfigService

View File

@@ -49,7 +49,7 @@ if TYPE_CHECKING:
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
@dataclass

View File

@@ -38,7 +38,6 @@ from src.core.exceptions import (
ProviderTimeoutException,
)
from src.core.logger import logger
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.behavior import get_provider_behavior
from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream,
@@ -46,6 +45,7 @@ from src.services.provider.stream_policy import (
resolve_upstream_is_stream,
)
from src.services.provider.transport import build_provider_url
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.system.config import SystemConfigService
from src.utils.sse_parser import SSEEventParser
from src.utils.timeout import read_first_chunk_with_ttfb_timeout

View File

@@ -29,7 +29,6 @@ from src.core.exceptions import (
ThinkingSignatureException,
)
from src.core.logger import logger
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.behavior import get_provider_behavior
from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream,
@@ -37,6 +36,7 @@ from src.services.provider.stream_policy import (
resolve_upstream_is_stream,
)
from src.services.provider.transport import build_provider_url
from src.services.scheduling.aware_scheduler import ProviderCandidate
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol

View File

@@ -693,36 +693,22 @@ class GeminiCliResponseParser(GeminiResponseParser):
self.api_format = "gemini:cli"
# 解析器注册表
_PARSERS: dict[str, type[ResponseParser]] = {
"claude:chat": ClaudeResponseParser,
"claude:cli": ClaudeCliResponseParser,
"openai:chat": OpenAIResponseParser,
"openai:cli": OpenAICliResponseParser,
"gemini:chat": GeminiResponseParser,
"gemini:cli": GeminiCliResponseParser,
}
# 注册解析器到 core 层注册表(供 services 层通过 format_id 获取)
from src.core.stream_types import get_parser_for_format, register_parser
def get_parser_for_format(format_id: str) -> ResponseParser:
"""
根据格式 ID 获取 ResponseParser
def register_default_parsers() -> None:
register_parser("claude:chat", ClaudeResponseParser)
register_parser("claude:cli", ClaudeCliResponseParser)
register_parser("openai:chat", OpenAIResponseParser)
register_parser("openai:cli", OpenAICliResponseParser)
register_parser("gemini:chat", GeminiResponseParser)
register_parser("gemini:cli", GeminiCliResponseParser)
Args:
format_id: endpoint signature"claude:chat", "openai:cli"
Returns:
ResponseParser 实例
Raises:
KeyError: 格式不存在
"""
from src.core.api_format.signature import normalize_signature_key
normalized = normalize_signature_key(format_id)
if normalized not in _PARSERS:
raise KeyError(f"Unknown format: {normalized}")
return _PARSERS[normalized]()
# 模块加载时自动注册(保证 import parsers 即可用,测试也不需要手动初始化)
# main.py lifespan 中的显式调用是冗余但无害的安全保障dict 覆盖幂等)
register_default_parsers()
__all__ = [
@@ -732,6 +718,7 @@ __all__ = [
"ClaudeCliResponseParser",
"GeminiResponseParser",
"GeminiCliResponseParser",
"register_default_parsers",
"get_parser_for_format",
"is_cli_format",
]

View File

@@ -14,17 +14,10 @@
from __future__ import annotations
import copy
import json
import re
import time
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from typing import Any
import httpx
from sqlalchemy.orm import object_session
from src.clients.redis_client import get_redis_client
from src.core.api_format import (
UPSTREAM_DROP_HEADERS,
HeaderBuilder,
@@ -32,99 +25,9 @@ from src.core.api_format import (
make_signature_key,
)
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
from src.core.provider_auth_types import ProviderAuthInfo
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
if TYPE_CHECKING:
from src.models.database import ProviderAPIKey, ProviderEndpoint
# ==============================================================================
# Service Account 认证结果类型
# ==============================================================================
@dataclass
class ProviderAuthInfo:
"""Provider 认证信息(用于 Service Account 等异步认证场景)"""
auth_header: str
auth_value: str
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
decrypted_auth_config: dict[str, Any] | None = None
def as_tuple(self) -> tuple[str, str]:
"""返回 (auth_header, auth_value) 元组"""
return (self.auth_header, self.auth_value)
# ==============================================================================
# OAuth Token Refresh helpers
# ==============================================================================
async def _acquire_refresh_lock(key_id: str) -> tuple[Any, bool]:
"""尝试获取 OAuth refresh 分布式锁。
返回 ``(redis_client | None, got_lock)``。调用方在刷新完成后
必须调用 :func:`_release_refresh_lock` 释放锁。
"""
redis = await get_redis_client(require_redis=False)
lock_key = f"provider_oauth_refresh_lock:{key_id}"
got_lock = False
if redis is not None:
try:
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
except Exception:
got_lock = False
return redis, got_lock
async def _release_refresh_lock(redis: Any, key_id: str) -> None:
"""释放 OAuth refresh 分布式锁best-effort"""
if redis is not None:
try:
await redis.delete(f"provider_oauth_refresh_lock:{key_id}")
except Exception:
pass
def _persist_refreshed_token(
key: Any,
access_token: str,
token_meta: dict[str, Any],
) -> None:
"""将刷新后的 access_token 和 auth_config 持久化到数据库。"""
key.api_key = crypto_service.encrypt(access_token)
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
sess = object_session(key)
if sess is not None:
sess.add(key)
sess.commit()
else:
logger.warning(
"[OAUTH_REFRESH] key {} refreshed but cannot persist (no session); "
"next request will refresh again",
key.id,
)
def _get_proxy_config(key: Any, endpoint: Any = None) -> Any:
"""获取有效代理配置Key 级别优先于 Provider 级别)。"""
try:
from src.services.proxy_node.resolver import resolve_effective_proxy
provider = getattr(key, "provider", None) or (
getattr(endpoint, "provider", None) if endpoint else None
)
provider_proxy = getattr(provider, "proxy", None)
key_proxy = getattr(key, "proxy", None)
return resolve_effective_proxy(provider_proxy, key_proxy)
except Exception:
return None
from src.services.provider.auth import get_provider_auth
# ==============================================================================
# 统一的头部配置常量
@@ -1003,303 +906,3 @@ def build_passthrough_request(
endpoint,
key,
)
# ==============================================================================
# OAuth Token Refresh logic (Kiro / Generic)
# ==============================================================================
async def _refresh_kiro_token(
key: Any,
endpoint: Any,
token_meta: dict[str, Any],
) -> dict[str, Any]:
"""Kiro OAuth refresh: validate + call Kiro-specific refresh endpoint."""
from src.core.exceptions import InvalidRequestException
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
from src.services.provider.adapters.kiro.token_manager import (
refresh_access_token,
validate_refresh_token,
)
cfg = KiroAuthConfig.from_dict(token_meta or {})
if not (cfg.refresh_token or "").strip():
raise InvalidRequestException(
"Kiro auth_config missing refresh_token; please re-import credentials."
)
proxy_config = _get_proxy_config(key, endpoint)
validate_refresh_token(cfg.refresh_token)
access_token, new_cfg = await refresh_access_token(
cfg,
proxy_config=proxy_config,
)
new_meta = new_cfg.to_dict()
new_meta["updated_at"] = int(time.time())
_persist_refreshed_token(key, access_token, new_meta)
return new_meta
async def _refresh_generic_oauth_token(
key: Any,
endpoint: Any,
template: Any,
provider_type: str,
refresh_token: str,
token_meta: dict[str, Any],
) -> dict[str, Any]:
"""Generic OAuth refresh via template (Codex, Antigravity, ClaudeCode, etc.)."""
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scopes = getattr(template.oauth, "scopes", None) or []
scope_str = " ".join(scopes) if scopes else ""
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": str(refresh_token),
}
if scope_str:
body["scope"] = scope_str
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": str(refresh_token),
}
if scope_str:
form["scope"] = scope_str
if template.oauth.client_secret:
form["client_secret"] = template.oauth.client_secret
headers = {
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
}
data = form
json_body = None
proxy_config = _get_proxy_config(key, endpoint)
resp = await post_oauth_token(
provider_type=provider_type,
token_url=token_url,
headers=headers,
data=data,
json_body=json_body,
proxy_config=proxy_config,
timeout_seconds=30.0,
)
if 200 <= resp.status_code < 300:
token = resp.json()
access_token = str(token.get("access_token") or "")
new_refresh_token = str(token.get("refresh_token") or "")
expires_in = token.get("expires_in")
new_expires_at: int | None = None
try:
if expires_in is not None:
new_expires_at = int(time.time()) + int(expires_in)
except Exception:
new_expires_at = None
if access_token:
token_meta["token_type"] = token.get("token_type")
if new_refresh_token:
token_meta["refresh_token"] = new_refresh_token
token_meta["expires_at"] = new_expires_at
token_meta["scope"] = token.get("scope")
token_meta["updated_at"] = int(time.time())
token_meta = await enrich_auth_config(
provider_type=provider_type,
auth_config=token_meta,
token_response=token,
access_token=access_token,
proxy_config=proxy_config,
)
_persist_refreshed_token(key, access_token, token_meta)
else:
logger.warning(
"OAuth token refresh failed: provider={}, key_id={}, status={}",
provider_type,
getattr(key, "id", "?"),
resp.status_code,
)
return token_meta
# ==============================================================================
# Service Account 认证支持
# ==============================================================================
async def get_provider_auth(
endpoint: "ProviderEndpoint",
key: "ProviderAPIKey",
*,
force_refresh: bool = False,
) -> ProviderAuthInfo | None:
"""
获取 Provider 的认证信息
对于标准 API Key返回 None由 build_headers 自动处理)。
对于 Service Account异步获取 Access Token 并返回认证信息。
Args:
endpoint: 端点配置
key: Provider API Key
Returns:
Service Account 场景: ProviderAuthInfo 对象(包含认证信息和解密后的配置)
API Key 场景: None由 build_headers 处理)
Raises:
InvalidRequestException: 认证配置无效或认证失败
"""
from src.core.exceptions import InvalidRequestException
auth_type = getattr(key, "auth_type", "api_key")
if auth_type == "oauth":
# OAuth token 保存在 key.api_key加密refresh_token/expires_at 等在 auth_config加密 JSON中。
# 在请求前做一次懒刷新:接近过期时刷新 access_token并用 Redis lock 避免并发风暴。
encrypted_auth_config = getattr(key, "auth_config", None)
if encrypted_auth_config:
try:
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
token_meta = json.loads(decrypted_config)
except Exception:
token_meta = {}
else:
token_meta = {}
expires_at = token_meta.get("expires_at")
refresh_token = token_meta.get("refresh_token")
provider_type = str(token_meta.get("provider_type") or "")
cached_access_token = str(token_meta.get("access_token") or "").strip()
# 120s skew (or force refresh when upstream returns 401)
should_refresh = False
try:
if expires_at is not None:
should_refresh = int(time.time()) >= int(expires_at) - 120
except Exception:
should_refresh = False
if force_refresh:
should_refresh = True
# Kiro 特殊处理:如果没有缓存的 access_token 或 key.api_key 是占位符,强制刷新
if provider_type == "kiro" and not should_refresh:
if not cached_access_token:
should_refresh = True
elif crypto_service.decrypt(key.api_key) == "__placeholder__":
should_refresh = True
if should_refresh and refresh_token and provider_type:
try:
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
from src.core.provider_templates.types import ProviderType
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
redis, got_lock = await _acquire_refresh_lock(key.id)
if got_lock or redis is None:
try:
if provider_type == ProviderType.KIRO.value:
token_meta = await _refresh_kiro_token(key, endpoint, token_meta)
elif template:
token_meta = await _refresh_generic_oauth_token(
key, endpoint, template, provider_type, refresh_token, token_meta
)
finally:
if got_lock:
await _release_refresh_lock(redis, key.id)
except Exception:
# 刷新失败不阻断请求;后续由上游返回 401 再触发管理端处理
pass
# 获取最终使用的 access_token
# Kiro 优先使用 token_meta 中缓存的 access_token刷新后会更新到 token_meta
if provider_type == "kiro":
refreshed_token = str(token_meta.get("access_token") or "").strip()
effective_token = refreshed_token or crypto_service.decrypt(key.api_key)
else:
effective_token = crypto_service.decrypt(key.api_key)
decrypted_auth_config: dict[str, Any] | None = None
if isinstance(token_meta, dict) and token_meta:
decrypted_auth_config = token_meta
return ProviderAuthInfo(
auth_header="Authorization",
auth_value=f"Bearer {effective_token}",
decrypted_auth_config=decrypted_auth_config,
)
if auth_type == "vertex_ai":
from src.core.vertex_auth import VertexAuthError, VertexAuthService
try:
# 优先从 auth_config 读取,兼容从 api_key 读取(过渡期)
encrypted_auth_config = getattr(key, "auth_config", None)
if encrypted_auth_config:
# auth_config 可能是加密字符串或未加密的 dict
if isinstance(encrypted_auth_config, dict):
# 已经是 dict直接使用兼容未加密存储的情况
sa_json = encrypted_auth_config
else:
# 是加密字符串,需要解密
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
sa_json = json.loads(decrypted_config)
else:
# 兼容旧数据:从 api_key 读取
decrypted_key = crypto_service.decrypt(key.api_key)
# 检查是否是占位符(表示 auth_config 丢失)
if decrypted_key == "__placeholder__":
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
sa_json = json.loads(decrypted_key)
if not isinstance(sa_json, dict):
raise InvalidRequestException("Service Account JSON 无效,请重新添加该密钥。")
# 获取 Access Token
service = VertexAuthService(sa_json)
access_token = await service.get_access_token()
# Vertex AI 使用 Bearer token
return ProviderAuthInfo(
auth_header="Authorization",
auth_value=f"Bearer {access_token}",
decrypted_auth_config=sa_json,
)
except InvalidRequestException:
raise
except VertexAuthError as e:
raise InvalidRequestException(f"Vertex AI 认证失败:{e}")
except json.JSONDecodeError:
raise InvalidRequestException("Service Account JSON 格式无效,请重新添加该密钥。")
except Exception:
raise InvalidRequestException("Vertex AI 认证失败,请检查 Key 的 auth_config")
# 其他认证类型可在此扩展
# elif auth_type == "oauth2":
# ...
# 标准 API Key返回 None由 build_headers 处理
return None

View File

@@ -1,176 +1,19 @@
"""
响应解析器基类 - 定义统一的响应解析接口
响应解析器基类 - re-export from src.core.stream_types
实际定义已下沉到 src/core/stream_types.py此文件保留向后兼容的 re-export。
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any
from src.core.stream_types import (
ParsedChunk,
ParsedResponse,
ResponseParser,
StreamStats,
)
@dataclass
class ParsedChunk:
"""解析后的流式数据块"""
# 原始数据
raw_line: str
event_type: str | None = None
data: dict[str, Any] | None = None
# 提取的内容
text_delta: str = ""
is_done: bool = False
is_error: bool = False
error_message: str | None = None
# 使用量信息(通常在最后一个 chunk 中)
input_tokens: int = 0
output_tokens: int = 0
cache_creation_tokens: int = 0
cache_read_tokens: int = 0
# 响应 ID
response_id: str | None = None
@dataclass
class StreamStats:
"""流式响应统计信息"""
# 计数
chunk_count: int = 0
data_count: int = 0
# Token 使用量
input_tokens: int = 0
output_tokens: int = 0
cache_creation_tokens: int = 0
cache_read_tokens: int = 0
# 内容
collected_text: str = ""
response_id: str | None = None
# 状态
has_completion: bool = False
status_code: int = 200
error_message: str | None = None
# Provider 信息
provider_name: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
# 响应头和完整响应
response_headers: dict[str, str] = field(default_factory=dict)
final_response: dict[str, Any] | None = None
@dataclass
class ParsedResponse:
"""解析后的非流式响应"""
# 原始响应
raw_response: dict[str, Any]
status_code: int
# 提取的内容
text_content: str = ""
response_id: str | None = None
# 使用量
input_tokens: int = 0
output_tokens: int = 0
cache_creation_tokens: int = 0
cache_read_tokens: int = 0
# 错误信息
is_error: bool = False
error_type: str | None = None
error_message: str | None = None
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
embedded_status_code: int | None = None
class ResponseParser(ABC):
"""
响应解析器基类
定义统一的接口来解析不同 API 格式的响应。
子类需要实现具体的解析逻辑。
"""
# 解析器名称(用于日志)
name: str = "base"
# 支持的 API 格式
api_format: str = "UNKNOWN"
@abstractmethod
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
"""
解析单行 SSE 数据
Args:
line: SSE 行数据
stats: 流统计对象(会被更新)
Returns:
解析后的数据块,如果行不包含有效数据则返回 None
"""
pass
@abstractmethod
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
"""
解析非流式响应
Args:
response: 响应 JSON
status_code: HTTP 状态码
Returns:
解析后的响应对象
"""
pass
@abstractmethod
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
"""
从响应中提取 token 使用量
Args:
response: 响应 JSON
Returns:
包含 input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens 的字典
"""
pass
@abstractmethod
def extract_text_content(self, response: dict[str, Any]) -> str:
"""
从响应中提取文本内容
Args:
response: 响应 JSON
Returns:
提取的文本内容
"""
pass
def is_error_response(self, response: dict[str, Any]) -> bool:
"""
判断响应是否为错误响应
Args:
response: 响应 JSON
Returns:
是否为错误响应
"""
return "error" in response
def create_stats(self) -> StreamStats:
"""创建新的流统计对象"""
return StreamStats()
__all__ = [
"ParsedChunk",
"ParsedResponse",
"ResponseParser",
"StreamStats",
]

View File

@@ -6,7 +6,6 @@ Video Handler 基类
from __future__ import annotations
import re
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any
@@ -19,75 +18,20 @@ from src.config.settings import config
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.core.video_utils import (
extract_short_id_from_operation,
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
if TYPE_CHECKING:
import httpx
from src.services.candidate.submit import SubmitOutcome
# 敏感信息匹配正则(预编译提升性能)
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
re.IGNORECASE,
)
def sanitize_error_message(message: str, max_length: int = 200) -> str:
"""
移除错误消息中可能包含的敏感信息
Args:
message: 原始错误消息
max_length: 最大长度,默认 200
Returns:
脱敏后的消息
"""
if not message:
return "Request failed"
# 先脱敏再截断,确保敏感信息不会因截断位置而泄露
sanitized = _SENSITIVE_PATTERN.sub("[REDACTED]", message)
return sanitized[:max_length]
def extract_short_id_from_operation(operation_id: str) -> str:
"""
从 operation ID 中提取短 ID
我们对外暴露的 operation name 格式是:
- models/{model}/operations/{short_id}
此函数提取最后一部分作为 short_id用于在数据库中查找任务。
Args:
operation_id: 原始 operation ID"models/veo-3.1/operations/abc123"
Returns:
short_id"abc123"
"""
# 格式: models/{model}/operations/{short_id}
# 或者直接是 short_id
if "/" in operation_id:
# 提取最后一部分
return operation_id.rsplit("/", 1)[-1]
return operation_id
def normalize_gemini_operation_id(operation_id: str) -> str:
"""
规范化 Gemini operation ID保留用于向后兼容
Args:
operation_id: 原始 operation ID
Returns:
规范化后的 operation ID原样返回
"""
return operation_id
class VideoHandlerBase(ABC):
"""视频处理器基类"""

View File

@@ -7,13 +7,9 @@ Gemini 图像生成模型请求适配
from typing import Any
from src.core.video_utils import is_image_gen_model
def is_image_gen_model(model: str | None) -> bool:
"""判断是否为图像生成模型(模式匹配,覆盖 gemini-*-image / imagen-* 系列)"""
if not model:
return False
m = model.lower()
return "image" in m and ("gemini" in m or "imagen" in m)
__all__ = ["is_image_gen_model", "adapt_request_for_image_gen"]
def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]:

View File

@@ -41,7 +41,7 @@ from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.usage.service import UsageService

View File

@@ -39,7 +39,7 @@ from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.usage.service import UsageService

View File

@@ -36,9 +36,9 @@ from src.core.logger import logger
from src.database import create_session
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
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
@dataclass