refactor: 全局适配 ApiFamily/EndpointKind 结构化标识体系

将新的 (ApiFamily, EndpointKind) / `family:kind` 签名体系应用到整个代码库:
- API Handlers: 所有 adapter/handler 使用新的签名格式
- Services: provider, model, usage, cache, auth 等服务层适配
- Database: ProviderEndpoint 新增 api_family/endpoint_kind 字段
- Frontend: Provider 管理、Usage 表格等组件适配
- Tests: 更新所有相关测试用例
This commit is contained in:
fawney19
2026-02-01 17:28:00 +08:00
parent c246ccfc91
commit 7b66505634
219 changed files with 4732 additions and 2545 deletions

View File

@@ -11,7 +11,6 @@ from datetime import datetime, timezone
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
# 安全策略配置:当 Redis 不可用时的行为
# True = fail-closed安全优先拒绝访问
# False = fail-open可用性优先允许访问
@@ -71,7 +70,12 @@ class JWTBlacklistService:
await redis_client.setex(redis_key, ttl_seconds, reason)
token_fp = JWTBlacklistService._get_token_hash(token)[:12]
logger.info("Token 已加入黑名单: token_fp={} (原因: {}, TTL: {}s)", token_fp, reason, ttl_seconds)
logger.info(
"Token 已加入黑名单: token_fp={} (原因: {}, TTL: {}s)",
token_fp,
reason,
ttl_seconds,
)
return True
except Exception as e:
@@ -94,7 +98,9 @@ class JWTBlacklistService:
if redis_client is None:
# Redis 不可用时,根据安全策略决定行为
if BLACKLIST_FAIL_CLOSED:
logger.warning("Redis 不可用,采用 fail-closed 策略拒绝访问(可通过 JWT_BLACKLIST_FAIL_CLOSED=false 改变)")
logger.warning(
"Redis 不可用,采用 fail-closed 策略拒绝访问(可通过 JWT_BLACKLIST_FAIL_CLOSED=false 改变)"
)
return True # 返回 True 表示在黑名单中,拒绝访问
else:
logger.debug("Redis 不可用,采用 fail-open 策略允许访问")

View File

@@ -172,7 +172,9 @@ class LDAPService:
}
@staticmethod
def authenticate_with_config(config: dict[str, Any], username: str, password: str) -> dict | None:
def authenticate_with_config(
config: dict[str, Any], username: str, password: str
) -> dict | None:
"""
LDAP bind 验证
@@ -186,7 +188,7 @@ class LDAPService:
"""
try:
import ldap3
from ldap3 import Server, Connection, SUBTREE
from ldap3 import SUBTREE, Connection, Server
from ldap3.core.exceptions import LDAPBindError, LDAPSocketOpenError
except ImportError:
logger.error("ldap3 库未安装")
@@ -272,9 +274,7 @@ class LDAPService:
# 提取用户属性(优先用 LDAP 提供的值,不合法则回退默认)
ldap_username = _get_attr_value(user_entry, config["username_attr"], username)
email = _get_attr_value(
user_entry, config["email_attr"], f"{username}@ldap.local"
)
email = _get_attr_value(user_entry, config["email_attr"], f"{username}@ldap.local")
display_name = _get_attr_value(user_entry, config["display_name_attr"], username)
logger.info(f"LDAP 认证成功: {username}")
@@ -315,7 +315,7 @@ class LDAPService:
"""
try:
import ldap3
from ldap3 import Server, Connection
from ldap3 import Connection, Server
except ImportError:
return False, "ldap3 库未安装"

View File

@@ -3,4 +3,3 @@
from .service import OAuthService
__all__ = ["OAuthService"]

View File

@@ -98,7 +98,9 @@ class OAuthProviderBase(ABC):
timeout_seconds: float = 5.0,
headers: dict[str, str] | None = None,
) -> httpx.Response:
async with httpx.AsyncClient(timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()) as client:
async with httpx.AsyncClient(
timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()
) as client:
return await client.post(url, data=data, headers=headers)
async def _http_get(
@@ -108,5 +110,7 @@ class OAuthProviderBase(ABC):
timeout_seconds: float = 5.0,
headers: dict[str, str] | None = None,
) -> httpx.Response:
async with httpx.AsyncClient(timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()) as client:
async with httpx.AsyncClient(
timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()
) as client:
return await client.get(url, headers=headers)

View File

@@ -29,4 +29,3 @@ class OAuthFlowError(Exception):
super().__init__(error_code)
self.error_code = error_code
self.detail = detail

View File

@@ -3,4 +3,3 @@
from .linuxdo import LinuxDoOAuthProvider
__all__ = ["LinuxDoOAuthProvider"]

View File

@@ -93,4 +93,3 @@ def get_oauth_provider_registry() -> OAuthProviderRegistry:
if _registry is None:
_registry = OAuthProviderRegistry()
return _registry

View File

@@ -17,13 +17,13 @@ from src.core.logger import logger
from src.core.modules import get_module_registry
from src.models.database import OAuthProvider, User, UserOAuthLink
from src.services.auth.ldap import LDAPService
from src.services.auth.oauth.base import OAuthProviderBase
from src.services.auth.oauth.models import OAuthFlowError, OAuthUserInfo
from src.services.auth.oauth.registry import get_oauth_provider_registry
from src.services.auth.oauth.state import consume_oauth_state, create_oauth_state
from src.services.auth.service import AuthService
from src.services.cache.user_cache import UserCacheService
from src.services.system.config import SystemConfigService
from src.services.auth.oauth.base import OAuthProviderBase
from src.utils.ssl_utils import get_ssl_context
@@ -34,7 +34,9 @@ class OAuthService:
def _require_module_active(db: Session) -> None:
registry = get_module_registry()
if not registry.is_active("oauth", db):
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="OAuth 模块未启用")
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="OAuth 模块未启用"
)
@staticmethod
def _get_provider_impl(provider_type: str) -> OAuthProviderBase | None:
@@ -138,7 +140,9 @@ class OAuthService:
if redis is None:
raise HTTPException(status_code=503, detail="Redis 不可用")
state = await create_oauth_state(redis, provider_type=provider_type, action="login", user_id=None)
state = await create_oauth_state(
redis, provider_type=provider_type, action="login", user_id=None
)
return provider.get_authorization_url(config, state)
@staticmethod
@@ -163,7 +167,9 @@ class OAuthService:
if redis is None:
raise HTTPException(status_code=503, detail="Redis 不可用")
state = await create_oauth_state(redis, provider_type=provider_type, action="bind", user_id=user.id)
state = await create_oauth_state(
redis, provider_type=provider_type, action="bind", user_id=user.id
)
return provider.get_authorization_url(config, state)
@staticmethod
@@ -319,7 +325,9 @@ class OAuthService:
if state_data.action == "bind":
try:
await OAuthService._handle_bind(db, user_id=state_data.user_id or "", config=config, oauth_user=oauth_user)
await OAuthService._handle_bind(
db, user_id=state_data.user_id or "", config=config, oauth_user=oauth_user
)
except OAuthFlowError as exc:
return OAuthService._build_frontend_error_redirect(
frontend_callback_url, error_code=exc.error_code, error_detail=exc.detail
@@ -348,7 +356,10 @@ class OAuthService:
}
)
refresh_token = AuthService.create_refresh_token(
data={"user_id": user.id, "created_at": user.created_at.isoformat() if user.created_at else None}
data={
"user_id": user.id,
"created_at": user.created_at.isoformat() if user.created_at else None,
}
)
return OAuthService._build_frontend_login_success_redirect(
@@ -356,7 +367,9 @@ class OAuthService:
)
@staticmethod
async def _handle_login(db: Session, *, config: OAuthProvider, oauth_user: OAuthUserInfo) -> User:
async def _handle_login(
db: Session, *, config: OAuthProvider, oauth_user: OAuthUserInfo
) -> User:
now = datetime.now(timezone.utc)
# 1) 已绑定账号:直接登录
@@ -381,7 +394,9 @@ class OAuthService:
return linked_user
# 2) 未绑定账号:可能需要新建用户(受注册开关控制)
enable_registration = SystemConfigService.get_config(db, "enable_registration", default=False)
enable_registration = SystemConfigService.get_config(
db, "enable_registration", default=False
)
if not enable_registration:
raise OAuthFlowError("registration_disabled")
@@ -399,7 +414,11 @@ class OAuthService:
raise OAuthFlowError("email_is_ldap")
raise OAuthFlowError("email_is_oauth")
base_username = oauth_user.username or (email.split("@", 1)[0] if email else None) or f"user_{uuid.uuid4().hex[:8]}"
base_username = (
oauth_user.username
or (email.split("@", 1)[0] if email else None)
or f"user_{uuid.uuid4().hex[:8]}"
)
default_quota = SystemConfigService.get_config(db, "default_user_quota_usd", default=10.0)
# 生成唯一用户名 + 创建用户(简单重试)
@@ -469,7 +488,9 @@ class OAuthService:
existing_link.last_login_at = now
db.commit()
assert existing_user.id is not None
await UserCacheService.invalidate_user_cache(existing_user.id, existing_user.email)
await UserCacheService.invalidate_user_cache(
existing_user.id, existing_user.email
)
return existing_user
raise OAuthFlowError("oauth_already_bound")
raise OAuthFlowError("provider_error", "link_create_failed")
@@ -724,7 +745,11 @@ class OAuthService:
OAuthService._validate_redirect_uri(data.redirect_uri)
# 覆盖端点:必须 https 且 hostname 命中 provider 白名单
for field_name in ("authorization_url_override", "token_url_override", "userinfo_url_override"):
for field_name in (
"authorization_url_override",
"token_url_override",
"userinfo_url_override",
):
value = getattr(data, field_name)
if value:
OAuthService._validate_url_override(provider, value)
@@ -784,7 +809,9 @@ class OAuthService:
async def _reachable(url: str) -> bool:
try:
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()) as client:
async with httpx.AsyncClient(
timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()
) as client:
await client.get(url)
return True
except Exception:
@@ -799,7 +826,9 @@ class OAuthService:
if cfg.client_secret_encrypted:
# 使用无效 code 做一次 token 请求(仅做粗略判定)
try:
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), verify=get_ssl_context()) as client:
async with httpx.AsyncClient(
timeout=httpx.Timeout(5.0), verify=get_ssl_context()
) as client:
resp = await client.post(
token_url,
data={
@@ -859,7 +888,9 @@ class OAuthService:
async def _reachable(url: str) -> bool:
try:
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()) as client:
async with httpx.AsyncClient(
timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()
) as client:
await client.get(url)
return True
except Exception:
@@ -874,7 +905,9 @@ class OAuthService:
if client_secret:
# 使用无效 code 做一次 token 请求(仅做粗略判定)
try:
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), verify=get_ssl_context()) as client:
async with httpx.AsyncClient(
timeout=httpx.Timeout(5.0), verify=get_ssl_context()
) as client:
resp = await client.post(
token_url,
data={
@@ -926,7 +959,10 @@ class OAuthService:
if not link:
raise InvalidRequestException("未绑定该 Provider")
total_links = db.query(func.count(UserOAuthLink.id)).filter(UserOAuthLink.user_id == user.id).scalar() or 0
total_links = (
db.query(func.count(UserOAuthLink.id)).filter(UserOAuthLink.user_id == user.id).scalar()
or 0
)
if user.auth_source == AuthSource.OAUTH and total_links <= 1:
raise InvalidRequestException("OAUTH 用户必须至少保留一个 OAuth 绑定")
@@ -935,7 +971,11 @@ class OAuthService:
if user.auth_source == AuthSource.LOCAL and not user.password_hash and total_links <= 1:
raise InvalidRequestException("请先设置密码后再解绑")
if LDAPService.is_ldap_exclusive(db) and user.auth_source == AuthSource.LOCAL and user.role != UserRole.ADMIN:
if (
LDAPService.is_ldap_exclusive(db)
and user.auth_source == AuthSource.LOCAL
and user.role != UserRole.ADMIN
):
# ldap_exclusive=true 时,普通本地用户解绑最后一个 OAuth 会锁死(密码登录被禁用)
if total_links <= 1:
raise InvalidRequestException("当前处于 LDAP 专属模式,解绑后将无法登录")

View File

@@ -3,9 +3,9 @@ from __future__ import annotations
import json
import secrets
import time
from collections.abc import Awaitable
from dataclasses import dataclass
from typing import Any, cast
from collections.abc import Awaitable
from redis.asyncio import Redis

View File

@@ -20,20 +20,20 @@ from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session, joinedload
from src.config import config
from src.core.logger import logger
from src.core.enums import AuthSource
from src.core.exceptions import ForbiddenException
from src.core.logger import logger
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
from src.models.database import ManagementToken
from src.models.database import ApiKey, User, UserRole
from src.services.auth.jwt_blacklist import JWTBlacklistService
from src.services.auth.ldap import LDAPService
from src.services.cache.user_cache import UserCacheService
from src.services.user.apikey import ApiKeyService
# API Key last_used_at 更新节流配置
# 同一个 API Key 在此时间间隔内只会更新一次 last_used_at
_LAST_USED_UPDATE_INTERVAL = 60 # 秒
@@ -243,9 +243,8 @@ class AuthService:
# 登录校验必须读取密码哈希,不能使用不包含 password_hash 的缓存对象
# 支持邮箱或用户名登录
from sqlalchemy import or_
user = db.query(User).filter(
or_(User.email == email, User.username == email)
).first()
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
if not user:
logger.warning(f"登录失败 - 用户不存在: {email}")
@@ -295,7 +294,9 @@ class AuthService:
注意:使用 with_for_update() 防止并发首次登录创建重复用户
"""
ldap_dn = (ldap_user.get("ldap_dn") or "").strip() or None
ldap_username = (ldap_user.get("ldap_username") or ldap_user.get("username") or "").strip() or None
ldap_username = (
ldap_user.get("ldap_username") or ldap_user.get("username") or ""
).strip() or None
email = ldap_user["email"]
# 优先用稳定标识查找,避免邮箱变更/用户名冲突导致重复建号
@@ -333,11 +334,7 @@ class AuthService:
# 同步邮箱LDAP 侧邮箱变更时更新;若新邮箱已被占用则拒绝)
if user.email != email:
email_taken = (
db.query(User)
.filter(User.email == email, User.id != user.id)
.first()
)
email_taken = db.query(User).filter(User.email == email, User.id != user.id).first()
if email_taken:
logger.warning(f"LDAP 登录拒绝 - 新邮箱已被占用: {email}")
return None
@@ -370,7 +367,9 @@ class AuthService:
logger.info(f"LDAP 用户名冲突,使用新用户名: {ldap_user['username']} -> {username}")
# 读取系统配置的默认配额
default_quota = SystemConfigService.get_config(db, "default_user_quota_usd", default=10.0)
default_quota = SystemConfigService.get_config(
db, "default_user_quota_usd", default=10.0
)
# 创建新用户
user = User(
@@ -408,7 +407,9 @@ class AuthService:
logger.error(f"LDAP 用户创建失败(用户名冲突重试耗尽): {username}")
return None
username = f"{base_username}_ldap_{int(time.time())}{uuid.uuid4().hex[:4]}"
logger.warning(f"LDAP 用户创建用户名冲突,重试 ({attempt + 1}/{max_retries}): {username}")
logger.warning(
f"LDAP 用户创建用户名冲突,重试 ({attempt + 1}/{max_retries}): {username}"
)
else:
# 其他约束冲突,不重试
logger.error(f"LDAP 用户创建失败 - 未知数据库约束冲突: {e}")
@@ -458,8 +459,10 @@ class AuthService:
if not is_balance_ok:
# 获取剩余余额用于日志
remaining_balance = ApiKeyService.get_remaining_balance(key_record)
logger.warning(f"API认证失败 - 余额不足 "
f"(已用: ${key_record.balance_used_usd:.4f}, 剩余: ${remaining_balance:.4f})")
logger.warning(
f"API认证失败 - 余额不足 "
f"(已用: ${key_record.balance_used_usd:.4f}, 剩余: ${remaining_balance:.4f})"
)
return None
# 获取用户
@@ -492,7 +495,9 @@ class AuthService:
# 检查美元配额
if user.used_usd + estimated_cost > user.quota_usd:
logger.warning(f"用户配额不足: {user.email} (已用: ${user.used_usd:.2f}, 配额: ${user.quota_usd:.2f})")
logger.warning(
f"用户配额不足: {user.email} (已用: ${user.used_usd:.2f}, 配额: ${user.quota_usd:.2f})"
)
return False
return True
@@ -509,7 +514,9 @@ class AuthService:
if role_rank.get(user.role, -1) >= role_rank.get(required_role, 999):
return True
logger.warning(f"权限不足: 用户 {user.email} 角色 {user.role.value} < 需要 {required_role.value}")
logger.warning(
f"权限不足: 用户 {user.email} 角色 {user.role.value} < 需要 {required_role.value}"
)
return False
@staticmethod
@@ -668,9 +675,7 @@ class AuthService:
# 检查 IP 白名单
if not token_record.is_ip_allowed(client_ip):
logger.warning(
f"Management Token IP 限制 - Token: {token_record.id}, IP: {client_ip}"
)
logger.warning(f"Management Token IP 限制 - Token: {token_record.id}, IP: {client_ip}")
AuditService.log_event(
db=db,
event_type=AuditEventType.MANAGEMENT_TOKEN_IP_BLOCKED,

View File

@@ -31,7 +31,6 @@ from src.config.constants import CacheTTL
from src.core.logger import logger
class CacheAffinity(NamedTuple):
"""缓存亲和性信息"""
@@ -78,7 +77,9 @@ class CacheAffinityManager:
# 默认缓存TTL- 使用统一常量
DEFAULT_CACHE_TTL = CacheTTL.CACHE_AFFINITY
def __init__(self, redis_client: Any | None = None, default_ttl: int = DEFAULT_CACHE_TTL) -> None:
def __init__(
self, redis_client: Any | None = None, default_ttl: int = DEFAULT_CACHE_TTL
) -> None:
"""
初始化缓存亲和性管理器
@@ -114,7 +115,9 @@ class CacheAffinityManager:
if self.redis:
logger.debug("CacheAffinityManager: 使用Redis存储")
else:
logger.debug("CacheAffinityManager: Redis不可用回退到内存存储(仅适用于单实例/开发环境)")
logger.debug(
"CacheAffinityManager: Redis不可用回退到内存存储(仅适用于单实例/开发环境)"
)
def _is_memory_backend(self) -> bool:
"""是否处于内存模式"""
@@ -172,8 +175,7 @@ class CacheAffinityManager:
清理的条目数量
"""
expired_keys = [
key for key, (expire_at, _) in self._l1_cache.items()
if current_time > expire_at
key for key, (expire_at, _) in self._l1_cache.items() if current_time > expire_at
]
for key in expired_keys:
self._l1_cache.pop(key, None)
@@ -181,8 +183,7 @@ class CacheAffinityManager:
# 如果缓存仍然过大,按过期时间排序移除最旧的条目
if len(self._l1_cache) > self._l1_max_size:
sorted_items = sorted(
self._l1_cache.items(),
key=lambda x: x[1][0] # 按 expire_at 排序
self._l1_cache.items(), key=lambda x: x[1][0] # 按 expire_at 排序
)
# 移除最旧的 20% 条目
remove_count = len(self._l1_cache) - int(self._l1_max_size * 0.8)
@@ -191,7 +192,9 @@ class CacheAffinityManager:
expired_keys.extend([k for k, _ in sorted_items[:remove_count]])
if expired_keys:
logger.debug(f"L1 缓存清理: 移除 {len(expired_keys)} 个条目,当前 {len(self._l1_cache)}")
logger.debug(
f"L1 缓存清理: 移除 {len(expired_keys)} 个条目,当前 {len(self._l1_cache)}"
)
return len(expired_keys)
@@ -371,16 +374,20 @@ class CacheAffinityManager:
or existing_affinity.key_id != key_id
):
self._stats["key_switches"] += 1
logger.debug(f"Key {affinity_key[:8]}... 在 {api_format} 格式下切换后端: "
logger.debug(
f"Key {affinity_key[:8]}... 在 {api_format} 格式下切换后端: "
f"[{existing_affinity.provider_id[:8]}.../{existing_affinity.endpoint_id[:8]}.../"
f"{existing_affinity.key_id[:8]}...] → "
f"[{provider_id[:8]}.../{endpoint_id[:8]}.../{key_id[:8]}...], 重置计数器")
f"[{provider_id[:8]}.../{endpoint_id[:8]}.../{key_id[:8]}...], 重置计数器"
)
created_at = current_time
request_count = 1
else:
logger.debug(f"刷新缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
logger.debug(
f"刷新缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
f"provider={provider_id[:8]}..., endpoint={endpoint_id[:8]}..., "
f"provider_key={key_id[:8]}..., ttl+={ttl}s")
f"provider_key={key_id[:8]}..., ttl+={ttl}s"
)
else:
created_at = current_time
request_count = 1
@@ -399,9 +406,11 @@ class CacheAffinityManager:
await self._save_affinity_dict(cache_key, ttl, affinity_dict)
logger.debug(f"设置缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
logger.debug(
f"设置缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, provider={provider_id[:8]}..., endpoint={endpoint_id[:8]}..., "
f"provider_key={key_id[:8]}..., ttl={ttl}s")
f"provider_key={key_id[:8]}..., ttl={ttl}s"
)
except Exception as e:
logger.exception(f"设置缓存亲和性失败: {e}")
@@ -443,8 +452,10 @@ class CacheAffinityManager:
should_invalidate = False
if not should_invalidate:
logger.debug(f"跳过失效: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, 过滤条件不匹配 (key={key_id}, provider={provider_id}, endpoint={endpoint_id})")
logger.debug(
f"跳过失效: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, 过滤条件不匹配 (key={key_id}, provider={provider_id}, endpoint={endpoint_id})"
)
return
try:
@@ -454,10 +465,12 @@ class CacheAffinityManager:
self._stats["cache_invalidations"] += 1
logger.debug(f"失效缓存亲和性: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
logger.debug(
f"失效缓存亲和性: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
f"model={model_name}, provider={existing_affinity.provider_id[:8]}..., "
f"endpoint={existing_affinity.endpoint_id[:8]}..., "
f"provider_key={existing_affinity.key_id[:8]}...")
f"provider_key={existing_affinity.key_id[:8]}..."
)
except Exception as e:
logger.exception(f"删除缓存亲和性失败: {e}")
@@ -493,8 +506,10 @@ class CacheAffinityManager:
self._stats["cache_invalidations"] += 1
if invalidated_count > 0:
logger.debug(f"批量失效Provider缓存亲和性: provider={provider_id[:8]}..., "
f"失效数量={invalidated_count}")
logger.debug(
f"批量失效Provider缓存亲和性: provider={provider_id[:8]}..., "
f"失效数量={invalidated_count}"
)
return invalidated_count

View File

@@ -30,7 +30,6 @@
from __future__ import annotations
import hashlib
import random
import re
@@ -40,7 +39,8 @@ from typing import TYPE_CHECKING, Any
from sqlalchemy.orm import Session, selectinload
from src.core.api_format import APIFormat
from src.core.api_format.enums import EndpointKind
from src.core.api_format.signature import make_signature_key, parse_signature_key
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import (
@@ -60,7 +60,7 @@ from src.services.cache.affinity_manager import (
)
from src.services.cache.model_cache import ModelCacheService
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_api_format
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import (
AdaptiveReservationManager,
get_adaptive_reservation_manager,
@@ -194,7 +194,7 @@ class CacheAwareScheduler:
self,
db: Session,
affinity_key: str,
api_format: str | APIFormat,
api_format: str,
model_name: str,
excluded_endpoints: list[str] | None = None,
excluded_keys: list[str] | None = None,
@@ -222,14 +222,14 @@ class CacheAwareScheduler:
excluded_endpoints_set = set(excluded_endpoints or [])
excluded_keys_set = set(excluded_keys or [])
normalized_format = normalize_api_format(api_format)
normalized_format = normalize_endpoint_signature(api_format)
logger.debug(
f"[CacheAwareScheduler] select_with_cache_affinity: "
f"affinity_key={affinity_key[:8]}..., api_format={normalized_format.value}, model={model_name}"
f"affinity_key={affinity_key[:8]}..., api_format={normalized_format}, model={model_name}"
)
self._metrics["last_api_format"] = normalized_format.value
self._metrics["last_api_format"] = normalized_format
self._metrics["last_model_name"] = model_name
provider_offset = 0
@@ -290,22 +290,17 @@ class CacheAwareScheduler:
f"并发状态[{snapshot.describe()}]"
)
if key.cache_ttl_minutes > 0 and global_model_id:
ttl = key.cache_ttl_minutes * 60 if key.cache_ttl_minutes > 0 else None
api_format_str = (
normalized_format.value
if isinstance(normalized_format, APIFormat)
else normalized_format
)
await self.set_cache_affinity(
affinity_key=affinity_key,
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
api_format=api_format_str,
global_model_id=global_model_id,
ttl=int(ttl) if ttl is not None else None,
)
if key.cache_ttl_minutes > 0 and global_model_id:
ttl = key.cache_ttl_minutes * 60 if key.cache_ttl_minutes > 0 else None
await self.set_cache_affinity(
affinity_key=affinity_key,
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
api_format=normalized_format,
global_model_id=global_model_id,
ttl=int(ttl) if ttl is not None else None,
)
if is_cached_user:
self._metrics["cache_hits"] += 1
@@ -549,7 +544,7 @@ class CacheAwareScheduler:
async def list_all_candidates(
self,
db: Session,
api_format: str | APIFormat,
api_format: str,
model_name: str,
affinity_key: str | None = None,
user_api_key: ApiKey | None = None,
@@ -584,7 +579,13 @@ class CacheAwareScheduler:
"""
await self._ensure_initialized()
target_format = normalize_api_format(api_format)
target_format = normalize_endpoint_signature(api_format)
logger.debug(
"[Scheduler] list_all_candidates: model=%s, api_format=%s",
model_name,
target_format,
)
# 0. 解析 model_name 到 GlobalModel支持直接匹配和映射名匹配使用 ModelCacheService
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
@@ -595,6 +596,12 @@ class CacheAwareScheduler:
logger.warning(f"GlobalModel not found: {model_name}")
raise ModelNotSupportedException(model=model_name)
logger.debug(
"[Scheduler] GlobalModel resolved: id=%s, name=%s",
global_model.id,
global_model.name,
)
# 使用 GlobalModel.id 作为缓存亲和性的模型标识,确保映射名和规范名都能命中同一个缓存
global_model_id: str = str(global_model.id)
@@ -613,11 +620,10 @@ class CacheAwareScheduler:
# 0.1 检查 API 格式是否被允许
if allowed_api_formats is not None:
# 统一转为大写比较,兼容数据库中存储的大小写
allowed_upper = {f.upper() for f in allowed_api_formats}
if target_format.value.upper() not in allowed_upper:
allowed_norm = {normalize_endpoint_signature(f) for f in allowed_api_formats if f}
if target_format not in allowed_norm:
logger.debug(
f"API Key {user_api_key.id[:8] if user_api_key else 'N/A'}... 不允许使用 API 格式 {target_format.value}, "
f"API Key {user_api_key.id[:8] if user_api_key else 'N/A'}... 不允许使用 API 格式 {target_format}, "
f"允许的格式: {allowed_api_formats}"
)
return [], global_model_id
@@ -642,6 +648,20 @@ class CacheAwareScheduler:
provider_limit=provider_limit,
)
logger.debug(
"[Scheduler] Found %d active providers",
len(providers),
)
for p in providers:
logger.debug(
"[Scheduler] Provider: id=%s, name=%s, is_active=%s, endpoints=%d, models=%d",
p.id[:8] if p.id else "N/A",
p.name,
p.is_active,
len(p.endpoints) if p.endpoints else 0,
len(p.models) if p.models else 0,
)
if not providers:
return [], global_model_id
@@ -660,6 +680,7 @@ class CacheAwareScheduler:
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤)
from src.config.settings import config
global_conversion_enabled = config.format_conversion_enabled
candidates = await self._build_candidates(
db=db,
@@ -675,7 +696,7 @@ class CacheAwareScheduler:
)
# 3. 应用优先级模式排序
candidates = self._apply_priority_mode_sort(candidates, affinity_key, target_format.value)
candidates = self._apply_priority_mode_sort(candidates, affinity_key, target_format)
# 更新指标
self._metrics["total_candidates"] += len(candidates)
@@ -683,7 +704,7 @@ class CacheAwareScheduler:
logger.debug(
f"预先获取到 {len(candidates)} 个可用组合 "
f"(api_format={target_format.value}, model={model_name})"
f"(api_format={target_format}, model={model_name})"
)
# 4. 根据调度模式应用不同的排序策略
@@ -698,7 +719,7 @@ class CacheAwareScheduler:
)
elif self.scheduling_mode == self.SCHEDULING_MODE_LOAD_BALANCE:
# 负载均衡模式:忽略缓存,同优先级内随机轮换
candidates = self._apply_load_balance(candidates, target_format.value)
candidates = self._apply_load_balance(candidates, target_format)
for candidate in candidates:
candidate.is_cached = False
else:
@@ -877,11 +898,14 @@ class CacheAwareScheduler:
mapping_api_formats = raw.get("api_formats")
if api_format and mapping_api_formats:
if (
isinstance(mapping_api_formats, list)
and api_format not in mapping_api_formats
):
continue
# 新模式endpoint signaturefamily:kind按小写 canonical 比较
if isinstance(mapping_api_formats, list):
target = str(api_format).strip().lower()
allowed = {
str(fmt).strip().lower() for fmt in mapping_api_formats if fmt
}
if target not in allowed:
continue
provider_model_names.add(name.strip())
@@ -982,7 +1006,7 @@ class CacheAwareScheduler:
self,
db: Session,
providers: list[Provider],
client_format: APIFormat,
client_format: str,
model_name: str,
affinity_key: str | None,
model_mappings: list[str] | None = None,
@@ -1014,9 +1038,21 @@ class CacheAwareScheduler:
from src.core.api_format.conversion.compatibility import is_format_compatible
candidates: list[ProviderCandidate] = []
client_format_str = client_format.value
client_format_str = normalize_endpoint_signature(client_format)
client_sig = parse_signature_key(client_format_str)
client_family, client_kind = client_sig.api_family, client_sig.endpoint_kind
# chat/cli 互相可回退用于同协议族下的端点变体video/image 等不跨类回退
if client_kind in {EndpointKind.CHAT, EndpointKind.CLI}:
allowed_kinds = {EndpointKind.CHAT, EndpointKind.CLI}
else:
allowed_kinds = {client_kind}
for provider in providers:
logger.debug(
"[Scheduler] Checking provider: %s, endpoints=%d",
provider.name,
len(provider.endpoints) if provider.endpoints else 0,
)
# 按端点格式分别判断兼容性与模型/Key 可用性:
# - 同格式端点优先needs_conversion=False
# - 跨格式端点次之needs_conversion=True
@@ -1026,14 +1062,62 @@ class CacheAwareScheduler:
exact_candidates: list[ProviderCandidate] = []
convertible_candidates: list[ProviderCandidate] = []
for endpoint in provider.endpoints:
if not endpoint.is_active:
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
# - chat/cli 请求允许互相回退(优先同 kind
# - video 等请求只允许同 kind
endpoints = list(provider.endpoints or [])
allowed_kind_values = {k.value for k in allowed_kinds}
preferred: list[ProviderEndpoint] = []
preferred_other_family: list[ProviderEndpoint] = []
fallback: list[ProviderEndpoint] = []
fallback_other_family: list[ProviderEndpoint] = []
for ep in endpoints:
if not getattr(ep, "is_active", False):
continue
endpoint_format_str = (
endpoint.api_format
if isinstance(endpoint.api_format, str)
else endpoint.api_format.value
raw_family = getattr(ep, "api_family", None)
raw_kind = getattr(ep, "endpoint_kind", None)
if not isinstance(raw_family, str) or not raw_family.strip():
continue
if not isinstance(raw_kind, str) or not raw_kind.strip():
continue
ep_family = raw_family.strip().lower()
ep_kind = raw_kind.strip().lower()
if allowed_kind_values and ep_kind not in allowed_kind_values:
continue
same_family = ep_family == client_family.value
same_kind = ep_kind == client_kind.value
if same_kind and same_family:
preferred.append(ep)
elif same_kind:
preferred_other_family.append(ep)
elif same_family:
fallback.append(ep)
else:
fallback_other_family.append(ep)
endpoints = preferred + preferred_other_family + fallback + fallback_other_family
for endpoint in endpoints:
logger.debug(
"[Scheduler] Checking endpoint: family=%s, kind=%s, is_active=%s, base_url=%s",
getattr(endpoint, "api_family", None),
getattr(endpoint, "endpoint_kind", None),
getattr(endpoint, "is_active", None),
(endpoint.base_url[:50] if endpoint.base_url else "N/A"),
)
if not endpoint.is_active:
logger.debug("[Scheduler] Endpoint skipped: not active")
continue
endpoint_format_str = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
@@ -1043,6 +1127,13 @@ class CacheAwareScheduler:
is_stream,
global_conversion_enabled,
)
logger.debug(
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, reason=%s",
client_format_str,
endpoint_format_str,
is_compatible,
_compat_reason,
)
if not is_compatible:
continue
@@ -1059,6 +1150,13 @@ class CacheAwareScheduler:
supports_model, skip_reason, _model_caps, provider_model_names = (
model_support_cache[endpoint_format_str]
)
logger.debug(
"[Scheduler] Model support: provider=%s, model=%s, supports=%s, reason=%s",
provider.name,
model_name,
supports_model,
skip_reason,
)
if not supports_model:
logger.debug(
f"Provider {provider.name} 端点 {endpoint_format_str} 不支持模型 {model_name}: {skip_reason}"
@@ -1111,7 +1209,7 @@ class CacheAwareScheduler:
skip_reason=key_skip_reason,
mapping_matched_model=mapping_matched_model,
needs_conversion=needs_conversion,
provider_api_format=str(endpoint_format_str or "").upper(),
provider_api_format=str(endpoint_format_str or ""),
)
if needs_conversion:
@@ -1132,7 +1230,7 @@ class CacheAwareScheduler:
self,
candidates: list[ProviderCandidate],
affinity_key: str,
api_format: APIFormat,
api_format: str,
global_model_id: str,
) -> list[ProviderCandidate]:
"""
@@ -1151,7 +1249,7 @@ class CacheAwareScheduler:
"""
try:
# 查询该亲和性标识符在当前 API 格式和模型下的缓存亲和性
api_format_str = api_format.value if isinstance(api_format, APIFormat) else api_format
api_format_str = str(api_format)
affinity = await self._affinity_manager.get_affinity(
affinity_key, api_format_str, global_model_id
)
@@ -1164,6 +1262,7 @@ class CacheAwareScheduler:
# 判断候选是否应该被降级(用于分组)
from src.config.settings import config
global_keep_priority = config.keep_priority_on_conversion
def should_demote(c: ProviderCandidate) -> bool:

View File

@@ -20,7 +20,6 @@ from collections import OrderedDict
from typing import Any
import redis.asyncio as aioredis
from src.core.logger import logger
from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger

View File

@@ -4,10 +4,10 @@
统一管理各种缓存的失效逻辑
"""
from __future__ import annotations
from typing import Any
from src.core.logger import logger

View File

@@ -493,9 +493,7 @@ class ModelCacheService:
model_mapping_conflict_total.inc()
# 按名称排序确保确定性
result_global_model = sorted(
mapping_matches, key=lambda gm: gm.name or ""
)[0]
result_global_model = sorted(mapping_matches, key=lambda gm: gm.name or "")[0]
global_model_dict = ModelCacheService._global_model_to_dict(result_global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
@@ -509,9 +507,7 @@ class ModelCacheService:
# 6. 完全未找到
resolution_method = "not_found"
# 未找到匹配,缓存负结果
await CacheService.set(
cache_key, "NOT_FOUND", ttl_seconds=ModelCacheService.CACHE_TTL
)
await CacheService.set(cache_key, "NOT_FOUND", ttl_seconds=ModelCacheService.CACHE_TTL)
logger.debug(f"GlobalModel 未找到(映射解析): {normalized_name}")
return None

View File

@@ -5,7 +5,6 @@ Provider 缓存服务 - 减少 Provider 和 ProviderAPIKey 查询
这些数据在 UsageService.record_usage() 中被频繁查询但变化不频繁。
"""
from __future__ import annotations
from sqlalchemy.orm import Session
@@ -44,9 +43,9 @@ class ProviderCacheService:
计算后的 rate_multiplier
"""
if api_format and rate_multipliers:
format_upper = api_format.upper()
if format_upper in rate_multipliers:
return float(rate_multipliers[format_upper])
format_key = str(api_format).strip().lower()
if format_key in rate_multipliers:
return float(rate_multipliers[format_key])
return 1.0
@staticmethod
@@ -67,13 +66,15 @@ class ProviderCacheService:
rate_multiplier 或 None如果找不到
"""
# 缓存键包含 api_format
format_suffix = api_format.upper() if api_format else "default"
format_suffix = str(api_format).strip().lower() if api_format else "default"
cache_key = f"provider_api_key:rate_multiplier:{provider_api_key_id}:{format_suffix}"
# 1. 尝试从缓存获取
cached_data = await CacheService.get(cache_key)
if cached_data is not None:
logger.debug(f"ProviderAPIKey rate_multiplier 缓存命中: {provider_api_key_id[:8]}... format={format_suffix}")
logger.debug(
f"ProviderAPIKey rate_multiplier 缓存命中: {provider_api_key_id[:8]}... format={format_suffix}"
)
# 缓存的 "NOT_FOUND" 表示数据库中不存在
if cached_data == "NOT_FOUND":
return None
@@ -95,7 +96,9 @@ class ProviderCacheService:
await CacheService.set(
cache_key, rate_multiplier, ttl_seconds=ProviderCacheService.CACHE_TTL
)
logger.debug(f"ProviderAPIKey rate_multiplier 已缓存: {provider_api_key_id[:8]}... format={format_suffix} value={rate_multiplier}")
logger.debug(
f"ProviderAPIKey rate_multiplier 已缓存: {provider_api_key_id[:8]}... format={format_suffix} value={rate_multiplier}"
)
return rate_multiplier
else:
# 缓存负结果
@@ -133,9 +136,7 @@ class ProviderCacheService:
await CacheService.delete(cache_key)
# 2. 缓存未命中,查询数据库
provider = (
db.query(Provider.billing_type).filter(Provider.id == provider_id).first()
)
provider = db.query(Provider.billing_type).filter(Provider.id == provider_id).first()
# 3. 写入缓存
if provider:
@@ -196,7 +197,9 @@ class ProviderCacheService:
async def invalidate_provider_api_key_cache(provider_api_key_id: str) -> None:
"""清除 ProviderAPIKey 缓存(包括所有 API 格式的缓存)"""
# 使用模式匹配删除所有格式的缓存
await CacheService.delete_pattern(f"provider_api_key:rate_multiplier:{provider_api_key_id}:*")
await CacheService.delete_pattern(
f"provider_api_key:rate_multiplier:{provider_api_key_id}:*"
)
logger.debug(f"ProviderAPIKey 缓存已清除: {provider_api_key_id[:8]}...")
@staticmethod

View File

@@ -11,14 +11,12 @@
from __future__ import annotations
from typing import Any
import asyncio
import json
from collections.abc import Callable
from typing import Any
import redis.asyncio as aioredis
from src.core.logger import logger
from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger
@@ -69,9 +67,11 @@ class CacheSyncService:
self._listener_task = asyncio.create_task(self._listen())
self._running = True
logger.info("[CacheSync] 缓存同步服务已启动,订阅频道: "
logger.info(
"[CacheSync] 缓存同步服务已启动,订阅频道: "
f"{self.CHANNEL_GLOBAL_MODEL}, "
f"{self.CHANNEL_MODEL}, {self.CHANNEL_CLEAR_ALL}")
f"{self.CHANNEL_MODEL}, {self.CHANNEL_CLEAR_ALL}"
)
except Exception as e:
logger.error(f"[CacheSync] 启动失败: {e}")
raise
@@ -167,7 +167,9 @@ class CacheSyncService:
_cache_sync_service: CacheSyncService | None = None
async def get_cache_sync_service(redis_client: aioredis.Redis | None = None) -> CacheSyncService | None:
async def get_cache_sync_service(
redis_client: aioredis.Redis | None = None,
) -> CacheSyncService | None:
"""
获取缓存同步服务实例

View File

@@ -19,10 +19,10 @@
await UserCacheService.invalidate_user_cache(user_id, email)
"""
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL

View File

@@ -9,16 +9,15 @@
5. 显式传入 (用于重试升级)
"""
from collections.abc import Callable
from typing import Any
from collections.abc import Callable
from src.core.api_format import get_header_value
from src.core.key_capabilities import (
CAPABILITY_DEFINITIONS,
CapabilityConfigMode,
get_user_configurable_capabilities,
)
from src.core.api_format import get_header_value
from src.core.logger import logger
# Adapter 检测器类型:接受 headers 和可选的 request_body返回能力需求字典
@@ -113,9 +112,7 @@ class CapabilityResolver:
# 只有尚未设置的能力才从 Adapter 检测
if cap_name not in requirements:
requirements[cap_name] = cap_value
logger.debug(
f"[CapabilityResolver] 从 Adapter 检测到 {cap_name}={cap_value}"
)
logger.debug(f"[CapabilityResolver] 从 Adapter 检测到 {cap_name}={cap_value}")
# 5. 显式覆盖(重试时使用)
if explicit_requirements:

View File

@@ -64,7 +64,9 @@ class EmailSenderService:
"smtp_use_tls": SystemConfigService.get_config(db, "smtp_use_tls", default=True),
"smtp_use_ssl": SystemConfigService.get_config(db, "smtp_use_ssl", default=False),
"smtp_from_email": SystemConfigService.get_config(db, "smtp_from_email"),
"smtp_from_name": SystemConfigService.get_config(db, "smtp_from_name", default="Aether"),
"smtp_from_name": SystemConfigService.get_config(
db, "smtp_from_name", default="Aether"
),
}
return config
@@ -143,7 +145,11 @@ class EmailSenderService:
# 发送邮件
return await EmailSenderService._send_email(
config=config, to_email=to_email, subject=subject, html_body=html_body, text_body=text_body
config=config,
to_email=to_email,
subject=subject,
html_body=html_body,
text_body=text_body,
)
@staticmethod

View File

@@ -419,9 +419,7 @@ class EmailTemplate:
return EmailTemplate.html_to_text(html)
@staticmethod
def get_subject(
template_type: str = "verification", db: Session | None = None
) -> str:
def get_subject(template_type: str = "verification", db: Session | None = None) -> str:
"""
获取邮件主题

View File

@@ -63,7 +63,7 @@ def _extract_file_name_from_uri(file_uri: str) -> str | None:
# 完整 URL 格式
if "/files/" in file_uri:
idx = file_uri.rfind("/files/")
return file_uri[idx + 1:] # 提取 files/xxx 部分
return file_uri[idx + 1 :] # 提取 files/xxx 部分
# 短格式
if file_uri.startswith("files/"):
return file_uri

View File

@@ -21,7 +21,6 @@ from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, RequestCandidate
# 缓存配置
CACHE_TTL_SECONDS = 30 # 缓存 30 秒
CACHE_KEY_PREFIX = "health:endpoint:"
@@ -31,6 +30,7 @@ def _get_redis_client() -> Any:
"""获取 Redis 客户端,失败返回 None"""
try:
from src.clients.redis_client import redis_client
return redis_client
except Exception:
return None
@@ -77,15 +77,19 @@ class EndpointHealthService:
# 批量查询所有密钥(通过 provider_id 关联)
all_keys = (
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id.in_(all_provider_ids))
.all()
) if all_provider_ids else []
(
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id.in_(all_provider_ids))
.all()
)
if all_provider_ids
else []
)
# 按 api_format 分组密钥(通过 api_formats 字段)
keys_by_format: dict[str, list[ProviderAPIKey]] = defaultdict(list)
for key in all_keys:
for fmt in (key.api_formats or []):
for fmt in key.api_formats or []:
keys_by_format[fmt].append(key)
# 按 API 格式聚合
@@ -157,11 +161,14 @@ class EndpointHealthService:
result = []
for api_format, stats in format_stats.items():
timeline_data = timeline_data_map.get(api_format, {
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
})
timeline_data = timeline_data_map.get(
api_format,
{
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
},
)
timeline = timeline_data["timeline"]
time_range_start = timeline_data.get("time_range_start")
time_range_end = timeline_data.get("time_range_end")
@@ -271,28 +278,22 @@ class EndpointHealthService:
final_statuses = ["success", "failed", "skipped"]
segment_expr = func.floor(
func.extract('epoch', RequestCandidate.created_at - start_time) / segment_seconds
).label('segment_idx')
func.extract("epoch", RequestCandidate.created_at - start_time) / segment_seconds
).label("segment_idx")
candidate_stats = (
db.query(
RequestCandidate.key_id,
segment_expr,
func.count(RequestCandidate.id).label('total_count'),
func.sum(
case(
(RequestCandidate.status == "success", 1),
else_=0
)
).label('success_count'),
func.sum(
case(
(RequestCandidate.status == "failed", 1),
else_=0
)
).label('failed_count'),
func.min(RequestCandidate.created_at).label('min_time'),
func.max(RequestCandidate.created_at).label('max_time'),
func.count(RequestCandidate.id).label("total_count"),
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
"success_count"
),
func.sum(case((RequestCandidate.status == "failed", 1), else_=0)).label(
"failed_count"
),
func.min(RequestCandidate.created_at).label("min_time"),
func.max(RequestCandidate.created_at).label("max_time"),
)
.filter(
RequestCandidate.key_id.in_(all_key_ids),
@@ -311,13 +312,17 @@ class EndpointHealthService:
key_to_format[key_id] = api_format
# 按 api_format 和 segment 聚合数据
format_segment_data: dict[str, dict[int, dict]] = defaultdict(lambda: defaultdict(lambda: {
"total": 0,
"success": 0,
"failed": 0,
"min_time": None,
"max_time": None,
}))
format_segment_data: dict[str, dict[int, dict]] = defaultdict(
lambda: defaultdict(
lambda: {
"total": 0,
"success": 0,
"failed": 0,
"min_time": None,
"max_time": None,
}
)
)
for row in candidate_stats:
key_id = row.key_id
@@ -454,24 +459,34 @@ class EndpointHealthService:
db, format_key_mapping, now, lookback_hours, segments
)
return result.get("_single", {
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
})
return result.get(
"_single",
{
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
},
)
@staticmethod
def _format_display_name(api_format: str) -> str:
"""格式化 API 格式的显示名称"""
format_names = {
"CLAUDE": "Claude API",
"CLAUDE_CLI": "Claude CLI",
"CLAUDE_COMPATIBLE": "Claude 兼容",
"OPENAI": "OpenAI API",
"OPENAI_CLI": "OpenAI CLI",
"OPENAI_COMPATIBLE": "OpenAI 兼容",
}
return format_names.get(api_format, api_format)
raw = str(api_format or "").strip()
normalized = raw.lower()
if ":" not in normalized:
return raw or api_format
fam, kind = normalized.split(":", 1)
fam_label = {"claude": "Claude", "openai": "OpenAI", "gemini": "Gemini"}.get(fam, fam)
kind_label = {
"chat": "",
"cli": "CLI",
"video": "Video",
"image": "Image",
}.get(kind, kind)
if not kind_label:
return fam_label
return f"{fam_label} {kind_label}"
@staticmethod
def _get_from_cache(key: str) -> list[dict[str, Any]] | None:

View File

@@ -230,9 +230,9 @@ class HealthMonitor:
if state == CircuitState.HALF_OPEN:
# 半开状态:记录成功
circuit_data["half_open_successes"] = int(
circuit_data.get("half_open_successes") or 0
) + 1
circuit_data["half_open_successes"] = (
int(circuit_data.get("half_open_successes") or 0) + 1
)
if circuit_data["half_open_successes"] >= cls.HALF_OPEN_SUCCESS_THRESHOLD:
# 达到成功阈值,关闭熔断器
@@ -357,9 +357,9 @@ class HealthMonitor:
if state == CircuitState.HALF_OPEN:
# 半开状态:记录失败
circuit_data["half_open_failures"] = int(
circuit_data.get("half_open_failures") or 0
) + 1
circuit_data["half_open_failures"] = (
int(circuit_data.get("half_open_failures") or 0) + 1
)
if circuit_data["half_open_failures"] >= cls.HALF_OPEN_FAILURE_THRESHOLD:
# 达到失败阈值,重新打开熔断器
@@ -581,9 +581,7 @@ class HealthMonitor:
return cls._get_status_from_circuit_data(circuit_data)
@classmethod
def _get_status_from_circuit_data(
cls, circuit_data: dict[str, Any]
) -> tuple[bool, str | None]:
def _get_status_from_circuit_data(cls, circuit_data: dict[str, Any]) -> tuple[bool, str | None]:
"""从熔断器数据获取状态描述"""
if not circuit_data.get("open"):
return True, None
@@ -662,9 +660,7 @@ class HealthMonitor:
result["health_score"] = float(health_data.get("health_score") or 1.0)
result["error_rate"] = cls._calculate_error_rate_from_window(window, now_ts)
result["window_size"] = len(valid_window)
result["consecutive_failures"] = int(
health_data.get("consecutive_failures") or 0
)
result["consecutive_failures"] = int(health_data.get("consecutive_failures") or 0)
result["last_failure_at"] = health_data.get("last_failure_at")
result["circuit_breaker"] = {
"state": cls._get_circuit_state_from_data(circuit_data, now),
@@ -678,7 +674,7 @@ class HealthMonitor:
else:
# 返回所有格式的健康度数据
formats_health = {}
for fmt in (key.api_formats or []):
for fmt in key.api_formats or []:
health_data = health_by_format.get(fmt, _default_health_data())
circuit_data = circuit_by_format.get(fmt, _default_circuit_data())
window = health_data.get("request_results_window") or []
@@ -688,9 +684,7 @@ class HealthMonitor:
"health_score": float(health_data.get("health_score") or 1.0),
"error_rate": cls._calculate_error_rate_from_window(window, now_ts),
"window_size": len(valid_window),
"consecutive_failures": int(
health_data.get("consecutive_failures") or 0
),
"consecutive_failures": int(health_data.get("consecutive_failures") or 0),
"last_failure_at": health_data.get("last_failure_at"),
"circuit_breaker": {
"state": cls._get_circuit_state_from_data(circuit_data, now),
@@ -701,9 +695,7 @@ class HealthMonitor:
"half_open_successes": int(
circuit_data.get("half_open_successes") or 0
),
"half_open_failures": int(
circuit_data.get("half_open_failures") or 0
),
"half_open_failures": int(circuit_data.get("half_open_failures") or 0),
},
}
@@ -711,9 +703,7 @@ class HealthMonitor:
# 计算整体健康度(取最低值)
if formats_health:
result["health_score"] = min(
h["health_score"] for h in formats_health.values()
)
result["health_score"] = min(h["health_score"] for h in formats_health.values())
result["any_circuit_open"] = any(
h["circuit_breaker"]["open"] for h in formats_health.values()
)
@@ -731,9 +721,7 @@ class HealthMonitor:
def get_endpoint_health(cls, db: Session, endpoint_id: str) -> dict[str, Any] | None:
"""获取 Endpoint 健康状态"""
try:
endpoint = (
db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
)
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
if not endpoint:
return None
@@ -782,6 +770,12 @@ class HealthMonitor:
db.rollback()
return False
@classmethod
def reset_open_circuit_count(cls) -> None:
"""重置进程级别的熔断计数器(批量恢复后调用)。"""
cls._open_circuit_keys = 0
health_open_circuits.set(0)
@classmethod
def manually_enable(cls, db: Session, key_id: str | None = None) -> bool:
"""手动启用 Key"""
@@ -915,18 +909,14 @@ class HealthMonitor:
# ==================== 便捷方法 ====================
@classmethod
def get_health_score(
cls, key: ProviderAPIKey, api_format: str | None = None
) -> float:
def get_health_score(cls, key: ProviderAPIKey, api_format: str | None = None) -> float:
"""获取指定格式的健康度分数"""
if not api_format:
# 返回所有格式中的最低健康度
health_by_format = key.health_by_format or {}
if not health_by_format:
return 1.0
return min(
float(h.get("health_score") or 1.0) for h in health_by_format.values()
)
return min(float(h.get("health_score") or 1.0) for h in health_by_format.values())
health_data = cls._get_health_data(key, api_format)
return float(health_data.get("health_score") or 1.0)

View File

@@ -2,9 +2,9 @@
from __future__ import annotations
from typing import Any
import ipaddress
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
@@ -159,9 +159,7 @@ class ManagementTokenService:
ValueError: 如果名称已存在或超过数量限制
"""
# 检查用户 Token 数量限制
token_count = (
db.query(ManagementToken).filter(ManagementToken.user_id == user_id).count()
)
token_count = db.query(ManagementToken).filter(ManagementToken.user_id == user_id).count()
max_tokens = config.management_token_max_per_user
if token_count >= max_tokens:
raise ValueError(f"已达到 Token 数量上限({max_tokens}")
@@ -335,9 +333,7 @@ class ManagementTokenService:
return token
@staticmethod
def delete_token(
db: Session, token_id: str, user_id: str | None = None
) -> bool:
def delete_token(db: Session, token_id: str, user_id: str | None = None) -> bool:
"""删除 Token
Args:

View File

@@ -4,10 +4,10 @@
包含模型管理、成本计算等功能。
"""
from src.services.model.availability import ModelAvailabilityQuery
from src.services.model.cost import ModelCostService
from src.services.model.fetch_scheduler import ModelFetchScheduler, get_model_fetch_scheduler
from src.services.model.global_model import GlobalModelService
from src.services.model.availability import ModelAvailabilityQuery
from src.services.model.service import ModelService
__all__ = [

View File

@@ -8,9 +8,7 @@
- API Key/User 的请求级访问限制由 models_service.AccessRestrictions 处理
"""
from sqlalchemy import or_
from sqlalchemy import or_, tuple_
from sqlalchemy.orm import Query, Session, contains_eager
from src.core.logger import logger
@@ -21,6 +19,7 @@ from src.models.database import (
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.provider.format import normalize_endpoint_signature
class ModelAvailabilityQuery:
@@ -86,23 +85,45 @@ class ModelAvailabilityQuery:
Returns:
{provider_id: {format1, format2, ...}}
"""
target_formats = {f.upper() for f in api_formats}
target_pairs: list[tuple[str, str]] = []
for fmt in api_formats:
if not fmt:
continue
try:
norm = normalize_endpoint_signature(fmt)
fam, kind = norm.split(":", 1)
if fam and kind:
target_pairs.append((fam, kind))
except Exception:
continue
if not target_pairs:
return {}
endpoint_rows = (
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
db.query(
ProviderEndpoint.provider_id,
ProviderEndpoint.api_family,
ProviderEndpoint.endpoint_kind,
)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
Provider.is_active.is_(True),
ProviderEndpoint.api_format.in_(list(target_formats)),
ProviderEndpoint.is_active.is_(True),
ProviderEndpoint.api_family.isnot(None),
ProviderEndpoint.endpoint_kind.isnot(None),
tuple_(ProviderEndpoint.api_family, ProviderEndpoint.endpoint_kind).in_(
target_pairs
),
)
.all()
)
provider_to_formats: dict[str, set[str]] = {}
for provider_id, fmt in endpoint_rows:
if provider_id and fmt:
provider_to_formats.setdefault(provider_id, set()).add(str(fmt).upper())
for provider_id, fam, kind in endpoint_rows:
if provider_id and fam and kind:
provider_to_formats.setdefault(provider_id, set()).add(
normalize_endpoint_signature(f"{fam}:{kind}")
)
return provider_to_formats
@@ -123,7 +144,7 @@ class ModelAvailabilityQuery:
if not provider_ids:
return set()
target_formats = {f.upper() for f in api_formats}
target_formats = {normalize_endpoint_signature(f) for f in api_formats if f}
key_rows = (
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
@@ -144,17 +165,24 @@ class ModelAvailabilityQuery:
continue
# 类型兜底key_formats 是 JSON 字段
if not isinstance(key_formats, list):
if key_formats is not None:
logger.warning(
f"[ModelAvailability] Key api_formats 类型异常, "
f"provider_id={provider_id}, type={type(key_formats).__name__}"
)
if key_formats is None:
# None = 全支持(兼容历史数据)
key_formats_norm = set(endpoint_formats)
elif not isinstance(key_formats, list):
logger.warning(
"[ModelAvailability] Key api_formats 类型异常, provider_id=%s, type=%s",
provider_id,
type(key_formats).__name__,
)
continue
else:
key_formats_norm = {
normalize_endpoint_signature(str(f))
for f in key_formats
if isinstance(f, str) and f
}
key_formats_upper: set[str] = {str(f).upper() for f in key_formats if isinstance(f, str)}
if key_formats_upper & endpoint_formats & target_formats:
if key_formats_norm & endpoint_formats & target_formats:
available_provider_ids.add(provider_id)
return available_provider_ids
@@ -179,7 +207,7 @@ class ModelAvailabilityQuery:
if not provider_ids:
return {}
target_formats = {f.upper() for f in api_formats}
target_formats = {normalize_endpoint_signature(f) for f in api_formats if f}
key_rows = (
db.query(
@@ -205,16 +233,23 @@ class ModelAvailabilityQuery:
continue
# 类型兜底key_formats
if not isinstance(key_formats, list):
if key_formats is not None:
logger.warning(
f"[ModelAvailability] Key api_formats 类型异常, "
f"key_id={key_id}, type={type(key_formats).__name__}"
)
if key_formats is None:
key_formats_norm = set(endpoint_formats)
elif not isinstance(key_formats, list):
logger.warning(
"[ModelAvailability] Key api_formats 类型异常, key_id=%s, type=%s",
key_id,
type(key_formats).__name__,
)
continue
else:
key_formats_norm = {
normalize_endpoint_signature(str(f))
for f in key_formats
if isinstance(f, str) and f
}
key_formats_upper: set[str] = {str(f).upper() for f in key_formats if isinstance(f, str)}
usable_formats = key_formats_upper & endpoint_formats & target_formats
usable_formats = key_formats_norm & endpoint_formats & target_formats
if not usable_formats:
continue
@@ -259,4 +294,3 @@ class ModelAvailabilityQuery:
query = query.filter(Model.provider_id.in_(provider_ids))
return query

View File

@@ -17,13 +17,13 @@ from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import GlobalModel, Model, Provider
ProviderRef = str | Provider | None
@dataclass
class TieredPriceResult:
"""阶梯计费价格查询结果"""
input_price_per_1m: float
output_price_per_1m: float
cache_creation_price_per_1m: float | None = None
@@ -34,6 +34,7 @@ class TieredPriceResult:
@dataclass
class CostBreakdown:
"""成本明细"""
input_cost: float
output_cost: float
cache_creation_cost: float
@@ -58,10 +59,7 @@ class ModelCostService:
# ------------------------------------------------------------------
@staticmethod
def get_tier_for_tokens(
tiered_pricing: dict,
total_input_tokens: int
) -> dict | None:
def get_tier_for_tokens(tiered_pricing: dict, total_input_tokens: int) -> dict | None:
"""
根据总输入 token 数确定价格阶梯。
@@ -89,8 +87,7 @@ class ModelCostService:
@staticmethod
def get_cache_read_price_for_ttl(
tier: dict,
cache_ttl_minutes: int | None = None
tier: dict, cache_ttl_minutes: int | None = None
) -> float | None:
"""
根据缓存 TTL 获取缓存读取价格。
@@ -121,9 +118,7 @@ class ModelCostService:
# 使用默认的缓存读取价格
return tier.get("cache_read_price_per_1m")
async def get_tiered_pricing_async(
self, provider: ProviderRef, model: str
) -> dict | None:
async def get_tiered_pricing_async(self, provider: ProviderRef, model: str) -> dict | None:
"""
异步获取模型的阶梯计费配置。
@@ -186,21 +181,18 @@ class ModelCostService:
if model_obj:
# 判断定价来源
if model_obj.tiered_pricing is not None:
result = {
"pricing": model_obj.tiered_pricing,
"source": "provider"
}
result = {"pricing": model_obj.tiered_pricing, "source": "provider"}
elif global_model.default_tiered_pricing is not None:
result = {
"pricing": global_model.default_tiered_pricing,
"source": "global"
"source": "global",
}
else:
# Provider 没有实现该模型,直接使用 GlobalModel 的默认阶梯配置
if global_model.default_tiered_pricing is not None:
result = {
"pricing": global_model.default_tiered_pricing,
"source": "global"
"source": "global",
}
self._tiered_pricing_cache[cache_key] = result
@@ -250,10 +242,16 @@ class ModelCostService:
if model_obj.tiered_pricing is not None:
result = {"pricing": model_obj.tiered_pricing, "source": "provider"}
elif global_model.default_tiered_pricing is not None:
result = {"pricing": global_model.default_tiered_pricing, "source": "global"}
result = {
"pricing": global_model.default_tiered_pricing,
"source": "global",
}
else:
if global_model.default_tiered_pricing is not None:
result = {"pricing": global_model.default_tiered_pricing, "source": "global"}
result = {
"pricing": global_model.default_tiered_pricing,
"source": "global",
}
self._tiered_pricing_cache[cache_key] = result
return result.get("pricing") if result else None
@@ -322,8 +320,10 @@ class ModelCostService:
else:
input_price = model_obj.get_effective_input_price()
output_price = model_obj.get_effective_output_price()
logger.debug(f"找到模型价格配置: {provider_name}/{model} "
f"(输入: ${input_price}/M, 输出: ${output_price}/M)")
logger.debug(
f"找到模型价格配置: {provider_name}/{model} "
f"(输入: ${input_price}/M, 输出: ${output_price}/M)"
)
else:
# Provider 没有实现该模型,直接使用 GlobalModel 的默认价格
tiered = global_model.default_tiered_pricing
@@ -334,8 +334,10 @@ class ModelCostService:
else:
input_price = 0.0
output_price = 0.0
logger.debug(f"使用 GlobalModel 默认价格: {provider_name}/{model} "
f"(输入: ${input_price}/M, 输出: ${output_price}/M)")
logger.debug(
f"使用 GlobalModel 默认价格: {provider_name}/{model} "
f"(输入: ${input_price}/M, 输出: ${output_price}/M)"
)
# 如果没有找到价格配置,使用 0.0 并记录警告
if input_price is None:
@@ -348,7 +350,9 @@ class ModelCostService:
# 异步检查按次计费价格
price_per_request = await self.get_request_price_async(provider, model)
if price_per_request is None or price_per_request == 0.0:
logger.warning(f"未找到模型价格配置: {provider_name}/{model},请在 GlobalModel 中配置价格")
logger.warning(
f"未找到模型价格配置: {provider_name}/{model},请在 GlobalModel 中配置价格"
)
self._price_cache[cache_key] = {"input": input_price, "output": output_price}
return input_price, output_price

View File

@@ -110,9 +110,7 @@ def _get_upstream_models_cache_key(provider_id: str, api_key_id: str) -> str:
return f"upstream_models:{provider_id}:{api_key_id}"
async def get_upstream_models_from_cache(
provider_id: str, api_key_id: str
) -> list[dict] | None:
async def get_upstream_models_from_cache(provider_id: str, api_key_id: str) -> list[dict] | None:
"""从缓存获取上游模型列表"""
cache_key = _get_upstream_models_cache_key(provider_id, api_key_id)
cached = await CacheService.get(cache_key)

View File

@@ -597,9 +597,7 @@ class GlobalModelService:
db.commit()
except Exception as e:
db.rollback()
logger.error(
f"Failed to auto-disassociate Provider {provider.name}: {e}"
)
logger.error(f"Failed to auto-disassociate Provider {provider.name}: {e}")
# 清空 success记录整体错误
results["success"] = []
results["errors"].append(

View File

@@ -68,18 +68,20 @@ class ModelMapperMiddleware:
original_model = request.model
request.model = mapping.model.select_provider_model_name()
logger.debug(f"Applied model mapping for provider {provider.name}: "
f"{original_model} -> {request.model}")
logger.debug(
f"Applied model mapping for provider {provider.name}: "
f"{original_model} -> {request.model}"
)
else:
# 没有找到映射,使用原始模型名
logger.debug(f"No model mapping found for {source_model} with provider {provider.name}, "
f"forwarding with original model name")
logger.debug(
f"No model mapping found for {source_model} with provider {provider.name}, "
f"forwarding with original model name"
)
return request
async def get_mapping(
self, source_model: str, provider_id: str
) -> object | None:
async def get_mapping(self, source_model: str, provider_id: str) -> object | None:
"""
获取模型映射
@@ -129,8 +131,10 @@ class ModelMapperMiddleware:
},
)()
logger.debug(f"Found model mapping: {source_model} -> {model.provider_model_name} "
f"(provider={provider_id[:8]}...)")
logger.debug(
f"Found model mapping: {source_model} -> {model.provider_model_name} "
f"(provider={provider_id[:8]}...)"
)
# 缓存结果
self._cache[cache_key] = mapping
@@ -275,6 +279,15 @@ class ModelRoutingMiddleware:
选中的提供商如果没有找到返回None
"""
request_prefix = f"ID:{request_id} | " if request_id else ""
allowed_norm: set[str] | None = None
if allowed_api_formats:
from src.services.provider.format import normalize_endpoint_signature
allowed_norm = {
normalize_endpoint_signature(str(fmt))
for fmt in allowed_api_formats
if isinstance(fmt, str) and fmt
}
# 1. 如果指定了提供商,直接使用
if preferred_provider:
@@ -286,18 +299,26 @@ class ModelRoutingMiddleware:
if provider:
# 检查API格式 - 从 endpoints 中检查
if allowed_api_formats:
if allowed_norm:
has_matching_endpoint = any(
ep.is_active and ep.api_format and ep.api_format in allowed_api_formats
ep.is_active
and ep.api_format
and str(ep.api_format).strip().lower() in allowed_norm
for ep in provider.endpoints
)
if not has_matching_endpoint:
logger.warning(f"Specified provider {provider.name} has no active endpoints with allowed API formats ({allowed_api_formats})")
logger.warning(
f"Specified provider {provider.name} has no active endpoints with allowed API formats ({allowed_api_formats})"
)
else:
logger.debug(f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}")
logger.debug(
f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}"
)
return provider
else:
logger.debug(f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}")
logger.debug(
f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}"
)
return provider
else:
logger.warning(f"Specified provider {preferred_provider} not found or inactive")
@@ -305,12 +326,12 @@ class ModelRoutingMiddleware:
# 2. 查找优先级最高的活动提供商
query = self.db.query(Provider).filter(Provider.is_active == True)
if allowed_api_formats:
if allowed_norm:
query = (
query.join(ProviderEndpoint)
.filter(
ProviderEndpoint.is_active == True,
ProviderEndpoint.api_format.in_(allowed_api_formats),
ProviderEndpoint.api_format.in_(sorted(allowed_norm)),
)
.distinct()
)
@@ -318,11 +339,15 @@ class ModelRoutingMiddleware:
best_provider = query.order_by(Provider.provider_priority.asc(), Provider.id.asc()).first()
if best_provider:
logger.debug(f" └─ {request_prefix}使用优先级最高提供商: {best_provider.name} (priority:{best_provider.provider_priority}) | 模型:{model_name}")
logger.debug(
f" └─ {request_prefix}使用优先级最高提供商: {best_provider.name} (priority:{best_provider.provider_priority}) | 模型:{model_name}"
)
return best_provider
if allowed_api_formats:
logger.error(f"No active providers found with allowed API formats {allowed_api_formats}.")
logger.error(
f"No active providers found with allowed API formats {allowed_api_formats}."
)
else:
logger.error("No active providers found.")
return None
@@ -392,13 +417,15 @@ class ModelRoutingMiddleware:
# 按总价格排序
cheapest = min(
models_with_providers,
key=lambda x: x[1].get_effective_input_price() + x[1].get_effective_output_price()
key=lambda x: x[1].get_effective_input_price() + x[1].get_effective_output_price(),
)
provider = cheapest[0]
model = cheapest[1]
logger.debug(f"Selected cheapest provider {provider.name} for model {model_name} "
f"(input: ${model.get_effective_input_price()}/M, output: ${model.get_effective_output_price()}/M)")
logger.debug(
f"Selected cheapest provider {provider.name} for model {model_name} "
f"(input: ${model.get_effective_input_price()}/M, output: ${model.get_effective_output_price()}/M)"
)
return provider

View File

@@ -16,6 +16,7 @@ from dataclasses import dataclass
@dataclass
class UsageTokens:
"""请求的 token 使用量"""
input_tokens: int = 0
output_tokens: int = 0
cache_creation_input_tokens: int = 0
@@ -25,6 +26,7 @@ class UsageTokens:
@dataclass
class PricingConfig:
"""价格配置"""
input_price_per_1m: float = 0.0
output_price_per_1m: float = 0.0
cache_creation_price_per_1m: float | None = None
@@ -37,6 +39,7 @@ class PricingConfig:
@dataclass
class CostResult:
"""计费结果"""
input_cost: float = 0.0
output_cost: float = 0.0
cache_creation_cost: float = 0.0

View File

@@ -10,16 +10,15 @@ 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.api.base.models_service import invalidate_models_list_cache
from src.services.cache.invalidation import get_cache_invalidation_service
from src.services.cache.model_cache import ModelCacheService
class ModelService:
"""模型管理服务"""
@@ -76,7 +75,9 @@ class ModelService:
.first()
)
logger.info(f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}")
logger.info(
f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}"
)
# 清除 Redis 缓存(异步执行,不阻塞返回)
# 重要:新增模型可能需要清除 resolver 的 NOT_FOUND 负缓存global_model:resolve:*
@@ -182,12 +183,16 @@ class ModelService:
# 添加调试日志
logger.debug(f"更新模型 {model_id} 收到的数据: {update_data}")
logger.debug(f"更新前的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}")
logger.debug(
f"更新前的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
)
for field, value in update_data.items():
setattr(model, field, value)
logger.debug(f"更新后的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}")
logger.debug(
f"更新后的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
)
try:
db.commit()
@@ -228,7 +233,9 @@ class ModelService:
# 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache())
logger.info(f"更新模型成功: id={model_id}, 最终 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}")
logger.info(
f"更新模型成功: id={model_id}, 最终 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
)
return model
except IntegrityError as e:
db.rollback()
@@ -260,8 +267,10 @@ class ModelService:
)
if other_implementations == 0:
logger.warning(f"警告:删除模型 {model_id}Provider: {model.provider_id[:8]}...)后,"
f"GlobalModel '{model.global_model_id}' 将没有任何活跃的关联提供商")
logger.warning(
f"警告:删除模型 {model_id}Provider: {model.provider_id[:8]}...)后,"
f"GlobalModel '{model.global_model_id}' 将没有任何活跃的关联提供商"
)
# 保存缓存清除所需的信息(删除后无法访问)
cache_info = {
@@ -290,13 +299,17 @@ class ModelService:
# 清除内存缓存
if cache_info["provider_id"] and cache_info["global_model_id"]:
cache_service = get_cache_invalidation_service()
cache_service.on_model_changed(cache_info["provider_id"], cache_info["global_model_id"])
cache_service.on_model_changed(
cache_info["provider_id"], cache_info["global_model_id"]
)
# 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache())
logger.info(f"删除模型成功: id={model_id}, provider_model_name={cache_info['provider_model_name']}, "
f"global_model_id={cache_info['global_model_id'][:8] if cache_info['global_model_id'] else 'None'}...")
logger.info(
f"删除模型成功: id={model_id}, provider_model_name={cache_info['provider_model_name']}, "
f"global_model_id={cache_info['global_model_id'][:8] if cache_info['global_model_id'] else 'None'}..."
)
except Exception as e:
db.rollback()
logger.error(f"删除模型失败: {str(e)}")

View File

@@ -10,7 +10,7 @@ import asyncio
import httpx
from src.core.api_format import APIFormat, get_extra_headers_from_endpoint
from src.core.api_format import get_extra_headers_from_endpoint
from src.core.logger import logger
from src.models.database import ProviderEndpoint
from src.utils.ssl_utils import get_ssl_context
@@ -18,8 +18,8 @@ from src.utils.ssl_utils import get_ssl_context
# 并发请求限制
MAX_CONCURRENT_REQUESTS = 5
# 只对这些基础格式获取模型列表CLI 格式使用相同的上游 API
MODEL_FETCH_FORMATS = [APIFormat.OPENAI, APIFormat.CLAUDE, APIFormat.GEMINI]
# 只对这些基础 endpoint signature 获取模型列表CLI 使用相同的上游 API
MODEL_FETCH_FORMATS = ["openai:chat", "claude:chat", "gemini:chat"]
def _get_adapter_for_format(api_format: str) -> type | None:
@@ -43,7 +43,7 @@ def build_all_format_configs(
"""
构建所有 API 格式的端点配置
所有 APIFormat 枚举值构建配置,如果该格式有专门的端点配置则使用,
基础 endpoint signature 列表构建配置,如果该格式有专门的端点配置则使用,
否则使用基础端点的 base_url 尝试。
Args:
@@ -59,9 +59,9 @@ def build_all_format_configs(
# 获取任意一个端点的 base_url 作为基础(用于尝试所有格式)
# 优先使用 OPENAI 格式的端点,因为它最通用
base_endpoint = (
format_to_endpoint.get("OPENAI")
or format_to_endpoint.get("CLAUDE")
or format_to_endpoint.get("GEMINI")
format_to_endpoint.get("openai:chat")
or format_to_endpoint.get("claude:chat")
or format_to_endpoint.get("gemini:chat")
or next(iter(format_to_endpoint.values()))
)
base_url = base_endpoint.base_url
@@ -70,7 +70,7 @@ def build_all_format_configs(
# 只对基础 API 格式获取模型CLI 格式使用相同的上游 API
endpoint_configs: list[dict] = []
for fmt in MODEL_FETCH_FORMATS:
fmt_value = fmt.value
fmt_value = fmt
# 如果该格式有专门的端点配置,使用其 base_url 和 headers
if fmt_value in format_to_endpoint:
ep = format_to_endpoint[fmt_value]
@@ -115,9 +115,7 @@ async def fetch_models_from_endpoints(
has_success = False
semaphore = asyncio.Semaphore(MAX_CONCURRENT_REQUESTS)
async def fetch_one(
client: httpx.AsyncClient, config: dict
) -> tuple[list, str | None, bool]:
async def fetch_one(client: httpx.AsyncClient, config: dict) -> tuple[list, str | None, bool]:
base_url = config["base_url"]
if not base_url:
return [], None, False

View File

@@ -10,14 +10,12 @@ from typing import Any
from sqlalchemy.orm import Session
from src.core.api_format import APIFormat
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
class CandidateResolver:
"""
候选解析器 - 负责获取和排序可用的 Provider 组合
@@ -45,7 +43,7 @@ class CandidateResolver:
async def fetch_candidates(
self,
api_format: APIFormat,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey | None = None,
@@ -78,6 +76,12 @@ class CandidateResolver:
provider_batch_size = 20
global_model_id: str | None = None
logger.debug(
"[CandidateResolver] fetch_candidates starting: model=%s, api_format=%s",
model_name,
api_format,
)
while True:
candidates, resolved_global_model_id = await self.cache_scheduler.list_all_candidates(
db=self.db,
@@ -91,6 +95,12 @@ class CandidateResolver:
capability_requirements=capability_requirements,
)
logger.debug(
"[CandidateResolver] list_all_candidates batch: offset=%d, returned=%d candidates",
provider_offset,
len(candidates),
)
if resolved_global_model_id and global_model_id is None:
global_model_id = resolved_global_model_id
@@ -100,6 +110,11 @@ class CandidateResolver:
all_candidates.extend(candidates)
provider_offset += provider_batch_size
logger.debug(
"[CandidateResolver] fetch_candidates completed: total=%d candidates",
len(all_candidates),
)
if not all_candidates:
logger.error(f" [{request_id}] 没有找到任何可用的 Provider/Endpoint/Key 组合")
request_type = "流式" if is_stream else "非流式"
@@ -196,7 +211,9 @@ class CandidateResolver:
candidate_record_map[(candidate_index, 0)] = record_id
else:
# max_retries 已从 Endpoint 迁移到 ProviderEndpoint 仍可能保留旧字段用于兼容)
max_retries_for_candidate = int(provider.max_retries or 2) if candidate.is_cached else 1
max_retries_for_candidate = (
int(provider.max_retries or 2) if candidate.is_cached else 1
)
for retry_index in range(max_retries_for_candidate):
record_id = str(uuid.uuid4())
@@ -226,7 +243,9 @@ class CandidateResolver:
)
self.db.flush()
logger.debug(f" [{request_id}] 批量插入完成: {len(candidate_records_to_insert)} 条记录")
logger.debug(
f" [{request_id}] 批量插入完成: {len(candidate_records_to_insert)} 条记录"
)
return candidate_record_map

View File

@@ -13,7 +13,7 @@ from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.core.api_format import APIFormat
from src.core.api_format.signature import make_signature_key
from src.core.exceptions import (
ConcurrencyLimitError,
ProviderAuthException,
@@ -28,7 +28,7 @@ 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_api_format
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
@@ -561,7 +561,7 @@ class ErrorClassifier:
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
affinity_key: str,
api_format: str | APIFormat,
api_format: str,
global_model_id: str,
request_id: str | None,
captured_key_concurrent: int | None,
@@ -618,18 +618,11 @@ class ErrorClassifier:
extra_data["error_response"] = error_response_text
# client_format用于缓存亲和性/缓存失效(用户视角)
client_format_str = (
normalize_api_format(api_format).value
if isinstance(api_format, (str, APIFormat))
else str(api_format)
)
client_format_str = normalize_endpoint_signature(api_format)
# provider_format用于健康度/熔断 bucketProvider 真实端点格式)
provider_api_format = getattr(endpoint, "api_format", None)
provider_format_str = (
provider_api_format.value
if isinstance(provider_api_format, APIFormat)
else str(provider_api_format or client_format_str)
).upper()
fam = str(getattr(endpoint, "api_family", "")).strip().lower()
kind = str(getattr(endpoint, "endpoint_kind", "")).strip().lower()
provider_format_str = make_signature_key(fam, kind) if fam and kind else client_format_str
# 处理客户端请求错误(不应重试,不失效缓存,不记录健康失败)
if isinstance(converted_error, UpstreamClientException):
@@ -704,7 +697,7 @@ class ErrorClassifier:
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
affinity_key: str,
api_format: str | APIFormat,
api_format: str,
global_model_id: str,
captured_key_concurrent: int | None,
elapsed_ms: int | None,
@@ -737,18 +730,11 @@ class ErrorClassifier:
)
# client_format用于缓存亲和性/缓存失效(用户视角)
client_format_str = (
normalize_api_format(api_format).value
if isinstance(api_format, (str, APIFormat))
else str(api_format)
)
client_format_str = normalize_endpoint_signature(api_format)
# provider_format用于健康度/熔断 bucketProvider 真实端点格式)
provider_api_format = getattr(endpoint, "api_format", None)
provider_format_str = (
provider_api_format.value
if isinstance(provider_api_format, APIFormat)
else str(provider_api_format or client_format_str)
).upper()
fam = str(getattr(endpoint, "api_family", "")).strip().lower()
kind = str(getattr(endpoint, "endpoint_kind", "")).strip().lower()
provider_format_str = make_signature_key(fam, kind) if fam and kind else client_format_str
# 处理限流错误
if isinstance(error, ProviderRateLimitException) and key:

View File

@@ -21,19 +21,16 @@
- 本类作为协调者,组合使用上述组件
"""
from __future__ import annotations
from typing import Any, NoReturn
from collections.abc import Callable
from typing import Any, NoReturn
import httpx
from redis import Redis
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.api_format import APIFormat
from src.core.api_format.conversion.exceptions import FormatConversionError
from src.core.error_utils import extract_error_message
from src.core.exceptions import (
@@ -51,7 +48,7 @@ from src.services.cache.aware_scheduler import (
get_cache_aware_scheduler,
)
from src.services.message.thinking_rectifier import ThinkingRectifier
from src.services.provider.format import normalize_api_format
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.concurrency_manager import get_concurrency_manager
from src.services.request.candidate import RequestCandidateService
@@ -172,7 +169,7 @@ class FallbackOrchestrator:
async def _fetch_all_candidates(
self,
api_format: APIFormat,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey | None = None,
@@ -256,7 +253,7 @@ class FallbackOrchestrator:
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None,
api_format: APIFormat,
api_format: str,
model_name: str,
affinity_key: str,
global_model_id: str,
@@ -343,8 +340,11 @@ class FallbackOrchestrator:
if not config.thinking_rectifier_enabled:
logger.info(f" [{request_id}] Thinking 错误:整流器已禁用,终止重试")
self._mark_thinking_error_failed(
candidate_record_id, converted_error, elapsed_ms,
captured_key_concurrent, serializable_extra_data
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
serializable_extra_data,
)
raise converted_error
@@ -352,8 +352,11 @@ class FallbackOrchestrator:
if request_body_ref is None:
logger.warning(f" [{request_id}] Thinking 错误:无法获取请求体引用,终止重试")
self._mark_thinking_error_failed(
candidate_record_id, converted_error, elapsed_ms,
captured_key_concurrent, serializable_extra_data
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
serializable_extra_data,
)
raise converted_error
@@ -361,8 +364,11 @@ class FallbackOrchestrator:
if request_body_ref.get("_rectified", False):
logger.warning(f" [{request_id}] Thinking 错误:已整流仍失败,终止重试")
self._mark_thinking_error_failed(
candidate_record_id, converted_error, elapsed_ms,
captured_key_concurrent, {**serializable_extra_data, "rectified": True}
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
{**serializable_extra_data, "rectified": True},
)
raise converted_error
@@ -383,8 +389,11 @@ class FallbackOrchestrator:
# 标记当前尝试为失败(整流前的状态)
# 注意:整流后重试会复用此记录 ID成功时会覆盖为 success 状态
self._mark_thinking_error_failed(
candidate_record_id, converted_error, elapsed_ms,
captured_key_concurrent, {**serializable_extra_data, "rectified": True}
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
{**serializable_extra_data, "rectified": True},
)
# 返回 continue在当前候选的重试循环中继续使用整流后的请求体重试
@@ -392,8 +401,11 @@ class FallbackOrchestrator:
else:
logger.warning(f" [{request_id}] Thinking 错误:无可整流内容")
self._mark_thinking_error_failed(
candidate_record_id, converted_error, elapsed_ms,
captured_key_concurrent, serializable_extra_data
candidate_record_id,
converted_error,
elapsed_ms,
captured_key_concurrent,
serializable_extra_data,
)
raise converted_error
@@ -425,7 +437,7 @@ class FallbackOrchestrator:
retry_index: int,
max_retries_for_candidate: int,
affinity_key: str,
api_format: APIFormat,
api_format: str,
global_model_id: str,
request_id: str | None,
attempt: int,
@@ -506,9 +518,7 @@ class FallbackOrchestrator:
"provider_id": str(provider.id),
"provider_endpoint_id": str(endpoint.id),
"provider_api_key_id": str(key.id),
"api_format": (
api_format.value if hasattr(api_format, "value") else str(api_format)
),
"api_format": str(api_format),
}
raise client_error
else:
@@ -588,9 +598,7 @@ class FallbackOrchestrator:
"provider_id": str(provider.id),
"provider_endpoint_id": str(endpoint.id),
"provider_api_key_id": str(key.id),
"api_format": (
api_format.value if hasattr(api_format, "value") else str(api_format)
),
"api_format": str(api_format),
}
raise converted_error
@@ -662,7 +670,7 @@ class FallbackOrchestrator:
user_api_key: ApiKey,
model_name: str,
is_stream: bool,
api_format_enum: APIFormat,
api_format: str,
) -> None:
"""创建 pending 状态的使用记录(用于实时状态追踪)"""
if not request_id:
@@ -681,7 +689,7 @@ class FallbackOrchestrator:
api_key=user_api_key,
model=model_name,
is_stream=is_stream,
api_format=api_format_enum.value,
api_format=api_format,
)
except Exception as e:
# 创建 pending 记录失败不应阻塞请求
@@ -694,7 +702,7 @@ class FallbackOrchestrator:
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None,
api_format_enum: APIFormat,
api_format: str,
model_name: str,
affinity_key: str,
global_model_id: str,
@@ -724,7 +732,7 @@ class FallbackOrchestrator:
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format_enum=api_format_enum,
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
@@ -735,9 +743,9 @@ class FallbackOrchestrator:
)
if result["success"]:
response: tuple[
Any, str, str | None, str | None, str | None, str | None
] = result["response"]
response: tuple[Any, str, str | None, str | None, str | None, str | None] = result[
"response"
]
return response
# 更新计数器和错误信息
@@ -746,14 +754,12 @@ class FallbackOrchestrator:
if result.get("error"):
last_error = result["error"]
if result.get("should_raise") and last_error is not None:
self._attach_metadata_to_error(
last_error, last_candidate, model_name, api_format_enum
)
self._attach_metadata_to_error(last_error, last_candidate, model_name, api_format)
raise last_error
# 所有组合都已尝试完毕,全部失败
self._raise_all_failed_exception(
request_id, max_attempts, last_candidate, model_name, api_format_enum, last_error
request_id, max_attempts, last_candidate, model_name, api_format, last_error
)
async def _try_candidate_with_retries(
@@ -764,7 +770,7 @@ class FallbackOrchestrator:
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None,
api_format_enum: APIFormat,
api_format: str,
model_name: str,
affinity_key: str,
global_model_id: str,
@@ -820,7 +826,7 @@ class FallbackOrchestrator:
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format=api_format_enum,
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,
@@ -839,7 +845,7 @@ class FallbackOrchestrator:
retry_index=retry_index,
max_retries_for_candidate=max_retries_for_candidate,
affinity_key=affinity_key,
api_format=api_format_enum,
api_format=api_format,
global_model_id=global_model_id,
request_id=request_id,
attempt=attempt_counter,
@@ -886,7 +892,7 @@ class FallbackOrchestrator:
error: Exception | None,
candidate: ProviderCandidate | None,
model_name: str,
api_format_enum: APIFormat,
api_format: str,
) -> None:
"""附加 candidate 信息到异常,以便记录 usage"""
if not error or not candidate:
@@ -915,7 +921,7 @@ class FallbackOrchestrator:
provider_api_key_id=(
getattr(existing_metadata, "provider_api_key_id", None) or str(candidate.key.id)
),
api_format=api_format_enum.value,
api_format=api_format,
)
# 使用 setattr 避免类型检查错误
setattr(error, "request_metadata", metadata)
@@ -926,7 +932,7 @@ class FallbackOrchestrator:
max_attempts: int,
last_candidate: ProviderCandidate | None,
model_name: str,
api_format_enum: APIFormat,
api_format: str,
last_error: Exception | None = None,
) -> NoReturn:
"""所有组合都失败时抛出异常"""
@@ -940,7 +946,7 @@ class FallbackOrchestrator:
"provider_id": str(last_candidate.provider.id),
"provider_endpoint_id": str(last_candidate.endpoint.id),
"provider_api_key_id": str(last_candidate.key.id),
"api_format": api_format_enum.value,
"api_format": api_format,
}
# 提取上游错误响应
@@ -988,7 +994,7 @@ class FallbackOrchestrator:
async def execute_with_fallback(
self,
api_format: str | APIFormat,
api_format: str,
model_name: str,
user_api_key: ApiKey,
request_func: Callable[[Provider, ProviderEndpoint, ProviderAPIKey], Any],
@@ -1002,7 +1008,7 @@ class FallbackOrchestrator:
执行请求,并在失败时自动故障转移(缓存感知)
Args:
api_format: API 格式(如 'CLAUDE', 'OPENAI'
api_format: endpoint signature'claude:chat', 'openai:cli'
model_name: 模型名称
user_api_key: 用户的 API Key对象
request_func: 请求函数,接收 (provider, endpoint, key) 参数,返回响应
@@ -1023,22 +1029,22 @@ class FallbackOrchestrator:
# 准备执行上下文
affinity_key = str(user_api_key.id)
user_id = str(user_api_key.user_id)
api_format_enum = normalize_api_format(api_format)
api_format_norm = normalize_endpoint_signature(api_format)
logger.debug(
f"[FallbackOrchestrator] execute_with_fallback 被调用: "
f"api_format={api_format_enum.value}, model_name={model_name}, "
f"api_format={api_format_norm}, model_name={model_name}, "
f"request_id={request_id}, is_stream={is_stream}"
)
# 创建 pending 状态的使用记录
self._create_pending_usage_record(
request_id, user_api_key, model_name, is_stream, api_format_enum
request_id, user_api_key, model_name, is_stream, api_format_norm
)
# 1. 收集所有候选(同时获取规范化的 global_model_id 用于缓存亲和性)
all_candidates, global_model_id = await self._fetch_all_candidates(
api_format=api_format_enum,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
@@ -1064,7 +1070,7 @@ class FallbackOrchestrator:
user_api_key=user_api_key,
request_func=request_func,
request_id=request_id,
api_format_enum=api_format_enum,
api_format=api_format_norm,
model_name=model_name,
affinity_key=affinity_key,
global_model_id=global_model_id,

View File

@@ -4,13 +4,11 @@
负责执行单个候选请求
"""
from typing import Any
from collections.abc import Callable
from typing import Any
from sqlalchemy.orm import Session
from src.core.api_format import APIFormat
from src.core.logger import logger
from src.models.database import ApiKey
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
@@ -18,7 +16,6 @@ from src.services.request.candidate import RequestCandidateService
from src.services.request.executor import RequestExecutor
class RequestDispatcher:
"""
请求分发器 - 负责执行单个候选请求
@@ -56,7 +53,7 @@ class RequestDispatcher:
user_api_key: ApiKey,
request_func: Callable[..., Any],
request_id: str | None,
api_format: APIFormat,
api_format: str,
model_name: str,
affinity_key: str,
global_model_id: str,
@@ -133,15 +130,12 @@ class RequestDispatcher:
# 设置缓存亲和性
if provider_supports_caching and self.cache_scheduler is not None:
try:
api_format_str = (
api_format.value if isinstance(api_format, APIFormat) else api_format
)
await self.cache_scheduler.set_cache_affinity(
affinity_key=affinity_key,
provider_id=provider_id,
endpoint_id=endpoint_id,
key_id=key_id,
api_format=api_format_str,
api_format=api_format,
global_model_id=global_model_id,
ttl=provider_cache_ttl_seconds,
)

View File

@@ -4,12 +4,12 @@ Provider 服务模块
包含 Provider 管理、格式处理、传输层等功能。
"""
from src.services.provider.format import normalize_api_format
from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.service import ProviderService
from src.services.provider.transport import build_provider_url
__all__ = [
"ProviderService",
"normalize_api_format",
"normalize_endpoint_signature",
"build_provider_url",
]

View File

@@ -1,18 +1,45 @@
"""
API 格式辅助函数,确保在调度/编排链路中使用统一的枚举值
Endpoint signature 辅助函数
调度/编排链路使用 endpoint signature key`family:kind`(如 "claude:chat", "openai:cli")。
"""
from __future__ import annotations
from src.core.api_format.enums import ApiFamily, EndpointKind
from src.core.api_format.signature import (
EndpointSignature,
make_signature_key,
normalize_signature_key,
)
DEFAULT_ENDPOINT_SIGNATURE: str = make_signature_key(ApiFamily.CLAUDE, EndpointKind.CHAT)
from src.core.api_format import APIFormat, resolve_api_format
def normalize_api_format(
value: str | APIFormat | None, default: APIFormat = APIFormat.CLAUDE
) -> APIFormat:
def normalize_endpoint_signature(
value: str | EndpointSignature | tuple[ApiFamily, EndpointKind] | None,
*,
default: str = DEFAULT_ENDPOINT_SIGNATURE,
) -> str:
"""
将任意字符串/枚举值归一化为 APIFormat
未识别的值回退到默认枚举(默认 CLAUDE
将任意输入归一化为 canonical signature key`family:kind`,小写)
不支持旧格式(如 "CLAUDE_CLI"),仅接受 `family:kind` 格式。
解析失败时返回默认值。
"""
resolved = resolve_api_format(value)
return resolved or default
if value is None:
return default
if isinstance(value, EndpointSignature):
return value.key
if isinstance(value, tuple) and len(value) == 2:
fam, kind = value
if isinstance(fam, ApiFamily) and isinstance(kind, EndpointKind):
return make_signature_key(fam, kind)
return default
if isinstance(value, str):
try:
return normalize_signature_key(value)
except ValueError:
# 如果解析失败,返回默认值
return default
return default

View File

@@ -3,8 +3,8 @@
负责提供商选择、模型映射和请求处理
"""
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import GlobalModel, Model, Provider

View File

@@ -13,8 +13,13 @@ import re
from typing import TYPE_CHECKING, Any
from urllib.parse import urlencode
from src.core.api_format import APIFormat, get_default_path, resolve_api_format
from src.core.api_format import (
EndpointKind,
get_default_path_for_endpoint,
make_signature_key,
)
from src.core.logger import logger
from src.services.provider.format import normalize_endpoint_signature
if TYPE_CHECKING:
from src.models.database import ProviderAPIKey, ProviderEndpoint
@@ -102,13 +107,29 @@ def build_provider_url(
decrypted_auth_config=decrypted_auth_config,
)
# 准备路径参数,添加 Gemini API 所需的 action 参数
effective_path_params = dict(path_params) if path_params else {}
# endpoint signature新模式
raw_family = getattr(endpoint, "api_family", None)
raw_kind = getattr(endpoint, "endpoint_kind", None)
endpoint_sig = ""
if isinstance(raw_family, str) and isinstance(raw_kind, str) and raw_family and raw_kind:
endpoint_sig = make_signature_key(raw_family, raw_kind)
else:
# 兜底:允许 api_format 已直接存 signature key 的情况
raw_format = getattr(endpoint, "api_format", None)
if isinstance(raw_format, str) and ":" in raw_format:
endpoint_sig = raw_format
# 为 Gemini API 格式自动添加 action 参数
resolved_format = resolve_api_format(endpoint.api_format)
if resolved_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
if "action" not in effective_path_params:
# endpoint_sig 为空时保持为空(更安全:默认路径回退到 "/",避免误判为 claude:chat
endpoint_sig = normalize_endpoint_signature(endpoint_sig) if endpoint_sig else ""
# 准备路径参数Gemini chat/cli 需要 action
effective_path_params = dict(path_params) if path_params else {}
if endpoint_sig.startswith("gemini:"):
try:
kind = EndpointKind(endpoint_sig.split(":", 1)[1])
except Exception:
kind = None
if kind in {EndpointKind.CHAT, EndpointKind.CLI} and "action" not in effective_path_params:
effective_path_params["action"] = (
"streamGenerateContent" if is_stream else "generateContent"
)
@@ -124,7 +145,7 @@ def build_provider_url(
pass
else:
# 使用 API 格式的默认路径
path = _resolve_default_path(endpoint.api_format)
path = _resolve_default_path(endpoint_sig)
if effective_path_params:
try:
path = path.format(**effective_path_params)
@@ -143,9 +164,9 @@ def build_provider_url(
# 合并查询参数
effective_query_params = dict(query_params) if query_params else {}
# Gemini 格式下清除可能存在的 key 参数(避免客户端传入的认证信息泄露到上游)
# Gemini family 下清除可能存在的 key 参数(避免客户端传入的认证信息泄露到上游)
# 上游认证始终使用 header 方式,不使用 URL 参数
if resolved_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
if endpoint_sig.startswith("gemini:"):
effective_query_params.pop("key", None)
# Gemini streamGenerateContent 官方支持 `?alt=sse` 返回 SSEdata: {...})。
# 网关侧统一使用 SSE 输出,优先向上游请求 SSE 以减少解析分支;同时保留 JSON-array 兜底解析。
@@ -161,16 +182,13 @@ def build_provider_url(
return url
def _resolve_default_path(api_format: str | None) -> str:
"""
根据 API 格式返回默认路径
"""
resolved = resolve_api_format(api_format)
if resolved:
return get_default_path(resolved)
logger.warning(f"Unknown api_format '{api_format}' for endpoint, fallback to '/'")
return "/"
def _resolve_default_path(endpoint_sig: str | None) -> str:
"""根据 endpoint signature 返回默认路径。"""
try:
return get_default_path_for_endpoint(endpoint_sig or "")
except Exception:
logger.warning(f"Unknown endpoint signature '{endpoint_sig}' for endpoint, fallback to '/'")
return "/"
# ==============================================================================
@@ -179,15 +197,15 @@ def _resolve_default_path(api_format: str | None) -> str:
# Vertex AI 模型前缀到 API 格式的映射
# 用于 auth_type=vertex_ai 时,根据模型名动态确定实际的请求/响应格式
# 格式:前缀 -> APIFormat 值
# 格式:前缀 -> endpoint signaturefamily:kind
VERTEX_AI_MODEL_FORMAT_MAPPING: dict[str, str] = {
"claude-": "CLAUDE", # Anthropic Claude 模型
"gemini-": "GEMINI", # Google Gemini 模型
"imagen-": "GEMINI", # Google Imagen 模型(使用 Gemini 格式)
"claude-": "claude:chat", # Anthropic Claude 模型
"gemini-": "gemini:chat", # Google Gemini 模型
"imagen-": "gemini:chat", # Google Imagen 模型(使用 Gemini chat 格式)
}
# Vertex AI 默认 API 格式(当模型前缀不匹配时)
VERTEX_AI_DEFAULT_FORMAT: str = "GEMINI"
# Vertex AI 默认 endpoint signature(当模型前缀不匹配时)
VERTEX_AI_DEFAULT_FORMAT: str = "gemini:chat"
def get_vertex_ai_effective_format(
@@ -220,7 +238,7 @@ def get_vertex_ai_effective_format(
auth_config: 解密后的认证配置(可选),可包含 model_format_mapping 和 default_format
Returns:
实际应使用的 API 格式(如 "CLAUDE", "GEMINI"
实际应使用的 endpoint signature"claude:chat", "gemini:chat"
"""
# 用户配置的模型-格式映射
user_format_mapping: dict[str, str] = {}
@@ -232,21 +250,39 @@ def get_vertex_ai_effective_format(
# 1. 用户配置:精确匹配
if model in user_format_mapping:
return user_format_mapping[model].upper()
try:
return normalize_endpoint_signature(user_format_mapping[model])
except Exception:
logger.warning(
"Invalid vertex_ai model_format_mapping value for model '%s': %r",
model,
user_format_mapping[model],
)
# 2. 用户配置:前缀匹配
for prefix, api_format in user_format_mapping.items():
if prefix.endswith("-") and model.startswith(prefix):
return api_format.upper()
try:
return normalize_endpoint_signature(api_format)
except Exception:
logger.warning(
"Invalid vertex_ai model_format_mapping value for prefix '%s': %r",
prefix,
api_format,
)
break
# 3. 内置配置:前缀匹配
for prefix, api_format in VERTEX_AI_MODEL_FORMAT_MAPPING.items():
if model.startswith(prefix):
return api_format
return normalize_endpoint_signature(api_format)
# 4. 用户默认格式
if user_default_format:
return user_default_format.upper()
try:
return normalize_endpoint_signature(user_default_format)
except Exception:
logger.warning("Invalid vertex_ai default_format: %r", user_default_format)
# 5. 内置默认格式
return VERTEX_AI_DEFAULT_FORMAT

View File

@@ -163,8 +163,15 @@ class NewApiBalanceAction(BalanceAction):
# 检查是否是认证失败(未登录、无权限等)- Cookie 已失效
auth_fail_indicators = [
"未登录", "请登录", "login", "unauthorized", "无权限", "权限不足",
"turnstile", "captcha", "验证码", # 需要人机验证
"未登录",
"请登录",
"login",
"unauthorized",
"无权限",
"权限不足",
"turnstile",
"captcha",
"验证码", # 需要人机验证
]
is_auth_fail = any(ind in message.lower() for ind in auth_fail_indicators)
if is_auth_fail:

View File

@@ -2,12 +2,12 @@
Provider 架构模块
"""
from src.services.provider_ops.architectures.anyrouter import AnyrouterArchitecture
from src.services.provider_ops.architectures.base import (
ProviderArchitecture,
ProviderConnector,
VerifyResult,
)
from src.services.provider_ops.architectures.anyrouter import AnyrouterArchitecture
from src.services.provider_ops.architectures.cubence import CubenceArchitecture
from src.services.provider_ops.architectures.generic_api import GenericApiArchitecture
from src.services.provider_ops.architectures.nekocode import NekoCodeArchitecture

View File

@@ -11,7 +11,6 @@ from typing import Any
import httpx
from src.core.logger import logger
from src.utils.ssl_utils import get_ssl_context
from src.services.provider_ops.actions import (
AnyrouterBalanceAction,
ProviderAction,
@@ -22,6 +21,7 @@ from src.services.provider_ops.architectures.base import (
VerifyResult,
)
from src.services.provider_ops.types import ConnectorAuthType, ProviderActionType
from src.utils.ssl_utils import get_ssl_context
# acw_sc__v2 算法常量
_XOR_KEY = "3000176000856006061501533003690027800375"

View File

@@ -3,23 +3,22 @@ Provider 架构抽象基类
"""
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from collections.abc import AsyncIterator
import httpx
from src.services.provider_ops.actions.base import ProviderAction
from src.utils.ssl_utils import get_ssl_context
from src.services.provider_ops.types import (
ConnectorAuthType,
ConnectorState,
ConnectorStatus,
ProviderActionType,
)
from src.utils.ssl_utils import get_ssl_context
# ==================== 连接器基类 ====================
@@ -488,8 +487,6 @@ class ProviderArchitecture(ABC):
for a in self.supported_actions
],
"default_connector": (
self.supported_connectors[0].auth_type.value
if self.supported_connectors
else None
self.supported_connectors[0].auth_type.value if self.supported_connectors else None
),
}

View File

@@ -80,7 +80,16 @@ class ProviderOpsService:
"""
# 凭据中需要加密的字段
SENSITIVE_FIELDS = {"api_key", "password", "session_token", "session_cookie", "token_cookie", "auth_cookie", "cookie_string", "cookie"}
SENSITIVE_FIELDS = {
"api_key",
"password",
"session_token",
"session_cookie",
"token_cookie",
"auth_cookie",
"cookie_string",
"cookie",
}
def __init__(self, db: Session):
self.db = db
@@ -385,9 +394,7 @@ class ProviderOpsService:
Returns:
操作结果
"""
result = await self.execute_action(
provider_id, ProviderActionType.QUERY_BALANCE, config
)
result = await self.execute_action(provider_id, ProviderActionType.QUERY_BALANCE, config)
# 成功或 auth_expired 时缓存auth_expired 带有 cookie_expired 信息供前端显示警告)
if result.status in (ActionStatus.SUCCESS, ActionStatus.AUTH_EXPIRED) and result.data:
@@ -485,7 +492,9 @@ class ProviderOpsService:
"response_time_ms": result.response_time_ms,
}
await CacheService.set(cache_key, cache_data, AUTH_FAILED_CACHE_TTL)
logger.info(f"余额缓存已写入(认证失败): provider_id={provider_id}, message={result.message}")
logger.info(
f"余额缓存已写入(认证失败): provider_id={provider_id}, message={result.message}"
)
async def _cache_balance(self, provider_id: str, result: ActionResult) -> None:
"""缓存余额结果"""
@@ -504,7 +513,9 @@ class ProviderOpsService:
}
await CacheService.set(cache_key, cache_data, BALANCE_CACHE_TTL)
logger.debug(f"余额缓存已写入: provider_id={provider_id}, extra={data.get('extra') if data else None}")
logger.debug(
f"余额缓存已写入: provider_id={provider_id}, extra={data.get('extra') if data else None}"
)
async def _cache_balance_from_verify(
self,
@@ -633,7 +644,9 @@ class ProviderOpsService:
if key in self.SENSITIVE_FIELDS and isinstance(value, str):
if value: # 只加密非空值
encrypted[key] = self.crypto.encrypt(value)
logger.debug(f"加密字段 {key}: 原始长度={len(value)}, 加密后长度={len(encrypted[key])}")
logger.debug(
f"加密字段 {key}: 原始长度={len(value)}, 加密后长度={len(encrypted[key])}"
)
else:
logger.warning(f"跳过空值字段 {key}")
encrypted[key] = value
@@ -705,8 +718,14 @@ class ProviderOpsService:
if saved_config:
saved_credentials = self._decrypt_credentials(saved_config.connector_credentials)
sensitive_fields = [
"api_key", "password", "session_token", "cookie_string", "cookie",
"token_cookie", "auth_cookie", "session_cookie", # Cookie 认证字段
"api_key",
"password",
"session_token",
"cookie_string",
"cookie",
"token_cookie",
"auth_cookie",
"session_cookie", # Cookie 认证字段
]
for field in sensitive_fields:
@@ -738,11 +757,7 @@ class ProviderOpsService:
if provider_ids is None:
# 查询所有已配置的 Provider
providers = self.db.query(Provider).filter(Provider.is_active.is_(True)).all()
provider_ids = [
p.id
for p in providers
if p.config and p.config.get("provider_ops")
]
provider_ids = [p.id for p in providers if p.config and p.config.get("provider_ops")]
if not provider_ids:
return {}

View File

@@ -4,8 +4,8 @@
包含自适应 RPM 控制、并发管理、IP限流等功能。
"""
from src.services.rate_limit.adaptive_rpm import AdaptiveConcurrencyManager # 向后兼容别名
from src.services.rate_limit.adaptive_rpm import (
AdaptiveConcurrencyManager, # 向后兼容别名
AdaptiveRPMManager,
get_adaptive_rpm_manager,
)

View File

@@ -17,7 +17,6 @@ from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any
from src.config.constants import AdaptiveReservationDefaults
if TYPE_CHECKING:
@@ -165,9 +164,7 @@ class AdaptiveReservationManager:
return request_count
def _calculate_load_ratio(
self, current_usage: int, effective_limit: int | None
) -> float:
def _calculate_load_ratio(self, current_usage: int, effective_limit: int | None) -> float:
"""计算当前负载率"""
if not effective_limit or effective_limit <= 0:
return 0.0

View File

@@ -110,9 +110,7 @@ class AdaptiveRPMManager:
is_adaptive = key.rpm_limit is None
if not is_adaptive:
logger.debug(
f"Key {key.id} 设置了固定 RPM 限制 ({key.rpm_limit}),跳过自适应调整"
)
logger.debug(f"Key {key.id} 设置了固定 RPM 限制 ({key.rpm_limit}),跳过自适应调整")
return int(key.rpm_limit) # type: ignore[arg-type]
# 更新429统计
@@ -318,7 +316,7 @@ class AdaptiveRPMManager:
# 限制采样数量
if len(samples) > self.UTILIZATION_WINDOW_SIZE:
samples = samples[-self.UTILIZATION_WINDOW_SIZE:]
samples = samples[-self.UTILIZATION_WINDOW_SIZE :]
# 更新到 key 对象
key.utilization_samples = samples # type: ignore[assignment]
@@ -538,7 +536,7 @@ class AdaptiveRPMManager:
# 保留最近N条记录
if len(history) > self.MAX_HISTORY_RECORDS:
history = history[-self.MAX_HISTORY_RECORDS:]
history = history[-self.MAX_HISTORY_RECORDS :]
key.adjustment_history = history # type: ignore[assignment]

View File

@@ -10,12 +10,12 @@ RPM 限制管理器 - 支持 Redis 或内存的 Key 级别 RPM 限制
from __future__ import annotations
from typing import Any
import asyncio
import math
import os
import time
from contextlib import asynccontextmanager
from typing import Any
import redis.asyncio as aioredis
@@ -75,7 +75,9 @@ class ConcurrencyManager:
if self._redis:
logger.info("[OK] ConcurrencyManager 已复用全局 Redis 客户端")
else:
logger.warning("[WARN] Redis 不可用RPM 限制降级为内存模式(仅在单实例环境下安全)")
logger.warning(
"[WARN] Redis 不可用RPM 限制降级为内存模式(仅在单实例环境下安全)"
)
# 内存模式下启动后台清理任务
self._start_background_cleanup()
except Exception as e:
@@ -191,14 +193,11 @@ class ConcurrencyManager:
evict_count = max(1, self._max_memory_rpm_entries // 5)
# 按 bucket时间排序删除最旧的
sorted_keys = sorted(
self._memory_key_rpm_counts.items(),
key=lambda x: x[1][0] # 按 bucket 排序
self._memory_key_rpm_counts.items(), key=lambda x: x[1][0] # 按 bucket 排序
)
for k, _ in sorted_keys[:evict_count]:
del self._memory_key_rpm_counts[k]
logger.warning(
f"[WARN] 内存 RPM 计数器达到上限,已淘汰 {evict_count} 个最旧条目"
)
logger.warning(f"[WARN] 内存 RPM 计数器达到上限,已淘汰 {evict_count} 个最旧条目")
self._memory_key_rpm_counts[key_id] = (bucket, count)
def _cleanup_expired_memory_rpm_counts(self, current_bucket: int, force: bool = False) -> None:
@@ -487,9 +486,7 @@ class ConcurrencyManager:
from src.core.exceptions import ConcurrencyLimitError
user_type = "缓存用户" if is_cached_user else "新用户"
raise ConcurrencyLimitError(
f"RPM 限制已达上限: key={key_id}, 类型={user_type}"
)
raise ConcurrencyLimitError(f"RPM 限制已达上限: key={key_id}, 类型={user_type}")
# 记录开始时间和状态
import time

View File

@@ -147,14 +147,12 @@ class RateLimitDetector:
and retry_after <= 30
):
is_likely_concurrent = True
concurrent_reason = f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
concurrent_reason = (
f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
)
# 条件 B无 remaining 头但 retry_after 很短(<= 5 秒)
elif (
requests_remaining is None
and retry_after is not None
and retry_after <= 5
):
elif requests_remaining is None and retry_after is not None and retry_after <= 5:
is_likely_concurrent = True
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
@@ -245,12 +243,10 @@ class RateLimitDetector:
and retry_after <= 30
):
is_likely_concurrent = True
concurrent_reason = f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
elif (
requests_remaining is None
and retry_after is not None
and retry_after <= 5
):
concurrent_reason = (
f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
)
elif requests_remaining is None and retry_after is not None and retry_after <= 5:
is_likely_concurrent = True
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
@@ -330,11 +326,7 @@ class RateLimitDetector:
):
is_likely_concurrent = True
concurrent_reason = f"remaining={remaining} > 0, retry_after={retry_after}s <= 30s"
elif (
remaining is None
and retry_after is not None
and retry_after <= 5
):
elif remaining is None and retry_after is not None and retry_after <= 5:
is_likely_concurrent = True
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"

View File

@@ -12,7 +12,6 @@ from src.clients.redis_client import get_redis_client
from src.core.logger import logger
class IPRateLimiter:
"""IP 速率限制服务"""
@@ -91,7 +90,9 @@ class IPRateLimiter:
allowed = count <= rate_limit
if not allowed:
logger.warning(f"IP 速率限制触发: {ip_address}, 类型: {endpoint_type}, 计数: {count}/{rate_limit}")
logger.warning(
f"IP 速率限制触发: {ip_address}, 类型: {endpoint_type}, 计数: {count}/{rate_limit}"
)
return allowed, remaining, ttl

View File

@@ -351,7 +351,9 @@ class RequestCandidateService:
候选自身的 TTFB毫秒如果计算失败则返回 global_first_byte_time_ms
"""
try:
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
candidate = (
db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
)
if candidate and candidate.started_at:
started_at = candidate.started_at
if started_at.tzinfo is None:

View File

@@ -5,18 +5,17 @@
from __future__ import annotations
import time
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
from collections.abc import Callable
from sqlalchemy.orm import Session
from src.core.api_format import APIFormat
from src.core.api_format.signature import make_signature_key
from src.core.exceptions import ConcurrencyLimitError
from src.core.logger import logger
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_api_format
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
from src.services.request.candidate import RequestCandidateService
@@ -64,7 +63,7 @@ class RequestExecutor:
user_api_key: Any,
request_func: Callable[..., Any],
request_id: str | None,
api_format: str | APIFormat,
api_format: str,
model_name: str,
is_stream: bool = False,
) -> ExecutionResult:
@@ -148,16 +147,11 @@ class RequestExecutor:
context.elapsed_ms = int((time.time() - context.start_time) * 1000)
provider_api_format = getattr(endpoint, "api_format", None)
provider_format_str = (
provider_api_format.value
if isinstance(provider_api_format, APIFormat)
else str(provider_api_format or "")
)
client_format_str = (
api_format.value if isinstance(api_format, APIFormat) else str(api_format)
)
health_format = normalize_api_format(provider_format_str or client_format_str).value
fam = str(getattr(endpoint, "api_family", "")).strip().lower()
kind = str(getattr(endpoint, "endpoint_kind", "")).strip().lower()
provider_format_str = make_signature_key(fam, kind) if fam and kind else ""
client_format_str = normalize_endpoint_signature(api_format)
health_format = provider_format_str or client_format_str
health_monitor.record_success(
db=self.db,
@@ -196,11 +190,7 @@ class RequestExecutor:
extra_data={
"is_cached_user": is_cached_user,
"model_name": model_name,
"api_format": (
api_format.value
if isinstance(api_format, APIFormat)
else api_format
),
"api_format": api_format,
},
)

View File

@@ -14,10 +14,11 @@
"""
from __future__ import annotations
from collections.abc import AsyncIterator
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
from collections.abc import AsyncIterator
class RequestStatus(Enum):
@@ -276,7 +277,9 @@ class RequestResult:
provider_id=get_meta_value(existing_metadata, "provider_id"),
provider_endpoint_id=get_meta_value(existing_metadata, "provider_endpoint_id"),
provider_api_key_id=get_meta_value(existing_metadata, "provider_api_key_id"),
provider_request_headers=get_meta_value(existing_metadata, "provider_request_headers", {}),
provider_request_headers=get_meta_value(
existing_metadata, "provider_request_headers", {}
),
provider_response_headers=get_meta_value(
existing_metadata, "provider_response_headers", {}
),

View File

@@ -6,12 +6,12 @@
from src.services.system.announcement import AnnouncementService
from src.services.system.audit import AuditService
from src.services.system.config import SystemConfigService
from src.services.system.maintenance_scheduler import CleanupScheduler # 兼容旧名称
from src.services.system.maintenance_scheduler import (
CleanupScheduler, # 兼容旧名称
MaintenanceScheduler,
get_maintenance_scheduler,
)
from src.services.system.config import SystemConfigService
from src.services.system.scheduler import APP_TIMEZONE, TaskScheduler, get_scheduler
from src.services.system.sync_stats import SyncStatsService

View File

@@ -14,7 +14,6 @@ from src.core.logger import logger
from src.models.database import Announcement, AnnouncementRead, User, UserRole
class AnnouncementService:
"""公告系统服务"""

View File

@@ -14,8 +14,6 @@ from src.core.logger import logger
from src.database import get_db
from src.models.database import AuditEventType, AuditLog
# 审计模型已移至 src/models/database.py

View File

@@ -61,7 +61,9 @@ class CacheWarmupService:
elapsed = time.time() - start_time
if error_count > 0:
logger.warning(f"缓存预热完成: {success_count}/3 成功, {error_count} 失败, 耗时 {elapsed:.2f}s")
logger.warning(
f"缓存预热完成: {success_count}/3 成功, {error_count} 失败, 耗时 {elapsed:.2f}s"
)
else:
logger.info(f"缓存预热完成: {success_count}/3 成功, 耗时 {elapsed:.2f}s")

View File

@@ -227,7 +227,9 @@ class SystemConfigService:
return result
@staticmethod
def set_config(db: Session, key: str, value: Any, description: str | None = None) -> SystemConfig:
def set_config(
db: Session, key: str, value: Any, description: str | None = None
) -> SystemConfig:
"""设置系统配置值"""
config = db.query(SystemConfig).filter(SystemConfig.key == key).first()

View File

@@ -14,9 +14,9 @@
from __future__ import annotations
from typing import Any
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import delete
from sqlalchemy.orm import Session
@@ -196,10 +196,12 @@ class MaintenanceScheduler:
logger.info("开始执行统计数据聚合...")
from src.models.database import StatsDaily, User as DBUser
from src.services.system.scheduler import APP_TIMEZONE
from zoneinfo import ZoneInfo
from src.models.database import StatsDaily
from src.models.database import User as DBUser
from src.services.system.scheduler import APP_TIMEZONE
# 使用业务时区计算日期,确保与定时任务触发时间一致
# 定时任务在 Asia/Shanghai 凌晨 1 点触发,此时应聚合 Asia/Shanghai 的"昨天"
app_tz = ZoneInfo(APP_TIMEZONE)
@@ -227,9 +229,9 @@ class MaintenanceScheduler:
from src.models.database import StatsDailyModel, StatsDailyProvider
yesterday_business_date = today_local.date() - timedelta(days=1)
max_backfill_days: int = SystemConfigService.get_config(
db, "max_stats_backfill_days", 30
) or 30
max_backfill_days: int = (
SystemConfigService.get_config(db, "max_stats_backfill_days", 30) or 30
)
# 计算回填检查的起始日期
check_start_date = yesterday_business_date - timedelta(
@@ -287,7 +289,9 @@ class MaintenanceScheduler:
# 需要回填 StatsDailyProvider 的日期
missing_provider_dates = all_dates - existing_provider_dates
# 合并所有需要处理的日期
dates_to_process = missing_daily_dates | missing_model_dates | missing_provider_dates
dates_to_process = (
missing_daily_dates | missing_model_dates | missing_provider_dates
)
if dates_to_process:
sorted_dates = sorted(dates_to_process)
@@ -298,9 +302,7 @@ class MaintenanceScheduler:
f"StatsDailyProvider 缺失 {len(missing_provider_dates)} 天)"
)
users = (
db.query(DBUser.id).filter(DBUser.is_active.is_(True)).all()
)
users = db.query(DBUser.id).filter(DBUser.is_active.is_(True)).all()
failed_dates = 0
failed_users = 0
@@ -496,15 +498,9 @@ class MaintenanceScheduler:
# 获取所有已配置 provider_ops 的活跃 Provider只查询需要的字段
providers = (
db.query(Provider.id, Provider.config)
.filter(Provider.is_active.is_(True))
.all()
db.query(Provider.id, Provider.config).filter(Provider.is_active.is_(True)).all()
)
provider_ids = [
p.id
for p in providers
if p.config and p.config.get("provider_ops")
]
provider_ids = [p.id for p in providers if p.config and p.config.get("provider_ops")]
if not provider_ids:
logger.info("无已配置的 Provider跳过签到任务")
@@ -548,9 +544,7 @@ class MaintenanceScheduler:
# 统计结果
success_count = sum(1 for _, success, _ in results if success)
logger.info(
f"Provider 签到完成: {success_count}/{len(provider_ids)} 成功"
)
logger.info(f"Provider 签到完成: {success_count}/{len(provider_ids)} 成功")
# 记录详细结果
for provider_id, success, message in results:
@@ -694,12 +688,12 @@ class MaintenanceScheduler:
.values(
request_body=null(),
response_body=null(),
request_body_compressed=compress_json(req_body)
if req_body
else None,
response_body_compressed=compress_json(resp_body)
if resp_body
else None,
request_body_compressed=(
compress_json(req_body) if req_body else None
),
response_body_compressed=(
compress_json(resp_body) if resp_body else None
),
)
)
if result.rowcount > 0:
@@ -858,9 +852,7 @@ class MaintenanceScheduler:
break
total_cleaned += rows_updated
logger.debug(
f"已清理 {rows_updated} 条记录的 header 字段,累计 {total_cleaned}"
)
logger.debug(f"已清理 {rows_updated} 条记录的 header 字段,累计 {total_cleaned}")
await asyncio.sleep(0.1)

View File

@@ -43,9 +43,7 @@ def _get_business_day_range(date: datetime) -> tuple[datetime, datetime]:
app_tz = ZoneInfo(APP_TIMEZONE)
# 取日期部分,构造业务时区的当天 00:00:00
day_start_local = datetime(
date.year, date.month, date.day, 0, 0, 0, tzinfo=app_tz
)
day_start_local = datetime(date.year, date.month, date.day, 0, 0, 0, tzinfo=app_tz)
day_end_local = day_start_local + timedelta(days=1)
# 转换为 UTC
@@ -91,11 +89,9 @@ class StatsAggregatorService:
"unique_providers": 0,
}
error_requests = (
base_query.filter(
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
).count()
)
error_requests = base_query.filter(
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
).count()
aggregated = (
db.query(
@@ -159,13 +155,17 @@ class StatsAggregatorService:
"error_requests": error_requests,
"input_tokens": int(aggregated.input_tokens or 0) if aggregated else 0,
"output_tokens": int(aggregated.output_tokens or 0) if aggregated else 0,
"cache_creation_tokens": int(aggregated.cache_creation_tokens or 0) if aggregated else 0,
"cache_creation_tokens": (
int(aggregated.cache_creation_tokens or 0) if aggregated else 0
),
"cache_read_tokens": int(aggregated.cache_read_tokens or 0) if aggregated else 0,
"total_cost": float(aggregated.total_cost or 0) if aggregated else 0.0,
"actual_total_cost": float(aggregated.actual_total_cost or 0) if aggregated else 0.0,
"input_cost": float(aggregated.input_cost or 0) if aggregated else 0.0,
"output_cost": float(aggregated.output_cost or 0) if aggregated else 0.0,
"cache_creation_cost": float(aggregated.cache_creation_cost or 0) if aggregated else 0.0,
"cache_creation_cost": (
float(aggregated.cache_creation_cost or 0) if aggregated else 0.0
),
"cache_read_cost": float(aggregated.cache_read_cost or 0) if aggregated else 0.0,
"avg_response_time_ms": float(aggregated.avg_response_time or 0) if aggregated else 0.0,
"fallback_count": fallback_count,
@@ -219,7 +219,9 @@ class StatsAggregatorService:
db.commit()
# 日志使用业务日期(输入参数),而不是 UTC 日期
logger.info(f"[StatsAggregator] 聚合日期 {date.date()} 完成: {computed['total_requests']} 请求")
logger.info(
f"[StatsAggregator] 聚合日期 {date.date()} 完成: {computed['total_requests']} 请求"
)
return stats
@staticmethod
@@ -259,16 +261,16 @@ class StatsAggregatorService:
existing = (
db.query(StatsDailyModel)
.filter(and_(StatsDailyModel.date == day_start, StatsDailyModel.model == stat.model))
.filter(
and_(StatsDailyModel.date == day_start, StatsDailyModel.model == stat.model)
)
.first()
)
if existing:
record = existing
else:
record = StatsDailyModel(
id=str(uuid.uuid4()), date=day_start, model=stat.model
)
record = StatsDailyModel(id=str(uuid.uuid4()), date=day_start, model=stat.model)
record.total_requests = stat.total_requests or 0
record.input_tokens = int(stat.input_tokens or 0)
@@ -283,9 +285,7 @@ class StatsAggregatorService:
results.append(record)
db.commit()
logger.info(
f"[StatsAggregator] 聚合日期 {date.date()} 模型统计完成: {len(results)} 个模型"
)
logger.info(f"[StatsAggregator] 聚合日期 {date.date()} 模型统计完成: {len(results)} 个模型")
return results
@staticmethod
@@ -322,7 +322,12 @@ class StatsAggregatorService:
for stat in provider_stats:
existing = (
db.query(StatsDailyProvider)
.filter(and_(StatsDailyProvider.date == day_start, StatsDailyProvider.provider_name == stat.provider_name))
.filter(
and_(
StatsDailyProvider.date == day_start,
StatsDailyProvider.provider_name == stat.provider_name,
)
)
.first()
)
@@ -384,24 +389,28 @@ class StatsAggregatorService:
date_utc = stat.date.astimezone(timezone.utc)
date_str = date_utc.astimezone(app_tz).date().isoformat()
result.append({
"date": date_str,
"model": stat.model,
"requests": stat.total_requests,
"tokens": (
stat.input_tokens + stat.output_tokens +
stat.cache_creation_tokens + stat.cache_read_tokens
),
"cost": stat.total_cost,
"avg_response_time": stat.avg_response_time_ms / 1000.0 if stat.avg_response_time_ms else 0,
})
result.append(
{
"date": date_str,
"model": stat.model,
"requests": stat.total_requests,
"tokens": (
stat.input_tokens
+ stat.output_tokens
+ stat.cache_creation_tokens
+ stat.cache_read_tokens
),
"cost": stat.total_cost,
"avg_response_time": (
stat.avg_response_time_ms / 1000.0 if stat.avg_response_time_ms else 0
),
}
)
return result
@staticmethod
def aggregate_user_daily_stats(
db: Session, user_id: str, date: datetime
) -> StatsUserDaily:
def aggregate_user_daily_stats(db: Session, user_id: str, date: datetime) -> StatsUserDaily:
"""聚合指定用户指定日期的统计数据"""
# 将业务日期转换为 UTC 时间范围
day_start, day_end = _get_business_day_range(date)
@@ -443,11 +452,9 @@ class StatsAggregatorService:
db.commit()
return stats
error_requests = (
base_query.filter(
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
).count()
)
error_requests = base_query.filter(
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
).count()
aggregated = (
db.query(
@@ -581,11 +588,9 @@ class StatsAggregatorService:
"actual_total_cost": 0.0,
}
error_requests = (
base_query.filter(
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
).count()
)
error_requests = base_query.filter(
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
).count()
aggregated = (
db.query(
@@ -624,8 +629,7 @@ class StatsAggregatorService:
return {
"total_requests": summary.all_time_requests + today_stats["total_requests"],
"success_requests": summary.all_time_success_requests
+ today_stats["success_requests"],
"success_requests": summary.all_time_success_requests + today_stats["success_requests"],
"error_requests": summary.all_time_error_requests + today_stats["error_requests"],
"input_tokens": summary.all_time_input_tokens + today_stats["input_tokens"],
"output_tokens": summary.all_time_output_tokens + today_stats["output_tokens"],

View File

@@ -3,7 +3,6 @@ API密钥统计同步服务
定期同步API密钥的统计数据确保与实际使用记录一致
"""
from __future__ import annotations
from sqlalchemy import func
@@ -13,7 +12,6 @@ from src.core.logger import logger
from src.models.database import ApiKey, Usage
class SyncStatsService:
"""API密钥统计同步服务"""
@@ -75,12 +73,16 @@ class SyncStatsService:
# 检查是否需要更新
needs_update = False
if api_key.total_requests != actual_requests:
logger.info(f"API密钥 {api_key.id} 请求数不一致: {api_key.total_requests} -> {actual_requests}")
logger.info(
f"API密钥 {api_key.id} 请求数不一致: {api_key.total_requests} -> {actual_requests}"
)
api_key.total_requests = actual_requests
needs_update = True
if abs(api_key.total_cost_usd - actual_cost) > 0.0001:
logger.info(f"API密钥 {api_key.id} 费用不一致: {api_key.total_cost_usd} -> {actual_cost}")
logger.info(
f"API密钥 {api_key.id} 费用不一致: {api_key.total_cost_usd} -> {actual_cost}"
)
api_key.total_cost_usd = actual_cost
needs_update = True
@@ -104,7 +106,9 @@ class SyncStatsService:
# 提交所有更改
db.commit()
logger.info(f"同步完成: 处理 {result['synced']} 个密钥, 更新 {result['updated']} 个, 错误 {result['errors']}")
logger.info(
f"同步完成: 处理 {result['synced']} 个密钥, 更新 {result['updated']} 个, 错误 {result['errors']}"
)
except Exception as e:
logger.error(f"同步统计数据时出错: {e}")

View File

@@ -0,0 +1,25 @@
"""
异步任务服务层
提供视频/图片/音频等异步任务的:
- 提交阶段故障转移AsyncTaskOrchestrator
- 终态计费与 Usage 写入VideoTelemetry 等)
"""
from .orchestrator import (
AllCandidatesFailedError,
AsyncTaskOrchestrator,
CandidateSubmissionError,
CandidateUnsupportedError,
SubmitOutcome,
UpstreamClientRequestError,
)
__all__ = [
"AsyncTaskOrchestrator",
"SubmitOutcome",
"AllCandidatesFailedError",
"UpstreamClientRequestError",
"CandidateUnsupportedError",
"CandidateSubmissionError",
]

View File

@@ -0,0 +1,3 @@
"""Task telemetry implementations for concrete task types (video/image/audio)."""
__all__ = []

View File

@@ -0,0 +1,320 @@
"""
VideoTelemetryPhase3
将 Video 异步任务的“终态计费 + Usage 写入 + required 缺失告警”从 poller 中抽离出来,
便于未来 Image/Audio 复用相同框架。
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.video_handler_base import sanitize_error_message
from src.config.settings import config
from src.core.api_format.conversion.internal_video import VideoStatus
from src.core.logger import logger
from src.models.database import ApiKey, Provider, User, VideoTask
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.usage.service import UsageService
class VideoTelemetry:
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
self._formula_engine = FormulaEngine()
async def record_terminal_usage(self, task: VideoTask) -> None:
"""
为视频任务终态写入 Usage
- COMPLETED: 使用 FormulaEngine 计算 cost或 no_rule / incomplete -> cost=0
- FAILED: cost=0
该方法可能会在 strict_mode 缺失 required 维度时将任务降级为 FAILED 并隐藏产物。
"""
request_id = None
if isinstance(task.request_metadata, dict):
request_id = task.request_metadata.get("request_id")
request_id = request_id or task.id
# 计算异步任务总耗时ms
response_time_ms = None
if task.submitted_at and task.completed_at:
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
# 基础维度(无需 collectors 也可计费)
base_dimensions: dict[str, Any] = {
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size or "",
"retry_count": task.retry_count,
}
# collectors 可用的 metadata结构稳定便于配置 path
collector_metadata: dict[str, Any] = {
"task": {
"id": task.id,
"external_task_id": task.external_task_id,
"model": task.model,
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
"retry_count": task.retry_count,
"video_size_bytes": task.video_size_bytes,
},
"result": {
"video_url": task.video_url,
"video_urls": task.video_urls or [],
},
}
# 维度采集base + collectors 覆盖/补全
dims = DimensionCollectorService(self.db).collect_dimensions(
api_format=task.provider_api_format,
task_type="video",
request=task.original_request_body or {},
response=(
(task.request_metadata or {}).get("poll_raw_response")
if isinstance(task.request_metadata, dict)
else None
),
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# 取冻结的 rule_snapshot若缺失则回退 DB 查找(兼容旧任务)
rule_snapshot = None
if isinstance(task.request_metadata, dict):
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
billing_snapshot: dict[str, Any] = {
"status": "complete",
"missing_required": [],
"strict_mode": config.billing_strict_mode,
}
cost = 0.0
if task.status == VideoStatus.FAILED.value:
billing_snapshot["billed_reason"] = "task_failed"
else:
# COMPLETED计算成本
expression = None
variables = None
dimension_mappings = None
rule_id = None
rule_name = None
rule_scope = None
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
rule_id = rule_snapshot.get("rule_id")
rule_name = rule_snapshot.get("rule_name")
rule_scope = rule_snapshot.get("scope")
expression = rule_snapshot.get("expression")
variables = rule_snapshot.get("variables")
dimension_mappings = rule_snapshot.get("dimension_mappings")
else:
lookup = BillingRuleService.find_rule(
self.db,
provider_id=task.provider_id,
model_name=task.model,
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = rule.id
rule_name = rule.name
rule_scope = lookup.scope
expression = rule.expression
variables = rule.variables
dimension_mappings = rule.dimension_mappings
if not expression:
billing_snapshot["status"] = "no_rule"
billing_snapshot["cost_breakdown"] = {"total": 0.0}
logger.warning(
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
request_id,
task.model,
task.provider_id,
)
else:
billing_snapshot.update(
{
"rule_id": rule_id,
"rule_name": rule_name,
"rule_scope": rule_scope,
"expression": expression,
"variables": variables or {},
}
)
try:
result = self._formula_engine.evaluate(
expression=expression,
variables=variables or {},
dimensions=dims,
dimension_mappings=dimension_mappings or {},
strict_mode=config.billing_strict_mode,
)
billing_snapshot["status"] = result.status
billing_snapshot["missing_required"] = result.missing_required
billing_snapshot["resolved_values"] = result.resolved_values
if result.status == "complete":
cost = result.cost
else:
logger.error(
"Billing incomplete due to missing required dimensions "
"(request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
result.missing_required,
)
cost = 0.0
await self._maybe_alert_missing_required(
model=task.model,
missing_required=result.missing_required,
)
if result.error:
billing_snapshot["error"] = result.error
except BillingIncompleteError as exc:
logger.error(
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
exc.missing_required,
)
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["resolved_values"] = {}
billing_snapshot["error"] = "strict_mode_missing_required"
cost = 0.0
# strict_mode=true标记任务失败并隐藏产物避免"免费放行"
task.status = VideoStatus.FAILED.value
task.error_code = "billing_incomplete"
task.error_message = f"Missing required dimensions: {exc.missing_required}"
task.video_url = None
task.video_urls = None
await self._maybe_alert_missing_required(
model=task.model,
missing_required=exc.missing_required,
)
billing_snapshot["cost_breakdown"] = {"total": cost}
# 将 billing_snapshot 回写到 task.request_metadata 便于对账(不会影响 usage 的单独存档)
if task.request_metadata is None:
task.request_metadata = {}
if isinstance(task.request_metadata, dict):
task.request_metadata["billing_snapshot"] = billing_snapshot
usage_metadata: dict[str, Any] = {
"billing_snapshot": billing_snapshot,
"dimensions": dims,
"raw_response_ref": {
"video_task_id": task.id,
"field": "video_tasks.request_metadata.poll_raw_response",
},
}
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
api_key_obj = (
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
if task.api_key_id
else None
)
provider_obj = (
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
if task.provider_id
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
await UsageService.record_usage_with_custom_cost(
db=self.db,
user=user_obj,
api_key=api_key_obj,
provider=provider_name,
model=task.model,
request_type="video",
total_cost_usd=cost,
request_cost_usd=cost,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=task.client_api_format,
endpoint_api_format=task.provider_api_format,
has_format_conversion=bool(task.format_converted),
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
error_message=(
None
if task.status == VideoStatus.COMPLETED.value
else (task.error_message or task.error_code or "video_task_failed")
),
metadata=usage_metadata,
request_headers=(
(task.request_metadata or {}).get("request_headers")
if isinstance(task.request_metadata, dict)
else None
),
request_body=task.original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=task.provider_id,
provider_endpoint_id=task.endpoint_id,
provider_api_key_id=task.key_id,
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
target_model=None,
)
async def _maybe_alert_missing_required(
self, *, model: str, missing_required: list[str]
) -> None:
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
if not missing_required:
return
if not self.redis:
logger.error(
"Missing required billing dimensions (model=%s): %s", model, missing_required
)
return
# 按小时 bucket 聚合
now = datetime.now(timezone.utc)
hour_bucket = now.strftime("%Y%m%d%H")
for dim in missing_required:
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
try:
count = await self.redis.incr(key)
if count == 1:
await self.redis.expire(key, 3700)
if count >= 10:
logger.warning(
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
model,
dim,
count,
)
except Exception as exc:
logger.warning(
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
)
__all__ = ["VideoTelemetry"]

View File

@@ -0,0 +1,635 @@
"""
AsyncTaskOrchestrator
提交阶段故障转移(多候选尝试):
- 目标:拿到 external_task_id 后锁定 provider/endpoint/key后续轮询不再切换。
- 仅覆盖“提交阶段”;轮询阶段由各 task poller 使用已锁定的信息执行。
"""
from __future__ import annotations
import re
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Protocol, runtime_checkable
import httpx
from redis.asyncio import Redis
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
from src.services.orchestration.candidate_resolver import CandidateResolver
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.system.config import SystemConfigService
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
re.IGNORECASE,
)
def _sanitize(message: str, max_length: int = 200) -> str:
if not message:
return "request_failed"
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
@runtime_checkable
class SubmitFunc(Protocol):
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
@runtime_checkable
class ExtractExternalTaskIdFunc(Protocol):
def __call__(self, payload: dict[str, Any]) -> str | None: ...
class UpstreamClientRequestError(RuntimeError):
"""可判定为客户端请求问题(不应 failover的上游错误。"""
def __init__(
self,
*,
response: httpx.Response,
candidate_keys: list[dict[str, Any]],
) -> None:
self.response = response
self.candidate_keys = candidate_keys
super().__init__(f"Upstream client error: HTTP {response.status_code}")
class AllCandidatesFailedError(RuntimeError):
def __init__(
self,
*,
reason: str,
candidate_keys: list[dict[str, Any]],
last_status_code: int | None = None,
) -> None:
self.reason = reason
self.candidate_keys = candidate_keys
self.last_status_code = last_status_code
super().__init__(f"All candidates failed: {reason}")
class CandidateUnsupportedError(RuntimeError):
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
class CandidateSubmissionError(RuntimeError):
"""候选提交异常(网络/解密/解析等)。"""
@dataclass(slots=True)
class SubmitOutcome:
candidate: ProviderCandidate
candidate_keys: list[dict[str, Any]]
external_task_id: str
rule_lookup: BillingRuleLookupResult | None
upstream_payload: dict[str, Any] | None = None
class AsyncTaskOrchestrator:
"""
异步任务编排器:只负责提交阶段的候选遍历与错误处理策略。
"""
def __init__(self, db: Session, *, redis_client: Redis | None = None) -> None:
self.db = db
self.redis = redis_client
self._candidate_resolver: CandidateResolver | None = None
self._error_classifier: ErrorClassifier | None = None
self._cache_scheduler = None
# 候选记录映射:{candidate_index: RequestCandidate}
self._candidate_records: dict[int, RequestCandidate] = {}
def _create_candidate_records(
self,
candidates: list[ProviderCandidate],
request_id: str | None,
user_api_key: ApiKey,
) -> dict[int, RequestCandidate]:
"""
为所有候选预创建 RequestCandidate 记录。
Args:
candidates: 候选列表
request_id: 请求 ID
user_api_key: 用户 API Key
Returns:
{candidate_index: RequestCandidate} 映射
"""
if not request_id:
return {}
now = datetime.now(timezone.utc)
records: dict[int, RequestCandidate] = {}
for idx, cand in enumerate(candidates):
record = RequestCandidate(
id=str(uuid.uuid4()),
request_id=request_id,
candidate_index=idx,
retry_index=0,
user_id=user_api_key.user_id if user_api_key else None,
api_key_id=user_api_key.id if user_api_key else None,
provider_id=cand.provider.id,
endpoint_id=cand.endpoint.id,
key_id=cand.key.id,
status="available",
is_cached=bool(getattr(cand, "is_cached", False)),
created_at=now,
)
self.db.add(record)
records[idx] = record
try:
self.db.flush()
except Exception as exc:
logger.warning(
"[AsyncTaskOrchestrator] Failed to create candidate records: %s",
str(exc),
)
self.db.rollback()
return {}
return records
def _update_candidate_record(
self,
idx: int,
*,
status: str,
skip_reason: str | None = None,
status_code: int | None = None,
error_type: str | None = None,
error_message: str | None = None,
started_at: datetime | None = None,
finished_at: datetime | None = None,
) -> None:
"""更新候选记录状态。"""
record = self._candidate_records.get(idx)
if not record:
return
record.status = status
if skip_reason is not None:
record.skip_reason = skip_reason
if status_code is not None:
record.status_code = status_code
if error_type is not None:
record.error_type = error_type
if error_message is not None:
record.error_message = error_message
if started_at is not None:
record.started_at = started_at
if finished_at is not None:
record.finished_at = finished_at
try:
self.db.flush()
except Exception as exc:
logger.warning(
"[AsyncTaskOrchestrator] Failed to update candidate record %d: %s",
idx,
str(exc),
)
def _commit_candidate_records(self) -> None:
"""提交候选记录到数据库。"""
if not self._candidate_records:
return
try:
self.db.commit()
except Exception as exc:
logger.warning(
"[AsyncTaskOrchestrator] Failed to commit candidate records: %s",
str(exc),
)
self.db.rollback()
async def _ensure_initialized(self) -> None:
if self._cache_scheduler is not None:
return
# 使用 SystemConfigService 读取运行时调度策略(与 Chat/CLI 一致)
priority_mode = SystemConfigService.get_config(
self.db,
"provider_priority_mode",
"provider",
)
scheduling_mode = SystemConfigService.get_config(
self.db,
"scheduling_mode",
"cache_affinity",
)
self._cache_scheduler = await get_cache_aware_scheduler(
self.redis,
priority_mode=priority_mode,
scheduling_mode=scheduling_mode,
)
self._candidate_resolver = CandidateResolver(
db=self.db,
cache_scheduler=self._cache_scheduler,
)
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
"""
判断某个上游 HTTP 错误是否为“客户端错误”(不应 failover
规则:
- 401/403/429一般是 key/权限/限流问题,优先 failover
- 其他 4xx若 ErrorClassifier 判断为客户端请求错误,则停止
"""
if status_code in (401, 403, 429):
return False
if 400 <= status_code < 500:
assert self._error_classifier is not None
return self._error_classifier.is_client_error(error_text)
return False
async def submit_with_failover(
self,
*,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey,
request_id: str | None,
task_type: str,
submit_func: SubmitFunc,
extract_external_task_id: ExtractExternalTaskIdFunc,
supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False,
capability_requirements: dict[str, bool] | None = None,
max_candidates: int | None = None,
) -> SubmitOutcome:
"""
提交异步任务并在失败时自动尝试下一个候选,直到拿到 external_task_id。
Returns:
SubmitOutcome包含选中的候选 + external_task_id + candidate_keys + billing rule lookup
Raises:
UpstreamClientRequestError: 判定为客户端请求错误(不应 failover
ProviderNotAvailableException: 没有可用候选(调度器层面)
AllCandidatesFailedError: 有候选但全部提交失败
"""
await self._ensure_initialized()
assert self._candidate_resolver is not None
logger.info(
"[AsyncTaskOrchestrator] submit_with_failover: "
"api_format=%s, model=%s, task_type=%s, request_id=%s",
api_format,
model_name,
task_type,
request_id,
)
candidates, _global_model_id = await self._candidate_resolver.fetch_candidates(
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
is_stream=False,
capability_requirements=capability_requirements,
)
logger.info(
"[AsyncTaskOrchestrator] fetch_candidates returned %d candidates for model=%s",
len(candidates),
model_name,
)
# 如果没有候选,直接抛出异常
if not candidates:
logger.error(
"[AsyncTaskOrchestrator] No candidates returned from fetch_candidates for model=%s",
model_name,
)
raise ProviderNotAvailableException("No candidates available")
if max_candidates is not None and max_candidates > 0:
candidates = candidates[:max_candidates]
# 创建候选记录(用于链路追踪)
self._candidate_records = self._create_candidate_records(
candidates=candidates,
request_id=request_id,
user_api_key=user_api_key,
)
candidate_keys: list[dict[str, Any]] = []
eligible_count = 0
last_status_code: int | None = None
for idx, cand in enumerate(candidates):
submit_started_at = datetime.now(timezone.utc)
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
candidate_info: dict[str, Any] = {
"index": idx,
"provider_id": cand.provider.id,
"provider_name": cand.provider.name,
"endpoint_id": cand.endpoint.id,
"key_id": cand.key.id,
"key_name": cand.key.name,
"auth_type": auth_type,
"priority": getattr(cand.key, "priority", 0) or 0,
"is_cached": bool(getattr(cand, "is_cached", False)),
}
candidate_keys.append(candidate_info)
logger.info(
"[AsyncTaskOrchestrator] Checking candidate %d: provider=%s, is_skipped=%s, skip_reason=%s, needs_conversion=%s, auth_type=%s",
idx,
cand.provider.name,
getattr(cand, "is_skipped", False),
getattr(cand, "skip_reason", None),
getattr(cand, "needs_conversion", False),
auth_type,
)
# 调度器层面标记为跳过(健康/熔断/并发等)
if getattr(cand, "is_skipped", False):
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
candidate_info.update(
{
"skipped": True,
"skip_reason": skip_reason,
}
)
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
logger.info(
"[AsyncTaskOrchestrator] Candidate %d skipped: is_skipped=True, reason=%s",
idx,
cand.skip_reason,
)
continue
# 视频/图片等直连 upstream 的 handler 目前不支持跨格式转换
if not allow_format_conversion and bool(getattr(cand, "needs_conversion", False)):
candidate_info.update(
{"skipped": True, "skip_reason": "format_conversion_not_supported"}
)
self._update_candidate_record(
idx, status="skipped", skip_reason="format_conversion_not_supported"
)
logger.info("[AsyncTaskOrchestrator] Candidate %d skipped: needs_conversion", idx)
continue
# auth_type 过滤
if supported_auth_types is not None and auth_type not in supported_auth_types:
skip_reason = f"unsupported_auth_type:{auth_type}"
candidate_info.update(
{
"skipped": True,
"skip_reason": skip_reason,
}
)
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
logger.info(
"[AsyncTaskOrchestrator] Candidate %d skipped: unsupported_auth_type=%s",
idx,
auth_type,
)
continue
# billing rule 过滤(可选)
rule_lookup: BillingRuleLookupResult | None = None
has_billing_rule = True
if config.billing_require_rule:
logger.info(
"[AsyncTaskOrchestrator] Checking billing rule for candidate %d (billing_require_rule=True)",
idx,
)
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=cand.provider.id,
model_name=model_name,
task_type=task_type,
)
has_billing_rule = rule_lookup is not None
logger.info(
"[AsyncTaskOrchestrator] Billing rule lookup result: has_rule=%s",
has_billing_rule,
)
if not has_billing_rule:
candidate_info.update(
{
"has_billing_rule": False,
"skipped": True,
"skip_reason": "billing_rule_missing",
}
)
self._update_candidate_record(
idx, status="skipped", skip_reason="billing_rule_missing"
)
logger.info(
"[AsyncTaskOrchestrator] Candidate %d skipped: billing_rule_missing", idx
)
continue
candidate_info["has_billing_rule"] = has_billing_rule
logger.info("[AsyncTaskOrchestrator] Candidate %d eligible, attempting submit", idx)
eligible_count += 1
# 更新记录为 pending 状态(开始尝试)
self._update_candidate_record(idx, status="pending", started_at=submit_started_at)
# 尝试提交
try:
response = await submit_func(cand)
except Exception as exc:
finished_at = datetime.now(timezone.utc)
logger.error(
"[AsyncTaskOrchestrator] Candidate %d submit exception: %s: %s",
idx,
type(exc).__name__,
str(exc),
)
error_msg = _sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "exception",
"error_type": type(exc).__name__,
"error_message": error_msg,
}
)
self._update_candidate_record(
idx,
status="failed",
error_type=type(exc).__name__,
error_message=error_msg,
finished_at=finished_at,
)
continue
logger.info(
"[AsyncTaskOrchestrator] Candidate %d submit response: status_code=%d",
idx,
response.status_code,
)
last_status_code = int(getattr(response, "status_code", 0) or 0)
# 上游错误:决定是否停止
if response.status_code >= 400:
finished_at = datetime.now(timezone.utc)
error_text = ""
try:
error_text = response.text or ""
except Exception:
error_text = ""
error_msg = _sanitize(error_text)
candidate_info.update(
{
"attempt_status": "http_error",
"status_code": response.status_code,
"error_message": error_msg,
}
)
self._update_candidate_record(
idx,
status="failed",
status_code=response.status_code,
error_type="http_error",
error_message=error_msg,
finished_at=finished_at,
)
if self._should_stop_on_http_error(
status_code=response.status_code, error_text=error_text
):
self._commit_candidate_records()
raise UpstreamClientRequestError(
response=response,
candidate_keys=candidate_keys,
)
continue
# 解析任务 ID200 但缺字段也视为失败并 failover
payload: dict[str, Any] | None = None
try:
data = response.json()
if isinstance(data, dict):
payload = data
logger.info(
"[AsyncTaskOrchestrator] Candidate %d response payload: %s",
idx,
str(payload)[:500] if payload else "None",
)
except Exception as exc:
finished_at = datetime.now(timezone.utc)
logger.error(
"[AsyncTaskOrchestrator] Candidate %d invalid JSON: %s",
idx,
str(exc),
)
error_msg = _sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "invalid_json",
"error_type": type(exc).__name__,
"error_message": error_msg,
}
)
self._update_candidate_record(
idx,
status="failed",
status_code=response.status_code,
error_type="invalid_json",
error_message=error_msg,
finished_at=finished_at,
)
continue
external_task_id = extract_external_task_id(payload or {})
logger.info(
"[AsyncTaskOrchestrator] Candidate %d extracted task_id: %s",
idx,
external_task_id,
)
if not external_task_id:
finished_at = datetime.now(timezone.utc)
candidate_info.update(
{
"attempt_status": "empty_task_id",
"error_message": "Upstream returned empty task id",
}
)
self._update_candidate_record(
idx,
status="failed",
status_code=response.status_code,
error_type="empty_task_id",
error_message="Upstream returned empty task id",
finished_at=finished_at,
)
logger.warning(
"[AsyncTaskOrchestrator] Candidate %d: empty task_id, payload keys: %s",
idx,
list(payload.keys()) if payload else [],
)
continue
# 成功
finished_at = datetime.now(timezone.utc)
candidate_info.update({"attempt_status": "success", "selected": True})
self._update_candidate_record(
idx,
status="success",
status_code=response.status_code,
finished_at=finished_at,
)
self._commit_candidate_records()
return SubmitOutcome(
candidate=cand,
candidate_keys=candidate_keys,
external_task_id=str(external_task_id),
rule_lookup=rule_lookup,
upstream_payload=payload,
)
# 没有任何候选可尝试
if not candidates:
raise ProviderNotAvailableException("No candidates available")
# 提交所有候选记录
self._commit_candidate_records()
if eligible_count == 0:
reason = "no_eligible_candidates"
if config.billing_require_rule:
reason = "no_candidate_with_billing_rule"
raise AllCandidatesFailedError(
reason=reason,
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
raise AllCandidatesFailedError(
reason="all_candidates_failed",
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
__all__ = [
"AsyncTaskOrchestrator",
"SubmitOutcome",
"AllCandidatesFailedError",
"UpstreamClientRequestError",
"CandidateUnsupportedError",
"CandidateSubmissionError",
]

View File

@@ -4,7 +4,6 @@ Usage Redis Streams consumer.
高性能消费者实现,支持批量处理和单次提交多条记录。
"""
from __future__ import annotations
import asyncio
@@ -104,9 +103,7 @@ async def ensure_usage_stream_group() -> None:
id="0-0",
mkstream=True,
)
logger.info(
f"[usage-queue] Created consumer group {config.usage_queue_stream_group}"
)
logger.info(f"[usage-queue] Created consumer group {config.usage_queue_stream_group}")
except ResponseError as exc:
if "BUSYGROUP" in str(exc):
return
@@ -115,7 +112,7 @@ async def ensure_usage_stream_group() -> None:
class UsageQueueConsumer:
"""Usage 队列消费者
性能优化:
- 缓存配置值避免重复属性访问
- STREAMING 事件使用 pipeline 批量 ACK
@@ -259,14 +256,14 @@ class UsageQueueConsumer:
) -> None:
"""批量处理 STREAMING 事件(状态更新)"""
success_ids: list[str] = []
for message_id, event in messages:
try:
await self._apply_streaming_event(event)
success_ids.append(message_id)
except Exception as exc:
await self._handle_processing_error(redis_client, message_id, {}, exc)
# 使用 pipeline 批量 ACK 成功处理的消息
if success_ids:
pipe = redis_client.pipeline()
@@ -276,8 +273,8 @@ class UsageQueueConsumer:
async def _process_record_batch(
self,
redis_client: Any,
messages: list[tuple[str, dict[str, Any], UsageEvent]],
redis_client: Any,
messages: list[tuple[str, dict[str, Any], UsageEvent]],
) -> None:
"""批量处理记录类型的事件"""
db = create_session()
@@ -306,7 +303,9 @@ class UsageQueueConsumer:
except Exception as exc:
# 批量处理失败,回退到逐条处理(复用已创建的 db session
logger.warning(f"[usage-queue] Batch processing failed, falling back to individual: {exc}")
logger.warning(
f"[usage-queue] Batch processing failed, falling back to individual: {exc}"
)
try:
db.rollback() # 清理批量失败的事务状态
except Exception:
@@ -330,7 +329,9 @@ class UsageQueueConsumer:
else:
await self._handle_processing_error(redis_client, message_id, fields, ie)
except Exception as individual_exc:
await self._handle_processing_error(redis_client, message_id, fields, individual_exc)
await self._handle_processing_error(
redis_client, message_id, fields, individual_exc
)
# 批量 ACK 成功处理的消息
if success_ids:
pipe = redis_client.pipeline()
@@ -342,10 +343,10 @@ class UsageQueueConsumer:
async def _handle_processing_error(
self,
redis_client: Any,
message_id: str,
fields: dict[str, Any],
error: Exception,
redis_client: Any,
message_id: str,
fields: dict[str, Any],
error: Exception,
) -> None:
retries = await self._get_delivery_count(redis_client, message_id)
if retries >= self._max_retries:
@@ -415,9 +416,7 @@ class UsageQueueConsumer:
finally:
db.close()
async def _apply_record_event(
self, event: UsageEvent, db: Session | None = None
) -> None:
async def _apply_record_event(self, event: UsageEvent, db: Session | None = None) -> None:
"""处理记录类型事件(逐条写入,用于 fallback
Args:
@@ -501,9 +500,7 @@ class UsageQueueConsumer:
pending_count = int(pending.get("pending", 0))
elif isinstance(pending, (list, tuple)) and pending:
pending_count = int(pending[0])
logger.info(
f"[usage-queue] backlog={stream_len} pending={pending_count}"
)
logger.info(f"[usage-queue] backlog={stream_len} pending={pending_count}")
except Exception as exc:
logger.debug(f"[usage-queue] metrics log failed: {exc}")

View File

@@ -2,8 +2,8 @@
Usage 事件定义与序列化工具(用于 Redis Streams
"""
from __future__ import annotations
import json
import time
from dataclasses import dataclass

View File

@@ -10,8 +10,8 @@
from __future__ import annotations
from typing import Any
from datetime import datetime, timezone
from typing import Any
from src.core.enums import ProviderBillingType
from src.core.logger import logger

View File

@@ -36,7 +36,6 @@ from src.services.system.audit import audit_service
from src.services.usage.service import UsageService
class UsageRecorder:
"""
统一的 Usage 记录器
@@ -144,9 +143,11 @@ class UsageRecorder:
status_code=200,
)
logger.debug(f"[UsageRecorder] 成功记录: provider={metadata.provider}, "
logger.debug(
f"[UsageRecorder] 成功记录: provider={metadata.provider}, "
f"model={metadata.model}, api_format={metadata.api_format}, "
f"tokens={usage.input_tokens}+{usage.output_tokens}")
f"tokens={usage.input_tokens}+{usage.output_tokens}"
)
async def record_failure(
self,
@@ -213,9 +214,11 @@ class UsageRecorder:
error_message=result.error_message,
)
logger.debug(f"[UsageRecorder] 失败记录: provider={metadata.provider}, "
logger.debug(
f"[UsageRecorder] 失败记录: provider={metadata.provider}, "
f"model={metadata.model}, api_format={metadata.api_format}, "
f"status={result.status_code}, error={result.error_message[:100] if result.error_message else 'N/A'}")
f"status={result.status_code}, error={result.error_message[:100] if result.error_message else 'N/A'}"
)
async def record_from_exception(
self,

View File

@@ -12,7 +12,8 @@ from typing import Any
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.core.api_format.metadata import can_passthrough
from src.core.api_format.metadata import can_passthrough_endpoint
from src.core.api_format.signature import normalize_signature_key
from src.core.logger import logger
from src.models.database import (
ApiKey,
@@ -2607,7 +2608,12 @@ class UsageService:
# 兼容历史数据:当 streaming 状态已拿到两个格式但 has_format_conversion 为空时,回填推断结果
if has_format_conversion is None and api_format and endpoint_api_format:
has_format_conversion = not can_passthrough(api_format, endpoint_api_format)
client_raw = str(api_format).strip()
endpoint_raw = str(endpoint_api_format).strip()
if ":" in client_raw and ":" in endpoint_raw:
client_fmt = normalize_signature_key(client_raw)
endpoint_fmt = normalize_signature_key(endpoint_raw)
has_format_conversion = not can_passthrough_endpoint(client_fmt, endpoint_fmt)
item: dict[str, Any] = {
"id": r.id,

View File

@@ -7,8 +7,8 @@ from __future__ import annotations
import json
import re
from typing import Any
from collections.abc import AsyncIterator
from typing import Any
from sqlalchemy.orm import Session
@@ -21,7 +21,6 @@ from src.models.database import ApiKey, User
from src.services.usage.service import UsageService
class StreamUsageTracker:
"""流式响应用量跟踪器"""
@@ -64,7 +63,7 @@ class StreamUsageTracker:
provider_id: Provider ID用于记录真实成本
provider_endpoint_id: Endpoint ID用于记录真实成本
provider_api_key_id: API Key ID用于记录真实成本
api_format: API 格式CLAUDE, CLAUDE_CLI, OPENAI, OPENAI_CLI
api_format: endpoint signature"claude:chat", "openai:cli"
endpoint_api_format: 端点原生 API 格式
has_format_conversion: 是否发生了格式转换
"""
@@ -85,7 +84,7 @@ class StreamUsageTracker:
self.provider_api_key_id = provider_api_key_id
# API 格式和响应解析器
self.api_format = api_format or "CLAUDE"
self.api_format = api_format or "claude:chat"
self.endpoint_api_format = endpoint_api_format
self.has_format_conversion = has_format_conversion
self.response_parser = get_parser_for_format(self.api_format)
@@ -144,7 +143,9 @@ class StreamUsageTracker:
"""
self.status_code = status_code
self.error_message = error_message
logger.debug(f"ID:{self.request_id} | 流式响应错误状态已设置 | 状态码:{status_code} | 错误:{error_message[:100]}")
logger.debug(
f"ID:{self.request_id} | 流式响应错误状态已设置 | 状态码:{status_code} | 错误:{error_message[:100]}"
)
def _update_complete_response(self, chunk: dict[str, Any]) -> None:
"""根据响应块更新完整响应结构"""
@@ -465,7 +466,9 @@ class StreamUsageTracker:
messages = request_data.get("messages", [])
self.input_tokens = self.estimate_input_tokens(messages)
logger.debug(f"ID:{self.request_id} | 开始跟踪流式响应 | 估算输入tokens:{self.input_tokens}")
logger.debug(
f"ID:{self.request_id} | 开始跟踪流式响应 | 估算输入tokens:{self.input_tokens}"
)
chunk_count = 0
first_byte_time_ms = None # 预先记录 TTFB避免 yield 后计算不准确
@@ -479,7 +482,9 @@ class StreamUsageTracker:
if chunk_count == 1:
# 计算 TTFB使用请求原始开始时间或 track_stream 开始时间)
base_time = self.request_start_time or self.start_time
first_byte_time_ms = int((time.time() - base_time) * 1000) if base_time else None
first_byte_time_ms = (
int((time.time() - base_time) * 1000) if base_time else None
)
# 先返回原始块给客户端,确保 TTFB 不受数据库操作影响
yield chunk
@@ -535,8 +540,10 @@ class StreamUsageTracker:
# 流结束后记录使用量
self.end_time = time.time()
logger.debug(f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}")
logger.debug(
f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}"
)
# 检查是否收到了有效数据
# 情况1: 收到了原始数据但无法解析为有效的SSE JSON
@@ -564,7 +571,9 @@ class StreamUsageTracker:
await self._record_usage()
except Exception as e:
# 如果记录失败,至少输出基本的汇总日志
logger.exception(f"Failed to record stream usage for request {self.request_id}: {e}")
logger.exception(
f"Failed to record stream usage for request {self.request_id}: {e}"
)
# 尝试输出基本的汇总日志,使用多层防护
try:
# 计算响应时间,使用多层后备机制
@@ -581,13 +590,17 @@ class StreamUsageTracker:
total_response_time = 0
# 安全地输出汇总日志
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)")
logger.info(
f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)"
)
except Exception as log_error:
# 最后的防线:输出最简单的完成标记
logger.error(f"Failed to output summary log: {log_error}")
try:
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成")
logger.info(
f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成"
)
except Exception:
# 如果连最简单的日志都失败了,放弃
pass
@@ -686,7 +699,9 @@ class StreamUsageTracker:
user, api_key = _load_user_and_key(db_for_usage)
except InvalidRequestError:
# 会话处于不可用状态,需要回滚并重新开始
logger.warning(f"Session in invalid state for request {self.request_id}, rolling back and retrying")
logger.warning(
f"Session in invalid state for request {self.request_id}, rolling back and retrying"
)
try:
db_for_usage.rollback()
except Exception:
@@ -703,7 +718,9 @@ class StreamUsageTracker:
created_temp_session = True
user, api_key = _load_user_and_key(db_for_usage)
except Exception as session_error:
logger.exception(f"Failed to recover from invalid session for request {self.request_id}: {session_error}")
logger.exception(
f"Failed to recover from invalid session for request {self.request_id}: {session_error}"
)
return
# 根据状态码确定请求状态
@@ -774,8 +791,10 @@ class StreamUsageTracker:
else:
cost_str = "$0"
logger.info(f"{status_prefix} ID:{self.request_id} | {self.status_code} | 耗时:{total_response_time}ms | "
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:{cost_str}")
logger.info(
f"{status_prefix} ID:{self.request_id} | {self.status_code} | 耗时:{total_response_time}ms | "
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:{cost_str}"
)
# 记录提供商结果用于动态权重调整
# 记录提供商结果的健康监控已由 FallbackOrchestrator 自动处理
@@ -936,7 +955,9 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
messages = request_data.get("messages", [])
self.input_tokens = self.estimate_input_tokens(messages)
logger.debug(f"ID:{self.request_id} | 开始跟踪流式响应(Enhanced) | 估算输入tokens:{self.input_tokens}")
logger.debug(
f"ID:{self.request_id} | 开始跟踪流式响应(Enhanced) | 估算输入tokens:{self.input_tokens}"
)
chunk_count = 0
first_byte_time_ms = None # 预先记录 TTFB避免 yield 后计算不准确
@@ -950,7 +971,9 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
if chunk_count == 1:
# 计算 TTFB使用请求原始开始时间或 track_stream 开始时间)
base_time = self.request_start_time or self.start_time
first_byte_time_ms = int((time.time() - base_time) * 1000) if base_time else None
first_byte_time_ms = (
int((time.time() - base_time) * 1000) if base_time else None
)
# 先返回原始块给客户端,确保 TTFB 不受数据库操作影响
yield chunk
@@ -997,8 +1020,10 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
# 流结束后记录使用量
self.end_time = time.time()
logger.debug(f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}")
logger.debug(
f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}"
)
# 检查是否收到了有效数据
# 情况1: 收到了原始数据但无法解析为有效的SSE JSON
@@ -1026,7 +1051,9 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
await self._record_usage()
except Exception as e:
# 如果记录失败,至少输出基本的汇总日志
logger.exception(f"Failed to record stream usage for request {self.request_id}: {e}")
logger.exception(
f"Failed to record stream usage for request {self.request_id}: {e}"
)
# 尝试输出基本的汇总日志,使用多层防护
try:
# 计算响应时间,使用多层后备机制
@@ -1043,13 +1070,17 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
total_response_time = 0
# 安全地输出汇总日志
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)")
logger.info(
f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)"
)
except Exception as log_error:
# 最后的防线:输出最简单的完成标记
logger.error(f"Failed to output summary log: {log_error}")
try:
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成")
logger.info(
f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成"
)
except Exception:
# 如果连最简单的日志都失败了,放弃
pass
@@ -1096,7 +1127,7 @@ def create_stream_tracker(
provider_id: Provider ID用于记录真实成本
provider_endpoint_id: Endpoint ID用于记录真实成本
provider_api_key_id: API Key ID用于记录真实成本
api_format: API 格式CLAUDE, CLAUDE_CLI, OPENAI, OPENAI_CLI
api_format: endpoint signature"claude:chat", "openai:cli"
endpoint_api_format: 端点原生 API 格式
has_format_conversion: 是否发生了格式转换

View File

@@ -2,7 +2,6 @@
Telemetry writer abstraction for stream usage.
"""
from __future__ import annotations
import json
@@ -124,7 +123,7 @@ class QueueTelemetryWriter(TelemetryWriter):
"user_id": self.user_id,
"api_key_id": self.api_key_id,
}
# 可选字段 - 只添加非 None/非默认值,减少 payload 大小
# 注意:消费者端需要处理缺失字段的默认值
if kwargs.get("provider"):
@@ -133,7 +132,7 @@ class QueueTelemetryWriter(TelemetryWriter):
data["model"] = kwargs["model"]
if kwargs.get("target_model"):
data["target_model"] = kwargs["target_model"]
# Token 计数 - 0 是常见值,但仍需传递
input_tokens = kwargs.get("input_tokens", 0)
output_tokens = kwargs.get("output_tokens", 0)
@@ -141,7 +140,7 @@ class QueueTelemetryWriter(TelemetryWriter):
data["input_tokens"] = input_tokens
if output_tokens:
data["output_tokens"] = output_tokens
# 缓存 tokencache_creation_tokens -> cache_creation_input_tokens 映射)
cache_creation = kwargs.get("cache_creation_tokens", 0)
cache_read = kwargs.get("cache_read_tokens", 0)
@@ -149,20 +148,20 @@ class QueueTelemetryWriter(TelemetryWriter):
data["cache_creation_input_tokens"] = cache_creation
if cache_read:
data["cache_read_input_tokens"] = cache_read
# 时间指标
if kwargs.get("response_time_ms") is not None:
data["response_time_ms"] = kwargs["response_time_ms"]
if kwargs.get("first_byte_time_ms") is not None:
data["first_byte_time_ms"] = kwargs["first_byte_time_ms"]
# 状态信息
status_code = kwargs.get("status_code", 200)
if status_code != 200:
data["status_code"] = status_code
if kwargs.get("error_message"):
data["error_message"] = kwargs["error_message"]
# 格式信息
request_type = kwargs.get("request_type", "chat")
if request_type != "chat":
@@ -173,11 +172,11 @@ class QueueTelemetryWriter(TelemetryWriter):
data["endpoint_api_format"] = kwargs["endpoint_api_format"]
if kwargs.get("has_format_conversion"):
data["has_format_conversion"] = True
# 流式标记 - 默认 True只记录 False
if not kwargs.get("is_stream", True):
data["is_stream"] = False
# Provider 追踪
if kwargs.get("provider_id"):
data["provider_id"] = kwargs["provider_id"]
@@ -185,7 +184,7 @@ class QueueTelemetryWriter(TelemetryWriter):
data["provider_endpoint_id"] = kwargs["provider_endpoint_id"]
if kwargs.get("provider_api_key_id"):
data["provider_api_key_id"] = kwargs["provider_api_key_id"]
# 元数据
if kwargs.get("metadata"):
data["metadata"] = kwargs["metadata"]

View File

@@ -15,7 +15,6 @@ from src.core.logger import logger
from src.models.database import ApiKey, Usage
class ApiKeyService:
"""API密钥管理服务"""
@@ -86,8 +85,10 @@ class ApiKeyService:
db.commit()
db.refresh(api_key)
logger.info(f"创建API密钥: 用户ID {user_id}, 密钥名 {api_key.name}, "
f"独立Key={is_standalone}, 初始余额={initial_balance_usd}")
logger.info(
f"创建API密钥: 用户ID {user_id}, 密钥名 {api_key.name}, "
f"独立Key={is_standalone}, 初始余额={initial_balance_usd}"
)
return api_key, key # 返回密钥对象和明文密钥
@staticmethod
@@ -256,7 +257,9 @@ class ApiKeyService:
is_allowed = request_count < api_key.rate_limit
if not is_allowed:
logger.warning(f"API密钥速率限制: Key ID {api_key.id}, 请求数 {request_count}/{api_key.rate_limit}")
logger.warning(
f"API密钥速率限制: Key ID {api_key.id}, 请求数 {request_count}/{api_key.rate_limit}"
)
return is_allowed, api_key.rate_limit - request_count
@@ -289,7 +292,9 @@ class ApiKeyService:
if amount_usd < 0:
current = api_key.current_balance_usd or 0
if abs(amount_usd) > current:
logger.warning(f"余额扣除失败: 扣除金额 ${abs(amount_usd):.4f} 超过当前余额 ${current:.4f}")
logger.warning(
f"余额扣除失败: 扣除金额 ${abs(amount_usd):.4f} 超过当前余额 ${current:.4f}"
)
return None
# 调整当前余额
@@ -303,8 +308,10 @@ class ApiKeyService:
db.refresh(api_key)
action = "增加" if amount_usd > 0 else "扣除"
logger.info(f"余额调整成功: Key ID {key_id}, {action} ${abs(amount_usd):.4f}, "
f"余额 ${api_key.current_balance_usd:.4f}")
logger.info(
f"余额调整成功: Key ID {key_id}, {action} ${abs(amount_usd):.4f}, "
f"新余额 ${api_key.current_balance_usd:.4f}"
)
return api_key
@staticmethod
@@ -338,14 +345,18 @@ class ApiKeyService:
if should_delete:
# 物理删除Usage记录会保留因为是 SET NULL
db.delete(api_key)
logger.info(f"删除过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
f"过期时间 {api_key.expires_at}")
logger.info(
f"删除过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
f"过期时间 {api_key.expires_at}"
)
else:
# 仅禁用
api_key.is_active = False
api_key.updated_at = now
logger.info(f"禁用过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
f"过期时间 {api_key.expires_at}")
logger.info(
f"禁用过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
f"过期时间 {api_key.expires_at}"
)
count += 1
if count > 0:

View File

@@ -2,7 +2,6 @@
用户偏好设置服务
"""
from __future__ import annotations
from sqlalchemy.orm import Session
@@ -12,7 +11,6 @@ from src.core.logger import logger
from src.models.database import Provider, User, UserPreference
class PreferenceService:
"""用户偏好设置服务"""

View File

@@ -18,7 +18,6 @@ from src.services.cache.user_cache import UserCacheService
from src.utils.transaction_manager import retry_on_database_error, transactional
class UserService:
"""用户管理服务"""
@@ -198,7 +197,12 @@ class UserService:
]
# 允许设置为 None 的字段(表示无限制)
nullable_fields = ["quota_usd", "allowed_providers", "allowed_api_formats", "allowed_models"]
nullable_fields = [
"quota_usd",
"allowed_providers",
"allowed_api_formats",
"allowed_models",
]
for field, value in kwargs.items():
if field not in updatable_fields:
@@ -442,7 +446,9 @@ class UserService:
# 应用访问限制过滤
filtered_models = []
for model in all_models:
model_name = model.global_model.name if model.global_model else model.provider_model_name
model_name = (
model.global_model.name if model.global_model else model.provider_model_name
)
# 使用 AccessRestrictions.is_model_allowed 检查模型是否可访问
if restrictions.is_model_allowed(model_name, model.provider_id):
filtered_models.append(model)