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

@@ -1,14 +1,10 @@
"""
缓存服务模块
"""通用缓存模块。
包含缓存后端、缓存亲和性、缓存同步等能。
包含缓存后端、缓存失效与缓存同步等能backend/sync/*_cache
注意:由于循环依赖问题,部分类需要直接从子模块导入:
from src.services.cache.affinity_manager import CacheAffinityManager
from src.services.cache.aware_scheduler import CacheAwareScheduler
调度/候选/缓存亲和性相关逻辑已迁移到 `src.services.scheduling`。
"""
# 只导出不会导致循环依赖的基础类
from src.services.cache.backend import BaseCacheBackend, LocalCache, RedisCache, get_cache_backend
__all__ = [

View File

@@ -54,7 +54,7 @@ class CacheInvalidationService:
logger.error(f"[CacheInvalidation] 失效 ModelCacheService 缓存失败: {e}")
# 4. 清除 /v1/models 列表缓存
from src.api.base.models_service import invalidate_models_list_cache
from src.services.cache.model_list_cache import invalidate_models_list_cache
try:
await invalidate_models_list_cache()
@@ -79,7 +79,7 @@ class CacheInvalidationService:
self._refresh_provider_cache(provider_id)
# 清除 /v1/models 列表缓存allowed_models 变更会影响模型可用性)
from src.api.base.models_service import invalidate_models_list_cache
from src.services.cache.model_list_cache import invalidate_models_list_cache
try:
await invalidate_models_list_cache()

31
src/services/cache/model_list_cache.py vendored Normal file
View File

@@ -0,0 +1,31 @@
"""
/v1/models 列表缓存管理。
从 api/base/models_service.py 迁移到 services 层,
消除 services→api 的反向依赖。
"""
from __future__ import annotations
from src.core.cache_service import CacheService
from src.core.logger import logger
# 缓存 key 前缀models_service.py 也使用此常量)
MODELS_LIST_CACHE_PREFIX = "models:list"
async def invalidate_models_list_cache() -> None:
"""
清除所有 /v1/models 列表缓存
在模型创建、更新、删除时调用,确保模型列表实时更新
"""
try:
# 使用通配符删除所有 models:list:* 缓存(包括多格式组合的 key
deleted = await CacheService.delete_pattern(f"{MODELS_LIST_CACHE_PREFIX}:*")
if deleted > 0:
logger.info("[ModelsService] 已清除 {}{} 缓存", deleted, MODELS_LIST_CACHE_PREFIX)
else:
logger.debug("[ModelsService] 无 {} 缓存需要清除", MODELS_LIST_CACHE_PREFIX)
except Exception as e:
logger.warning("[ModelsService] 清除缓存失败: {}", e)

View File

@@ -13,9 +13,9 @@ from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import RequestCandidate
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
from src.services.request.candidate import RequestCandidateService
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.task.exceptions import StreamProbeError
from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult
from src.services.task.schema import ExecutionResult

View File

@@ -3,7 +3,7 @@ from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
CANDIDATE_KEY_SCHEMA_VERSION = "1.0"

View File

@@ -5,8 +5,8 @@ from typing import Any
from sqlalchemy.orm import Session
from src.models.database import ApiKey
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.scheduling.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
from src.services.system.config import SystemConfigService
from .recorder import CandidateRecorder

View File

@@ -6,7 +6,7 @@ from typing import Any, Protocol, runtime_checkable
import httpx
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
@runtime_checkable

View File

@@ -928,6 +928,7 @@ class ModelCostService:
# 获取对应 API 格式的 Adapter 实例来计算成本
# 优先检查 Chat Adapter然后检查 CLI Adapter
# TODO(arch): 引入 adapter 能力注册表,消除 services->api 依赖
from src.api.handlers.base.chat_adapter_base import get_adapter_instance
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_instance

View File

@@ -10,13 +10,13 @@ from sqlalchemy import and_
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.base.models_service import invalidate_models_list_cache
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.models.api import ModelCreate, ModelResponse, ModelUpdate
from src.models.database import Model, Provider
from src.services.cache.invalidation import get_cache_invalidation_service
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.model_list_cache import invalidate_models_list_cache
class ModelService:

View File

@@ -169,6 +169,7 @@ def merge_upstream_metadata(
def get_adapter_for_format(api_format: str) -> type | None:
"""根据 API 格式获取对应的 Adapter 类"""
# TODO(arch): 引入 adapter 能力注册表,消除 services->api 依赖
from src.api.handlers.base.chat_adapter_base import get_adapter_class
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class

View File

@@ -13,8 +13,8 @@ from sqlalchemy.orm import Session
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import ApiKey
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
from src.services.provider.format import normalize_endpoint_signature
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class CandidateResolver:

View File

@@ -25,9 +25,9 @@ from src.core.exceptions import (
)
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.cache.aware_scheduler import CacheAwareScheduler
from src.services.orchestration.error_handler import ErrorHandlerService
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
class ErrorAction(Enum):

View File

@@ -22,11 +22,11 @@ from src.core.exceptions import (
)
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.cache.aware_scheduler import CacheAwareScheduler
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
class ErrorHandlerService:

View File

@@ -11,9 +11,9 @@ from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
from src.services.request.candidate import RequestCandidateService
from src.services.request.executor import RequestExecutor
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class RequestDispatcher:

View File

@@ -11,7 +11,9 @@ import re
import threading
# ============== API 端点 ==============
PROD_BASE_URL = "https://cloudcode-pa.googleapis.com"
# 唯一定义在 core 层,此处 re-export 保持向后兼容
from src.core.provider_templates.fixed_providers import ANTIGRAVITY_PROD_URL as PROD_BASE_URL
DAILY_BASE_URL = "https://daily-cloudcode-pa.googleapis.com"
SANDBOX_BASE_URL = "https://daily-cloudcode-pa.sandbox.googleapis.com"
@@ -75,8 +77,7 @@ URL_UNAVAILABLE_TTL_SECONDS = 300 # 5 分钟
# ============== Thinking Signature ==============
# 统一从 core 层导入,避免多处定义
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE # noqa: E402
MIN_SIGNATURE_LENGTH = 50 # 与 Antigravity-Manager 对齐
from src.core.api_format.conversion.thinking_cache import MIN_SIGNATURE_LENGTH # noqa: E402, F401
# ============== Thinking Budget ==============
THINKING_BUDGET_AUTO_CAP = 24576

View File

@@ -521,7 +521,7 @@ def _inject_thought_signatures(inner_request: dict[str, Any], session_id: str |
"""
try:
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
from src.core.api_format.conversion.thinking_cache import signature_cache
except Exception:
return
@@ -881,7 +881,7 @@ def wrap_v1internal_request(
13. 注入 sessionId
14. 构建 v1internal 信封
"""
from src.api.handlers.gemini.image_gen import is_image_gen_model
from src.core.video_utils import is_image_gen_model
inner_request = dict(gemini_request)
inner_request.pop("model", None)
@@ -979,7 +979,7 @@ def cache_thought_signatures(model: str, response: dict[str, Any]) -> None:
同时缓存到 legacy (text) 层和 tool (Layer 1) 层。
"""
try:
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
from src.core.api_format.conversion.thinking_cache import signature_cache
except Exception:
return

View File

@@ -1,239 +1,13 @@
"""Antigravity thinking block signature cache (triple-layer).
"""Backward-compatible re-export for Antigravity thinking signature cache.
与 Antigravity-Manager 对齐的三层缓存设计:
Layer 1: tool_use_id → thoughtSignature (工具调用签名恢复)
Layer 2: signature → model_family (跨模型兼容校验)
Layer 3: session_id → latest signature (会话级签名追踪 + rewind 检测)
同时保留原有的 model:text → signature 兼容层。
The implementation moved to `src.core.api_format.conversion.thinking_cache` to eliminate
core → services reverse dependencies.
"""
from __future__ import annotations
import hashlib
import threading
import time
from typing import Any
from src.services.provider.adapters.antigravity.constants import (
DUMMY_THOUGHT_SIGNATURE,
from src.core.api_format.conversion.thinking_cache import (
MIN_SIGNATURE_LENGTH,
ThinkingSignatureCache,
signature_cache,
)
# TTL: 2 小时(与 Antigravity-Manager 对齐)
_SIGNATURE_TTL_SECONDS = 2 * 60 * 60
# 各层缓存上限
_TOOL_CACHE_LIMIT = 500
_FAMILY_CACHE_LIMIT = 200
_SESSION_CACHE_LIMIT = 1000
_TEXT_CACHE_LIMIT = 1000
class _CacheEntry:
"""带时间戳的缓存条目,支持 TTL 过期。"""
__slots__ = ("data", "created_at")
def __init__(self, data: Any) -> None:
self.data = data
self.created_at: float = time.monotonic()
def is_expired(self, now: float | None = None) -> bool:
return ((now or time.monotonic()) - self.created_at) > _SIGNATURE_TTL_SECONDS
class _SessionEntry:
"""Session 层缓存数据,包含消息计数用于 rewind 检测。"""
__slots__ = ("signature", "message_count")
def __init__(self, signature: str, message_count: int) -> None:
self.signature = signature
self.message_count = message_count
class ThinkingSignatureCache:
"""Triple-layer thinking signature cache.
Layer 1 (tool): tool_use_id → thoughtSignature
当客户端(如 OpenCode) 在 tool_result 中丢弃了 signature 时用于恢复。
Layer 2 (family): signature → model_family
防止跨模型签名污染Claude 签名不能用在 Gemini 上)。
Layer 3 (session): session_id → latest signature + message_count
会话级追踪,支持 rewind 检测(用户删除消息后不会注入来自"未来"的签名)。
Legacy (text): SHA256(model + text) → signature
向后兼容的 get_or_dummy() 接口。
"""
def __init__(self) -> None:
self._tool_sigs: dict[str, _CacheEntry] = {}
self._families: dict[str, _CacheEntry] = {}
self._sessions: dict[str, _CacheEntry] = {}
self._text_sigs: dict[str, _CacheEntry] = {}
self._lock = threading.Lock()
# ===== Layer 1: Tool Use ID → Signature =====
def cache_tool_signature(self, tool_use_id: str, signature: str) -> None:
"""缓存工具调用对应的 thinking signature。"""
if len(signature) < MIN_SIGNATURE_LENGTH:
return
with self._lock:
self._tool_sigs[tool_use_id] = _CacheEntry(signature)
if len(self._tool_sigs) > _TOOL_CACHE_LIMIT:
self._prune(self._tool_sigs, limit=_TOOL_CACHE_LIMIT)
def get_tool_signature(self, tool_use_id: str) -> str | None:
"""查找工具调用对应的 signature。"""
with self._lock:
entry = self._tool_sigs.get(tool_use_id)
if entry is None:
return None
if entry.is_expired():
self._tool_sigs.pop(tool_use_id, None)
return None
return entry.data
# ===== Layer 2: Signature → Model Family =====
def cache_thinking_family(self, signature: str, family: str) -> None:
"""记录 signature 所属的模型家族。"""
if len(signature) < MIN_SIGNATURE_LENGTH:
return
with self._lock:
self._families[signature] = _CacheEntry(family)
if len(self._families) > _FAMILY_CACHE_LIMIT:
self._prune(self._families, limit=_FAMILY_CACHE_LIMIT)
def get_signature_family(self, signature: str) -> str | None:
"""查找 signature 所属的模型家族。"""
with self._lock:
entry = self._families.get(signature)
if entry is None:
return None
if entry.is_expired():
self._families.pop(signature, None)
return None
return entry.data
# ===== Layer 3: Session ID → Latest Signature =====
def cache_session_signature(
self, session_id: str, signature: str, message_count: int = 0
) -> None:
"""存储会话的最新 thinking signature。
Rewind 检测:当 message_count 小于已缓存值时,说明用户删除了消息,
强制更新签名以避免注入来自"未来"的签名。
"""
if len(signature) < MIN_SIGNATURE_LENGTH:
return
with self._lock:
existing = self._sessions.get(session_id)
should_store = True
if existing and not existing.is_expired():
entry: _SessionEntry = existing.data
if message_count < entry.message_count:
# Rewind detected: 用户删除了消息,强制更新
pass
elif message_count == entry.message_count:
# 同一轮消息:仅当新签名更长(更完整)时才替换
should_store = len(signature) > len(entry.signature)
# else: 正常递增,更新
if should_store:
self._sessions[session_id] = _CacheEntry(_SessionEntry(signature, message_count))
if len(self._sessions) > _SESSION_CACHE_LIMIT:
self._prune(self._sessions, limit=_SESSION_CACHE_LIMIT)
def get_session_signature(self, session_id: str) -> str | None:
"""获取会话的最新 thinking signature。"""
with self._lock:
entry = self._sessions.get(session_id)
if entry is None:
return None
if entry.is_expired():
self._sessions.pop(session_id, None)
return None
return entry.data.signature
# ===== Legacy: model:text → signature向后兼容 =====
def get_or_dummy(self, model: str, thinking_text: str) -> str | None:
"""Legacy: 根据 model + thinking_text 查找 signature。
Gemini 模型在未命中时返回 DUMMY_THOUGHT_SIGNATURE跳过验证
"""
key = self._text_key(model, thinking_text)
with self._lock:
entry = self._text_sigs.get(key)
if entry is not None:
if entry.is_expired():
self._text_sigs.pop(key, None)
else:
return entry.data
if str(model).startswith("gemini-"):
return DUMMY_THOUGHT_SIGNATURE
return None
def cache(self, model: str, thinking_text: str, signature: str) -> None:
"""Legacy: 缓存 model + thinking_text → signature。"""
if len(signature) < MIN_SIGNATURE_LENGTH:
return
key = self._text_key(model, thinking_text)
with self._lock:
if key in self._text_sigs:
self._text_sigs[key] = _CacheEntry(signature)
return
if len(self._text_sigs) >= _TEXT_CACHE_LIMIT:
# FIFO 淘汰 1/4
evict_n = max(1, _TEXT_CACHE_LIMIT // 4)
for k in list(self._text_sigs.keys())[:evict_n]:
self._text_sigs.pop(k, None)
self._text_sigs[key] = _CacheEntry(signature)
# ===== Utilities =====
@staticmethod
def _text_key(model: str, thinking_text: str) -> str:
# 使用 \x00 作为分隔符避免 model 中含 ':' 时的歧义
content = f"{model}\x00{thinking_text}"
return hashlib.sha256(content.encode("utf-8")).hexdigest()[:32]
@staticmethod
def _prune(d: dict[str, _CacheEntry], *, limit: int | None = None) -> None:
"""Remove expired entries and optionally enforce a size limit."""
now = time.monotonic()
expired = [k for k, v in d.items() if v.is_expired(now)]
for k in expired:
d.pop(k, None)
if limit is None or len(d) <= limit:
return
# Evict oldest entries by created_at.
excess = len(d) - limit
for k, _entry in sorted(d.items(), key=lambda kv: kv[1].created_at)[:excess]:
d.pop(k, None)
def clear(self) -> None:
"""清空所有缓存层(用于测试或手动重置)。"""
with self._lock:
self._tool_sigs.clear()
self._families.clear()
self._sessions.clear()
self._text_sigs.clear()
signature_cache = ThinkingSignatureCache()
__all__ = ["ThinkingSignatureCache", "signature_cache"]
__all__ = ["MIN_SIGNATURE_LENGTH", "ThinkingSignatureCache", "signature_cache"]

View File

@@ -0,0 +1,394 @@
"""
Provider 认证逻辑OAuth / Service Account / Vertex AI
从 api/handlers/base/request_builder.py 迁移到 services 层,
消除 services→api 的反向依赖。
"""
from __future__ import annotations
import json
import time
from typing import TYPE_CHECKING, Any
from sqlalchemy.orm import object_session
from src.clients.redis_client import get_redis_client
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.core.provider_auth_types import ProviderAuthInfo
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
if TYPE_CHECKING:
from src.models.database import ProviderAPIKey, ProviderEndpoint
# ==============================================================================
# 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
# ==============================================================================
# Provider-specific refresh implementations
# ==============================================================================
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注入代理配置core 层不依赖 services
from src.services.proxy_node.resolver import build_proxy_client_kwargs
service = VertexAuthService(sa_json)
access_token = await service.get_access_token(
httpx_client_kwargs=build_proxy_client_kwargs(timeout=30),
)
# 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

@@ -64,7 +64,7 @@ async def resolve_oauth_access_token(
"""
# Local import to avoid circular imports during app startup.
from src.api.handlers.base.request_builder import get_provider_auth
from src.services.provider.auth import get_provider_auth
# Build detached key-like objects for get_provider_auth().
provider_obj = (

View File

@@ -28,8 +28,6 @@ class ConcurrencyManager:
_instance: ConcurrencyManager | None = None
_redis: aioredis.Redis | None = None
_key_rpm_bucket_seconds: int = 60
_key_rpm_key_ttl_seconds: int = 120 # 2 分钟,足够覆盖当前分钟与边界
def __new__(cls) -> "ConcurrencyManager":
"""单例模式"""
@@ -42,13 +40,18 @@ class ConcurrencyManager:
if hasattr(self, "_memory_initialized"):
return
from src.config.settings import config
self._key_rpm_bucket_seconds: int = config.rpm_bucket_seconds
self._key_rpm_key_ttl_seconds: int = config.rpm_key_ttl_seconds
self._memory_lock: asyncio.Lock = asyncio.Lock()
# Key RPM 计数器:{key_id: (bucket, count)}bucket = floor(now / 60)
self._memory_key_rpm_counts: dict[str, tuple[int, int]] = {}
self._owns_redis: bool = False
self._last_cleanup_bucket: int = 0 # 上次清理时的 bucket用于定期清理过期数据
self._last_cleanup_time: float = 0 # 上次清理的时间戳,用于强制定期清理
self._cleanup_interval_seconds: int = 300 # 强制清理间隔5 分钟)
self._cleanup_interval_seconds: int = config.rpm_cleanup_interval_seconds
self._cleanup_task: asyncio.Task | None = None # 后台清理任务
# 内存模式下的最大条目限制,防止内存泄漏(支持环境变量覆盖)
@@ -130,16 +133,14 @@ class ConcurrencyManager:
self._redis = None
self._owns_redis = False
@classmethod
def _get_rpm_bucket(cls, now_ts: float | None = None) -> int:
def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
"""获取当前 RPM 计数桶(按分钟)"""
ts = now_ts if now_ts is not None else time.time()
return int(ts // cls._key_rpm_bucket_seconds)
return int(ts // self._key_rpm_bucket_seconds)
@classmethod
def _get_key_key(cls, key_id: str, bucket: int | None = None) -> str:
def _get_key_key(self, key_id: str, bucket: int | None = None) -> str:
"""获取 ProviderAPIKey RPM 计数的 Redis Key按分钟桶"""
b = bucket if bucket is not None else cls._get_rpm_bucket()
b = bucket if bucket is not None else self._get_rpm_bucket()
return f"rpm:key:{key_id}:{b}"
def _get_memory_key_rpm_count(self, key_id: str, bucket: int) -> int:

View File

@@ -0,0 +1,30 @@
"""调度系统(候选/排序/亲和性/并发检查)。
原 `src.services.cache` 中与调度相关的代码已迁移到此包;
`src.services.cache` 现在只保留通用缓存 backend/sync/*_cache。
"""
from src.services.scheduling.affinity_manager import CacheAffinityManager, get_affinity_manager
from src.services.scheduling.aware_scheduler import (
CacheAwareScheduler,
ConcurrencySnapshot,
ProviderCandidate,
get_cache_aware_scheduler,
)
from src.services.scheduling.candidate_builder import CandidateBuilder
from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.concurrency_checker import ConcurrencyChecker
from src.services.scheduling.scheduling_config import SchedulingConfig
__all__ = [
"CacheAffinityManager",
"CandidateBuilder",
"CandidateSorter",
"CacheAwareScheduler",
"ConcurrencyChecker",
"ConcurrencySnapshot",
"ProviderCandidate",
"SchedulingConfig",
"get_affinity_manager",
"get_cache_aware_scheduler",
]

View File

@@ -31,7 +31,15 @@
from __future__ import annotations
import time
from typing import Any
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from src.services.scheduling.protocols import (
CacheAffinityManagerProtocol,
CandidateBuilderProtocol,
CandidateSorterProtocol,
ConcurrencyCheckerProtocol,
)
from sqlalchemy.orm import Session
@@ -47,32 +55,32 @@ from src.models.database import (
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache.affinity_manager import (
CacheAffinityManager,
get_affinity_manager,
)
from src.services.cache.candidate_builder import (
CandidateBuilder,
)
from src.services.cache.candidate_builder import (
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
)
from src.services.cache.candidate_sorter import CandidateSorter
from src.services.cache.concurrency_checker import ConcurrencyChecker
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.restriction_checker import get_effective_restrictions
from src.services.cache.scheduling_config import SchedulingConfig
from src.services.cache.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
from src.services.cache.schemas import ProviderCandidate as ProviderCandidate # re-export
from src.services.cache.utils import affinity_hash as _affinity_hash # re-export compat
from src.services.cache.utils import (
release_db_connection_before_await,
)
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import (
get_adaptive_reservation_manager,
)
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.scheduling.affinity_manager import (
CacheAffinityManager,
get_affinity_manager,
)
from src.services.scheduling.candidate_builder import (
CandidateBuilder,
)
from src.services.scheduling.candidate_builder import (
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
)
from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.concurrency_checker import ConcurrencyChecker
from src.services.scheduling.restriction_checker import get_effective_restrictions
from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.scheduling.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
from src.services.scheduling.schemas import ProviderCandidate as ProviderCandidate # re-export
from src.services.scheduling.utils import affinity_hash as _affinity_hash # re-export compat
from src.services.scheduling.utils import (
release_db_connection_before_await,
)
from src.services.system.config import SystemConfigService
@@ -102,6 +110,11 @@ class CacheAwareScheduler:
redis_client: Any | None = None,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
*,
candidate_builder: CandidateBuilderProtocol | None = None,
candidate_sorter: CandidateSorterProtocol | None = None,
concurrency_checker: ConcurrencyCheckerProtocol | None = None,
affinity_manager: CacheAffinityManagerProtocol | None = None,
) -> None:
"""
初始化调度器
@@ -117,9 +130,9 @@ class CacheAwareScheduler:
self.redis = redis_client
self._config = SchedulingConfig(priority_mode, scheduling_mode)
# 异步子组件(将在第一次使用时初始化)
self._affinity_manager: CacheAffinityManager | None = None
self._concurrency_checker: ConcurrencyChecker | None = None
# 异步子组件(将在第一次使用时初始化,可通过构造函数注入
self._affinity_manager: CacheAffinityManagerProtocol | None = affinity_manager
self._concurrency_checker: ConcurrencyCheckerProtocol | None = concurrency_checker
self._metrics: dict[str, Any] = {
"total_batches": 0,
@@ -139,9 +152,13 @@ class CacheAwareScheduler:
"last_reservation_result": None,
}
# 初始化子模块(不传 self解除反向引用
self._candidate_sorter = CandidateSorter(self._config)
self._candidate_builder = CandidateBuilder(self._candidate_sorter)
# 初始化子模块(不传 self解除反向引用,可通过构造函数注入
self._candidate_sorter: CandidateSorterProtocol = candidate_sorter or CandidateSorter(
self._config
)
self._candidate_builder: CandidateBuilderProtocol = candidate_builder or CandidateBuilder(
self._candidate_sorter
)
# ── 属性代理(保持外部访问兼容性)──────────────────────────

View File

@@ -28,15 +28,15 @@ from src.models.database import (
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache.quota_skipper import is_key_quota_exhausted
from src.services.cache.utils import release_db_connection_before_await
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature
from src.services.scheduling.quota_skipper import is_key_quota_exhausted
from src.services.scheduling.utils import release_db_connection_before_await
if TYPE_CHECKING:
from src.models.database import GlobalModel
from src.services.cache.candidate_sorter import CandidateSorter
from src.services.cache.schemas import ProviderCandidate
from src.services.scheduling.protocols import CandidateSorterProtocol
from src.services.scheduling.schemas import ProviderCandidate
from src.services.cache.model_cache import ModelCacheService
@@ -60,7 +60,7 @@ def _sort_endpoints_by_family_priority(
class CandidateBuilder:
"""候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。"""
def __init__(self, candidate_sorter: CandidateSorter) -> None:
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
self._sorter = candidate_sorter
def _query_providers(
@@ -380,7 +380,7 @@ class CandidateBuilder:
Returns:
候选列表
"""
from src.services.cache.schemas import ProviderCandidate
from src.services.scheduling.schemas import ProviderCandidate
candidates: list[ProviderCandidate] = []
client_format_str = normalize_endpoint_signature(client_format)

View File

@@ -13,15 +13,15 @@ import random
from collections import defaultdict
from typing import TYPE_CHECKING
from src.services.cache.scheduling_config import SchedulingConfig
from src.services.cache.utils import affinity_hash
from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.scheduling.utils import affinity_hash
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey
from src.services.cache.schemas import ProviderCandidate
from src.services.scheduling.schemas import ProviderCandidate
class CandidateSorter:

View File

@@ -11,9 +11,9 @@ from typing import Any
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.services.cache.schemas import ConcurrencySnapshot
from src.services.rate_limit.adaptive_reservation import AdaptiveReservationManager
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.scheduling.schemas import ConcurrencySnapshot
class ConcurrencyChecker:

View File

@@ -0,0 +1,137 @@
"""调度/候选子组件的协议接口。
目的:
- 用 `Protocol` 固化 CacheAwareScheduler 的子组件契约
- 便于单测注入 stub/mocks减少对具体实现类的耦合
说明:这里的协议面向“调度器内部协作”,因此保留了部分 `_` 前缀方法。
后续如果要对外暴露更稳定的 API可再抽出无下划线的 facade。
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Protocol
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from src.models.database import GlobalModel, Provider, ProviderAPIKey
from src.services.scheduling.affinity_manager import CacheAffinity
from src.services.scheduling.schemas import ConcurrencySnapshot, ProviderCandidate
class CandidateSorterProtocol(Protocol):
def _apply_priority_mode_sort(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]: ...
def _apply_load_balance(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]: ...
def shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]: ...
class CandidateBuilderProtocol(Protocol):
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[Provider]: ...
async def _build_candidates(
self,
db: Session,
providers: list[Provider],
client_format: str,
model_name: str,
affinity_key: str | None,
model_mappings: list[str] | None = None,
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
global_conversion_enabled: bool = True,
) -> list[ProviderCandidate]: ...
async def _check_model_support(
self,
db: Session,
provider: Provider,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]: ...
async def _check_model_support_for_global_model(
self,
db: Session,
provider: Provider,
global_model: GlobalModel,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]: ...
def _check_key_availability(
self,
key: ProviderAPIKey,
api_format: str | None,
model_name: str,
capability_requirements: dict[str, bool] | None = None,
model_mappings: list[str] | None = None,
candidate_models: set[str] | None = None,
*,
provider_type: str | None = None,
) -> tuple[bool, str | None, str | None]: ...
class ConcurrencyCheckerProtocol(Protocol):
async def check_available(
self,
key: ProviderAPIKey,
is_cached_user: bool = False,
) -> tuple[bool, ConcurrencySnapshot]: ...
def get_reservation_stats(self) -> dict[str, Any]: ...
class CacheAffinityManagerProtocol(Protocol):
async def get_affinity(
self, affinity_key: str, api_format: str, model_name: str
) -> CacheAffinity | None: ...
async def set_affinity(
self,
affinity_key: str,
provider_id: str,
endpoint_id: str,
key_id: str,
api_format: str,
model_name: str,
supports_caching: bool = True,
ttl: int | None = None,
) -> None: ...
async def invalidate_affinity(
self,
affinity_key: str,
api_format: str,
model_name: str,
key_id: str | None = None,
provider_id: str | None = None,
endpoint_id: str | None = None,
) -> None: ...
def get_stats(self) -> dict[str, Any]: ...

View File

@@ -72,7 +72,9 @@ class CacheWarmupService:
"""预热管理员仪表盘统计缓存"""
db = None
try:
from src.api.dashboard.routes import AdminDashboardStatsAdapter
from src.api.dashboard.routes import ( # TODO(arch): 提取 dashboard 统计计算到 services 层
AdminDashboardStatsAdapter,
)
from src.models.database import User as DBUser
db = create_session()
@@ -128,7 +130,9 @@ class CacheWarmupService:
"""预热每日统计缓存"""
db = None
try:
from src.api.dashboard.routes import DashboardDailyStatsAdapter
from src.api.dashboard.routes import ( # TODO(arch): 提取 dashboard 统计计算到 services 层
DashboardDailyStatsAdapter,
)
from src.models.database import User as DBUser
db = create_session()

View File

@@ -16,11 +16,6 @@ from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
from src.api.handlers.base.video_handler_base import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.clients.http_client import HTTPClientPool
from src.config.settings import config
from src.core.api_format import (
@@ -33,8 +28,14 @@ from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.core.provider_auth_types import ProviderAuthInfo
from src.core.video_utils import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.provider.auth import get_provider_auth
from src.services.task.service import TaskService

View File

@@ -6,7 +6,7 @@ from typing import Any, AsyncIterator, Protocol, runtime_checkable
import httpx
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.scheduling.aware_scheduler import ProviderCandidate
class AttemptKind(str, Enum):

View File

@@ -3,8 +3,8 @@ from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.candidate.schema import CandidateKey
from src.services.scheduling.aware_scheduler import ProviderCandidate
from .protocol import AttemptKind, AttemptResult

View File

@@ -30,10 +30,6 @@ from src.models.database import (
User,
VideoTask,
)
from src.services.cache.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryPolicy, SkipPolicy
from src.services.candidate.recorder import CandidateRecorder
@@ -43,6 +39,10 @@ from src.services.orchestration.request_dispatcher import RequestDispatcher
from src.services.provider.format import normalize_endpoint_signature
from src.services.request.candidate import RequestCandidateService
from src.services.request.result import RequestMetadata
from src.services.scheduling.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.system.config import SystemConfigService
from src.services.task.context import TaskMode
from src.services.task.exceptions import TaskNotFoundError
@@ -1560,7 +1560,6 @@ class TaskService:
from fastapi import HTTPException
from src.api.handlers.base.request_builder import get_provider_auth
from src.clients.http_client import HTTPClientPool
from src.core.api_format import (
build_upstream_headers_for_endpoint,
@@ -1569,6 +1568,7 @@ class TaskService:
)
from src.core.api_format.conversion.internal_video import VideoStatus
from src.core.crypto import crypto_service
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
try:

View File

@@ -28,7 +28,7 @@ from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.utils import filter_proxy_response_headers
from src.core.api_format import filter_response_headers as filter_proxy_response_headers
from src.core.logger import logger
from src.models.database import ApiKey, User
from src.services.request.result import RequestResult

View File

@@ -13,10 +13,9 @@ from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.response_parser import StreamStats
from src.core.exceptions import EmptyStreamException
from src.core.logger import logger
from src.core.stream_types import StreamStats, get_parser_for_format
from src.database.database import create_session
from src.models.database import ApiKey, User
from src.services.usage.service import UsageService

View File

@@ -0,0 +1,299 @@
"""
消息遥测记录器。
从 api/handlers/base/base_handler.py 迁移到 services 层,
消除 services→api 的反向依赖。
"""
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.services.system.audit import audit_service
from src.services.usage.service import UsageService
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(
"[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,
)

View File

@@ -8,11 +8,11 @@ import json
from abc import ABC, abstractmethod
from typing import Any
from src.api.handlers.base.base_handler import MessageTelemetry
from src.clients.redis_client import get_redis_client
from src.config.settings import config
from src.core.logger import logger
from src.services.usage.events import UsageEventType, build_usage_event
from src.services.usage.telemetry import MessageTelemetry
class TelemetryWriter(ABC):

View File

@@ -416,7 +416,7 @@ class UserService:
通过 GlobalModel + Model 关联查询用户可用模型
逻辑:使用 AccessRestrictions 统一处理 allowed_providers 和 allowed_models 限制
"""
from src.api.base.models_service import AccessRestrictions
from src.core.access_restrictions import AccessRestrictions
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)