perf: 全栈查询优化、前端缓存去重与页面可见性优化

后端:
- SQL count 查询统一改用 func.count() 子查询替代 query.count()
- Dashboard/Audit 等页面多次独立查询合并为单次聚合查询
- Provider summary 列表改为批量查询消除 N+1 问题
- DailyStats 逐天循环查询改为 CASE 分桶单次查询
- 使用 load_only() 减少不必要的列加载
- cache_decorator 支持嵌套属性路径解析(dotted vary_by)
- 多个管理/公共端点新增 @cache_result 缓存装饰

前端:
- cache.ts 新增 in-flight 请求复用、dedupedRequest、buildCacheKey
- 大量 API 调用添加前端缓存或去重
- 多个页面定时器在标签页隐藏时暂停、可见时恢复
- Auth 检查从 setInterval 改为 storage + visibilitychange 事件驱动
- 请求竞态防护(requestId 模式)

数据库:
- Usage 表新增 idx_usage_status_user_created 复合索引
This commit is contained in:
fawney19
2026-03-03 22:04:40 +08:00
parent 0a60492146
commit 97b0146ce9
66 changed files with 2306 additions and 657 deletions

View File

@@ -11,6 +11,7 @@ from typing import Any
from zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
@@ -366,7 +367,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
if self.is_active is not None:
query = query.filter(ApiKey.is_active == self.is_active)
total = query.count()
total = int(query.with_entities(func.count(ApiKey.id)).scalar() or 0)
api_keys = (
query.order_by(ApiKey.created_at.desc()).offset(self.skip).limit(self.limit).all()
)

View File

@@ -13,6 +13,7 @@ from typing import Any, Literal
from fastapi import APIRouter, Depends, Query, Request
from pydantic import BaseModel, Field
from sqlalchemy import func
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
@@ -251,7 +252,7 @@ class BillingRuleListAdapter(AdminApiAdapter):
if self.is_enabled is not None:
q = q.filter(BillingRule.is_enabled == self.is_enabled)
total = q.count()
total = int(q.with_entities(func.count(BillingRule.id)).scalar() or 0)
items = (
q.order_by(BillingRule.updated_at.desc())
.offset((self.page - 1) * self.page_size)
@@ -365,7 +366,7 @@ class DimensionCollectorListAdapter(AdminApiAdapter):
if self.is_enabled is not None:
q = q.filter(DimensionCollector.is_enabled == self.is_enabled)
total = q.count()
total = int(q.with_entities(func.count(DimensionCollector.id)).scalar() or 0)
items = (
q.order_by(DimensionCollector.updated_at.desc())
.offset((self.page - 1) * self.page_size)

View File

@@ -21,7 +21,7 @@ from typing import Any
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
from pydantic import BaseModel
from sqlalchemy import delete, func
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, load_only
from src.clients.http_client import HTTPClientPool
from src.core.crypto import crypto_service
@@ -106,26 +106,45 @@ async def list_file_mappings(
now = datetime.now(timezone.utc)
query = db.query(GeminiFileMapping)
count_query = db.query(func.count(GeminiFileMapping.id))
# 过滤过期
if not include_expired:
query = query.filter(GeminiFileMapping.expires_at > now)
active_filter = GeminiFileMapping.expires_at > now
query = query.filter(active_filter)
count_query = count_query.filter(active_filter)
# 搜索
if search:
search_pattern = f"%{search}%"
query = query.filter(
(GeminiFileMapping.file_name.ilike(search_pattern))
| (GeminiFileMapping.display_name.ilike(search_pattern))
search_filter = (GeminiFileMapping.file_name.ilike(search_pattern)) | (
GeminiFileMapping.display_name.ilike(search_pattern)
)
query = query.filter(search_filter)
count_query = count_query.filter(search_filter)
# 总数
total = query.count()
total = int(count_query.scalar() or 0)
# 分页
offset = (page - 1) * page_size
mappings = (
query.order_by(GeminiFileMapping.created_at.desc()).offset(offset).limit(page_size).all()
query.options(
load_only(
GeminiFileMapping.id,
GeminiFileMapping.file_name,
GeminiFileMapping.key_id,
GeminiFileMapping.user_id,
GeminiFileMapping.display_name,
GeminiFileMapping.mime_type,
GeminiFileMapping.created_at,
GeminiFileMapping.expires_at,
)
)
.order_by(GeminiFileMapping.created_at.desc())
.offset(offset)
.limit(page_size)
.all()
)
# 获取关联的 Key 和 User 信息
@@ -134,12 +153,22 @@ async def list_file_mappings(
keys_map = {}
if key_ids:
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.id.in_(key_ids)).all()
keys = (
db.query(ProviderAPIKey)
.options(load_only(ProviderAPIKey.id, ProviderAPIKey.name))
.filter(ProviderAPIKey.id.in_(key_ids))
.all()
)
keys_map = {str(k.id): k.name for k in keys}
users_map = {}
if user_ids:
users = db.query(User).filter(User.id.in_(user_ids)).all()
users = (
db.query(User)
.options(load_only(User.id, User.username))
.filter(User.id.in_(user_ids))
.all()
)
users_map = {str(u.id): u.username for u in users}
items = []
@@ -199,9 +228,11 @@ async def get_file_mapping_stats(
by_mime_type = {(mt or "unknown"): count for mt, count in mime_stats}
# 有 gemini_files 能力的 Key 数量
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
keys = db.query(ProviderAPIKey.capabilities).filter(ProviderAPIKey.is_active.is_(True)).all()
capable_keys_count = sum(
1 for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
1
for (capabilities,) in keys
if isinstance(capabilities, dict) and capabilities.get("gemini_files", False)
)
return FileMappingStatsResponse(
@@ -294,15 +325,28 @@ async def list_capable_keys(
"""获取所有具有 gemini_files 能力的 Key 列表"""
from src.models.database import Provider
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
key_rows = (
db.query(
ProviderAPIKey.id,
ProviderAPIKey.name,
ProviderAPIKey.provider_id,
ProviderAPIKey.capabilities,
)
.filter(ProviderAPIKey.is_active.is_(True))
.all()
)
capable_keys = [
key for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
key
for key in key_rows
if isinstance(key.capabilities, dict) and key.capabilities.get("gemini_files", False)
]
# 获取 Provider 名称
provider_ids = {key.provider_id for key in capable_keys}
providers = db.query(Provider).filter(Provider.id.in_(provider_ids)).all()
provider_map = {str(p.id): p.name for p in providers}
provider_ids = {key.provider_id for key in capable_keys if key.provider_id}
provider_map: dict[str, str] = {}
if provider_ids:
providers = db.query(Provider.id, Provider.name).filter(Provider.id.in_(provider_ids)).all()
provider_map = {str(provider_id): provider_name for provider_id, provider_name in providers}
return [
CapableKeyResponse(
@@ -454,6 +498,15 @@ async def upload_file(
with create_session() as db:
keys = (
db.query(ProviderAPIKey)
.options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.name,
ProviderAPIKey.api_key,
ProviderAPIKey.capabilities,
ProviderAPIKey.is_active,
)
)
.filter(
ProviderAPIKey.id.in_(key_id_list),
ProviderAPIKey.is_active.is_(True),
@@ -483,6 +536,7 @@ async def upload_file(
now = datetime.now(timezone.utc)
existing = (
db.query(GeminiFileMapping)
.options(load_only(GeminiFileMapping.key_id, GeminiFileMapping.file_name))
.filter(
GeminiFileMapping.source_hash == source_hash,
GeminiFileMapping.key_id.in_(capable_key_ids),

View File

@@ -7,6 +7,7 @@ from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi.responses import JSONResponse
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
@@ -190,7 +191,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
if self.is_active is not None:
query = query.filter(ManagementToken.is_active == self.is_active)
total = query.count()
total = int(query.with_entities(func.count(ManagementToken.id)).scalar() or 0)
tokens = (
query.order_by(ManagementToken.created_at.desc())
.offset(self.skip)

View File

@@ -276,17 +276,24 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
search: str | None
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import and_, case, func
from sqlalchemy import and_, case, func, or_
from src.models.database import Model, Provider
from src.models.database import GlobalModel, Model, Provider
models = GlobalModelService.list_global_models(
db=context.db,
skip=self.skip,
limit=self.limit,
is_active=self.is_active,
search=self.search,
)
query = context.db.query(GlobalModel)
if self.is_active is not None:
query = query.filter(GlobalModel.is_active == self.is_active)
if self.search:
search_pattern = f"%{self.search}%"
query = query.filter(
or_(
GlobalModel.name.ilike(search_pattern),
GlobalModel.display_name.ilike(search_pattern),
)
)
total = int(query.with_entities(func.count(GlobalModel.id)).scalar() or 0)
models = query.order_by(GlobalModel.name).offset(self.skip).limit(self.limit).all()
# 一次性查询所有 GlobalModel 的 provider_count优化 N+1 问题)
# 用条件聚合同时获取总数和活跃数,减少一次 DB 往返
@@ -331,7 +338,7 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
return GlobalModelListResponse(
models=model_responses,
total=len(models),
total=total,
)
@@ -349,10 +356,9 @@ class AdminGetGlobalModelAdapter(AdminApiAdapter):
global_model = GlobalModelService.get_global_model(context.db, self.global_model_id)
stats = GlobalModelService.get_global_model_stats(context.db, self.global_model_id)
# 查询 provider_count 和 active_provider_count与列表 API 一致)
count_row = (
# total_providers 已由 stats 提供,这里只查询活跃 provider 数量
active_count = (
context.db.query(
func.count(func.distinct(Model.provider_id)),
func.count(
func.distinct(
case(
@@ -366,17 +372,17 @@ class AdminGetGlobalModelAdapter(AdminApiAdapter):
else_=None,
)
)
),
)
)
.join(Provider, Model.provider_id == Provider.id)
.filter(Model.global_model_id == global_model.id)
.first()
.scalar()
or 0
)
total_count, active_count = count_row if count_row else (0, 0)
response = GlobalModelResponse.model_validate(global_model)
response.provider_count = total_count
response.active_provider_count = active_count
response.provider_count = stats["total_providers"]
response.active_provider_count = int(active_count)
return GlobalModelWithStats(
**response.model_dump(),

View File

@@ -7,13 +7,14 @@ from datetime import datetime, timedelta, timezone
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
from sqlalchemy import case, func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_query
from src.api.base.pagination import PaginationMeta, build_pagination_payload
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.core.logger import logger
from src.database import get_db
from src.models.database import (
@@ -26,6 +27,7 @@ from src.models.database import (
from src.models.database import User as DBUser
from src.services.health.monitor import HealthMonitor
from src.services.system.audit import audit_service
from src.utils.cache_decorator import cache_result
from src.utils.database_helpers import escape_like_pattern
router = APIRouter(prefix="/api/admin/monitoring", tags=["Admin - Monitoring"])
@@ -223,10 +225,26 @@ class AdminGetAuditLogsAdapter(AdminApiAdapter):
# 查看审计日志本身不应该产生审计记录,避免刷新页面时产生大量无意义的日志
audit_log_enabled: bool = False
@cache_result(
key_prefix="admin:monitoring:audit-logs",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=["username", "event_type", "days", "limit", "offset"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
cutoff_time = datetime.now(timezone.utc) - timedelta(days=self.days)
count_query = db.query(func.count(AuditLog.id)).filter(AuditLog.created_at >= cutoff_time)
if self.username:
escaped = escape_like_pattern(self.username)
count_query = count_query.outerjoin(DBUser, AuditLog.user_id == DBUser.id).filter(
DBUser.username.ilike(f"%{escaped}%", escape="\\")
)
if self.event_type:
count_query = count_query.filter(AuditLog.event_type == self.event_type)
total = int(count_query.scalar() or 0)
base_query = (
db.query(AuditLog, DBUser)
.outerjoin(DBUser, AuditLog.user_id == DBUser.id)
@@ -238,8 +256,12 @@ class AdminGetAuditLogsAdapter(AdminApiAdapter):
if self.event_type:
base_query = base_query.filter(AuditLog.event_type == self.event_type)
ordered_query = base_query.order_by(AuditLog.created_at.desc())
total, logs_with_users = paginate_query(ordered_query, self.limit, self.offset)
logs_with_users = (
base_query.order_by(AuditLog.created_at.desc())
.offset(self.offset)
.limit(self.limit)
.all()
)
items = [
{
@@ -287,39 +309,51 @@ class AdminGetAuditLogsAdapter(AdminApiAdapter):
class AdminSystemStatusAdapter(AdminApiAdapter):
@cache_result(
key_prefix="admin:monitoring:system-status",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
total_users = db.query(func.count(DBUser.id)).scalar()
active_users = db.query(func.count(DBUser.id)).filter(DBUser.is_active.is_(True)).scalar()
user_stats = db.query(
func.count(DBUser.id).label("total"),
func.sum(case((DBUser.is_active.is_(True), 1), else_=0)).label("active"),
).first()
total_users = int((user_stats.total if user_stats else 0) or 0)
active_users = int((user_stats.active if user_stats else 0) or 0)
total_providers = db.query(func.count(Provider.id)).scalar()
active_providers = (
db.query(func.count(Provider.id)).filter(Provider.is_active.is_(True)).scalar()
)
provider_stats = db.query(
func.count(Provider.id).label("total"),
func.sum(case((Provider.is_active.is_(True), 1), else_=0)).label("active"),
).first()
total_providers = int((provider_stats.total if provider_stats else 0) or 0)
active_providers = int((provider_stats.active if provider_stats else 0) or 0)
total_api_keys = db.query(func.count(ApiKey.id)).scalar()
active_api_keys = (
db.query(func.count(ApiKey.id)).filter(ApiKey.is_active.is_(True)).scalar()
)
api_key_stats = db.query(
func.count(ApiKey.id).label("total"),
func.sum(case((ApiKey.is_active.is_(True), 1), else_=0)).label("active"),
).first()
total_api_keys = int((api_key_stats.total if api_key_stats else 0) or 0)
active_api_keys = int((api_key_stats.active if api_key_stats else 0) or 0)
today_start = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
today_requests = (
db.query(func.count(Usage.id)).filter(Usage.created_at >= today_start).scalar()
)
today_tokens = (
db.query(func.sum(Usage.total_tokens)).filter(Usage.created_at >= today_start).scalar()
or 0
)
today_cost = (
db.query(func.sum(Usage.total_cost_usd))
today_stats = (
db.query(
func.count(Usage.id).label("requests"),
func.coalesce(func.sum(Usage.total_tokens), 0).label("tokens"),
func.coalesce(func.sum(Usage.total_cost_usd), 0.0).label("cost"),
)
.filter(Usage.created_at >= today_start)
.scalar()
or 0
.first()
)
today_requests = int((today_stats.requests if today_stats else 0) or 0)
today_tokens = int((today_stats.tokens if today_stats else 0) or 0)
today_cost = float((today_stats.cost if today_stats else 0.0) or 0.0)
recent_errors = (
db.query(AuditLog)
db.query(func.count(AuditLog.id))
.filter(
AuditLog.event_type.in_(
[
@@ -329,20 +363,21 @@ class AdminSystemStatusAdapter(AdminApiAdapter):
),
AuditLog.created_at >= datetime.now(timezone.utc) - timedelta(hours=1),
)
.count()
.scalar()
or 0
)
context.add_audit_metadata(
action="system_status_snapshot",
total_users=int(total_users or 0),
active_users=int(active_users or 0),
total_providers=int(total_providers or 0),
active_providers=int(active_providers or 0),
total_api_keys=int(total_api_keys or 0),
active_api_keys=int(active_api_keys or 0),
today_requests=int(today_requests or 0),
today_tokens=int(today_tokens or 0),
today_cost=float(today_cost or 0.0),
total_users=total_users,
active_users=active_users,
total_providers=total_providers,
active_providers=active_providers,
total_api_keys=total_api_keys,
active_api_keys=active_api_keys,
today_requests=today_requests,
today_tokens=today_tokens,
today_cost=today_cost,
recent_errors=int(recent_errors or 0),
)

View File

@@ -617,13 +617,68 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
.all()
)
# 先批量计算池化 Provider 的 Key 总数/启用数,避免每个 Provider 单独查询N+1
pool_provider_ids: list[str] = []
pool_enabled_map: dict[str, bool] = {}
for p in providers:
pid = str(p.id)
enabled = parse_pool_config(getattr(p, "config", None)) is not None
pool_enabled_map[pid] = enabled
if enabled:
pool_provider_ids.append(pid)
key_ids_by_provider: dict[str, list[str]] = {pid: [] for pid in pool_provider_ids}
key_stats_by_provider: dict[str, dict[str, int]] = {
pid: {"total": 0, "active": 0} for pid in pool_provider_ids
}
if pool_provider_ids:
key_rows = (
db.query(
ProviderAPIKey.provider_id,
ProviderAPIKey.id,
ProviderAPIKey.is_active,
)
.filter(ProviderAPIKey.provider_id.in_(pool_provider_ids))
.all()
)
for provider_id, key_id, is_active in key_rows:
pid = str(provider_id)
kid = str(key_id)
key_ids_by_provider.setdefault(pid, []).append(kid)
stats = key_stats_by_provider.setdefault(pid, {"total": 0, "active": 0})
stats["total"] += 1
if is_active:
stats["active"] += 1
# Redis 冷却状态并发获取,避免逐 Provider 串行等待。
cooldown_count_by_provider: dict[str, int] = {}
cooldown_targets = [
(pid, key_ids) for pid, key_ids in key_ids_by_provider.items() if key_ids
]
if cooldown_targets:
cooldown_results = await asyncio.gather(
*[
pool_redis.batch_get_cooldowns(pid, key_ids)
for pid, key_ids in cooldown_targets
],
return_exceptions=True,
)
for (pid, _key_ids), result in zip(cooldown_targets, cooldown_results, strict=False):
if isinstance(result, Exception):
logger.warning(
"池管理概览读取冷却状态失败",
extra={"provider_id": pid, "error": str(result)},
)
cooldown_count_by_provider[pid] = 0
else:
cooldown_count_by_provider[pid] = sum(
1 for value in result.values() if value is not None
)
items: list[PoolOverviewItem] = []
for p in providers:
pid = str(p.id)
pcfg = parse_pool_config(getattr(p, "config", None))
# Non-pool providers: skip Redis + key queries entirely.
if pcfg is None:
if not pool_enabled_map.get(pid, False):
items.append(
PoolOverviewItem(
provider_id=pid,
@@ -634,22 +689,16 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
)
continue
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == pid).all()
key_ids = [str(k.id) for k in keys]
cooldown_count = 0
if key_ids:
cooldowns = await pool_redis.batch_get_cooldowns(pid, key_ids)
cooldown_count = sum(1 for v in cooldowns.values() if v is not None)
key_stats = key_stats_by_provider.get(pid, {"total": 0, "active": 0})
items.append(
PoolOverviewItem(
provider_id=pid,
provider_name=p.name,
provider_type=str(getattr(p, "provider_type", "custom") or "custom"),
total_keys=len(keys),
active_keys=sum(1 for k in keys if k.is_active),
cooldown_count=cooldown_count,
total_keys=key_stats["total"],
active_keys=key_stats["active"],
cooldown_count=cooldown_count_by_provider.get(pid, 0),
pool_enabled=True,
)
)
@@ -688,7 +737,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
q = q.filter(ProviderAPIKey.is_active.is_(False))
# "cooldown" filtering is done post-query (Redis state)
total = q.count()
total = 0
# For cooldown filtering we need to fetch all, then filter, then paginate.
# Limit scan range to avoid loading the entire table into memory.
@@ -702,6 +751,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
offset = (self.page - 1) * self.page_size
keys = all_keys[offset : offset + self.page_size]
else:
total = int(q.with_entities(func.count(ProviderAPIKey.id)).scalar() or 0)
offset = (self.page - 1) * self.page_size
keys = (
q.order_by(ProviderAPIKey.created_at.desc())

View File

@@ -9,12 +9,14 @@ from typing import Any
from fastapi import APIRouter, Depends, Query, Request
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from sqlalchemy.orm import Session
from sqlalchemy import func
from sqlalchemy.orm import Session, load_only
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.models_service import invalidate_models_list_cache
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.core.enums import ProviderBillingType
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
@@ -27,6 +29,7 @@ from src.models.database import GlobalModel, Provider, ProviderAPIKey, ProviderE
from src.models.endpoint_models import ProviderWithEndpointsSummary
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.provider_cache import ProviderCacheService
from src.utils.cache_decorator import cache_result
from .summary import _build_provider_summary
@@ -341,9 +344,24 @@ class AdminListProvidersAdapter(AdminApiAdapter):
self.limit = limit
self.is_active = is_active
@cache_result(
key_prefix="admin:providers:list",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=["skip", "limit", "is_active"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(Provider)
query = db.query(Provider).options(
load_only(
Provider.id,
Provider.name,
Provider.provider_priority,
Provider.is_active,
Provider.created_at,
Provider.updated_at,
)
)
if self.is_active is not None:
query = query.filter(Provider.is_active == self.is_active)
providers = query.offset(self.skip).limit(self.limit).all()
@@ -723,11 +741,22 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context: ApiRequestContext) -> ProviderMappingPreviewResponse: # type: ignore[override]
@cache_result(
key_prefix="admin:providers:mapping-preview",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=["provider_id"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 获取 Provider
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
provider = (
db.query(Provider)
.options(load_only(Provider.id, Provider.name))
.filter(Provider.id == self.provider_id)
.first()
)
if not provider:
raise NotFoundException("提供商不存在", "provider")
@@ -735,19 +764,6 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
truncated_keys = 0
truncated_models = 0
# 获取该 Provider 有白名单配置的 Key 总数(用于截断统计)
from sqlalchemy import func
total_keys_with_allowed_models = (
db.query(func.count(ProviderAPIKey.id))
.filter(
ProviderAPIKey.provider_id == self.provider_id,
ProviderAPIKey.allowed_models.isnot(None),
)
.scalar()
or 0
)
# 获取该 Provider 有白名单配置的 Key只查询需要的字段
keys = (
db.query(
@@ -761,25 +777,22 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
ProviderAPIKey.provider_id == self.provider_id,
ProviderAPIKey.allowed_models.isnot(None),
)
.limit(MAPPING_PREVIEW_MAX_KEYS)
.limit(MAPPING_PREVIEW_MAX_KEYS + 1)
.all()
)
# 计算被截断的 Key 数量
if total_keys_with_allowed_models > MAPPING_PREVIEW_MAX_KEYS:
truncated_keys = total_keys_with_allowed_models - MAPPING_PREVIEW_MAX_KEYS
# 获取有 model_mappings 配置的 GlobalModel 总数(用于截断统计)
total_models_with_mappings = (
db.query(func.count(GlobalModel.id))
.filter(
GlobalModel.config.isnot(None),
GlobalModel.config["model_mappings"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_mappings"]) > 0,
if len(keys) > MAPPING_PREVIEW_MAX_KEYS:
keys = keys[:MAPPING_PREVIEW_MAX_KEYS]
total_keys_with_allowed_models = (
db.query(func.count(ProviderAPIKey.id))
.filter(
ProviderAPIKey.provider_id == self.provider_id,
ProviderAPIKey.allowed_models.isnot(None),
)
.scalar()
or 0
)
.scalar()
or 0
)
truncated_keys = total_keys_with_allowed_models - MAPPING_PREVIEW_MAX_KEYS
# 只查询有 model_mappings 配置的 GlobalModel使用 SQLAlchemy JSONB 操作符)
global_models = (
@@ -795,12 +808,22 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
GlobalModel.config["model_mappings"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_mappings"]) > 0,
)
.limit(MAPPING_PREVIEW_MAX_MODELS)
.limit(MAPPING_PREVIEW_MAX_MODELS + 1)
.all()
)
# 计算被截断的 GlobalModel 数量
if total_models_with_mappings > MAPPING_PREVIEW_MAX_MODELS:
if len(global_models) > MAPPING_PREVIEW_MAX_MODELS:
global_models = global_models[:MAPPING_PREVIEW_MAX_MODELS]
total_models_with_mappings = (
db.query(func.count(GlobalModel.id))
.filter(
GlobalModel.config.isnot(None),
GlobalModel.config["model_mappings"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_mappings"]) > 0,
)
.scalar()
or 0
)
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
# 构建有映射配置的 GlobalModel 映射
@@ -819,10 +842,10 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
keys=[],
total_keys=0,
total_matches=0,
truncated=False,
truncated_keys=0,
truncated_models=0,
)
truncated=truncated_keys > 0 or truncated_models > 0,
truncated_keys=truncated_keys,
truncated_models=truncated_models,
).model_dump()
key_infos: list[MappingMatchingKey] = []
total_matches = 0
@@ -837,18 +860,6 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
if not allowed_models_list:
continue
# 生成脱敏 Key
masked_key = "***"
if key.api_key:
try:
decrypted_key = crypto.decrypt(key.api_key, silent=True)
if len(decrypted_key) > 8:
masked_key = f"{decrypted_key[:4]}***{decrypted_key[-4:]}"
else:
masked_key = f"{decrypted_key[:2]}***"
except Exception:
pass
# 查找匹配的 GlobalModel
matching_global_models: list[MappingMatchingGlobalModel] = []
@@ -879,6 +890,18 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
total_matches += 1
if matching_global_models:
# 只有有匹配结果的 key 才做解密脱敏,减少 CPU 开销
masked_key = "***"
if key.api_key:
try:
decrypted_key = crypto.decrypt(key.api_key, silent=True)
if len(decrypted_key) > 8:
masked_key = f"{decrypted_key[:4]}***{decrypted_key[-4:]}"
else:
masked_key = f"{decrypted_key[:2]}***"
except Exception:
pass
key_infos.append(
MappingMatchingKey(
key_id=key.id or "",
@@ -901,7 +924,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
truncated=is_truncated,
truncated_keys=truncated_keys,
truncated_models=truncated_models,
)
).model_dump()
# ========== Claude Code Pool Management ==========

View File

@@ -8,12 +8,13 @@ from typing import Any
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import case, func
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, load_only
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.models_service import invalidate_models_list_cache
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.core.enums import ProviderBillingType
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
@@ -39,6 +40,7 @@ from src.models.endpoint_models import (
)
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.provider_cache import ProviderCacheService
from src.utils.cache_decorator import cache_result
router = APIRouter(tags=["Provider Summary"])
pipeline = ApiRequestPipeline()
@@ -130,11 +132,8 @@ async def get_provider_summary(
- `created_at`: 创建时间
- `updated_at`: 更新时间
"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException(f"Provider {provider_id} not found")
return _build_provider_summary(db, provider)
adapter = AdminProviderDetailAdapter(provider_id=provider_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/{provider_id}/health-monitor", response_model=ProviderEndpointHealthMonitorResponse)
@@ -319,12 +318,20 @@ def _extract_failover_rules_from_config(
def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndpointsSummary:
endpoints = db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
endpoints = (
db.query(ProviderEndpoint)
.options(
load_only(
ProviderEndpoint.id,
ProviderEndpoint.provider_id,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
.filter(ProviderEndpoint.provider_id == provider.id)
.all()
)
total_endpoints = len(endpoints)
active_endpoints = sum(1 for e in endpoints if e.is_active)
# Key 统计(合并为单个查询)
key_stats = (
db.query(
func.count(ProviderAPIKey.id).label("total"),
@@ -333,10 +340,9 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
.filter(ProviderAPIKey.provider_id == provider.id)
.first()
)
total_keys = key_stats.total or 0
total_keys = int(key_stats.total or 0)
active_keys = int(key_stats.active or 0)
# Model 统计(合并为单个查询)
model_stats = (
db.query(
func.count(Model.id).label("total"),
@@ -345,25 +351,62 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
.filter(Model.provider_id == provider.id)
.first()
)
total_models = model_stats.total or 0
total_models = int(model_stats.total or 0)
active_models = int(model_stats.active or 0)
# 活跃模型关联的全局模型 ID 列表
global_model_ids = [
row[0]
for row in db.query(Model.global_model_id)
.filter(
Model.provider_id == provider.id,
Model.is_active == True,
Model.global_model_id.isnot(None),
)
.distinct()
.all()
]
api_formats = [e.api_format for e in endpoints]
all_keys = (
db.query(ProviderAPIKey)
.options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.provider_id,
ProviderAPIKey.is_active,
ProviderAPIKey.api_formats,
ProviderAPIKey.health_by_format,
)
)
.filter(ProviderAPIKey.provider_id == provider.id)
.all()
)
# 优化: 一次性加载 Provider 的 keys避免 N+1 查询
all_keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == provider.id).all()
return _compose_provider_summary(
provider=provider,
endpoints=endpoints,
all_keys=all_keys,
total_keys=total_keys,
active_keys=active_keys,
total_models=total_models,
active_models=active_models,
global_model_ids=global_model_ids,
)
def _compose_provider_summary(
*,
provider: Provider,
endpoints: list[ProviderEndpoint],
all_keys: list[ProviderAPIKey],
total_keys: int,
active_keys: int,
total_models: int,
active_models: int,
global_model_ids: list[Any],
) -> ProviderWithEndpointsSummary:
total_endpoints = len(endpoints)
active_endpoints = sum(1 for e in endpoints if e.is_active)
api_formats = [e.api_format for e in endpoints]
# 按 api_formats 分组 keys通过 api_formats 关联)
format_to_endpoint_id: dict[str, str] = {e.api_format: e.id for e in endpoints}
@@ -379,9 +422,8 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
for endpoint in endpoints:
keys = keys_by_endpoint.get(endpoint.id, [])
if keys:
# 从 health_by_format 获取对应格式的健康度
api_fmt = endpoint.api_format
health_scores = []
health_scores: list[float] = []
for k in keys:
health_by_format = k.health_by_format or {}
if api_fmt in health_by_format:
@@ -389,7 +431,7 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
if score is not None:
health_scores.append(float(score))
else:
health_scores.append(1.0) # 默认健康度
health_scores.append(1.0)
avg_health = sum(health_scores) / len(health_scores) if health_scores else 1.0
endpoint_health_map[endpoint.id] = avg_health
else:
@@ -399,7 +441,6 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
avg_health_score = sum(all_health_scores) / len(all_health_scores) if all_health_scores else 1.0
unhealthy_endpoints = sum(1 for score in all_health_scores if score < 0.5)
# 计算每个端点的活跃密钥数量
active_keys_by_endpoint: dict[str, int] = {}
for endpoint_id, keys in keys_by_endpoint.items():
active_keys_by_endpoint[endpoint_id] = sum(1 for k in keys if k.is_active)
@@ -484,6 +525,119 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
)
def _build_provider_summaries_batch(
db: Session, providers: list[Provider]
) -> list[ProviderWithEndpointsSummary]:
if not providers:
return []
provider_ids = [provider.id for provider in providers]
endpoint_rows = (
db.query(ProviderEndpoint)
.options(
load_only(
ProviderEndpoint.id,
ProviderEndpoint.provider_id,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
.filter(ProviderEndpoint.provider_id.in_(provider_ids))
.all()
)
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
for endpoint in endpoint_rows:
endpoints_by_provider.setdefault(str(endpoint.provider_id), []).append(endpoint)
key_rows = (
db.query(ProviderAPIKey)
.options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.provider_id,
ProviderAPIKey.is_active,
ProviderAPIKey.api_formats,
ProviderAPIKey.health_by_format,
)
)
.filter(ProviderAPIKey.provider_id.in_(provider_ids))
.all()
)
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
for key in key_rows:
keys_by_provider.setdefault(str(key.provider_id), []).append(key)
key_stats_rows = (
db.query(
ProviderAPIKey.provider_id.label("provider_id"),
func.count(ProviderAPIKey.id).label("total"),
func.sum(case((ProviderAPIKey.is_active == True, 1), else_=0)).label("active"),
)
.filter(ProviderAPIKey.provider_id.in_(provider_ids))
.group_by(ProviderAPIKey.provider_id)
.all()
)
key_stats_by_provider: dict[str, dict[str, int]] = {
str(row.provider_id): {
"total": int(row.total or 0),
"active": int(row.active or 0),
}
for row in key_stats_rows
}
model_stats_rows = (
db.query(
Model.provider_id.label("provider_id"),
func.count(Model.id).label("total"),
func.sum(case((Model.is_active == True, 1), else_=0)).label("active"),
)
.filter(Model.provider_id.in_(provider_ids))
.group_by(Model.provider_id)
.all()
)
model_stats_by_provider: dict[str, dict[str, int]] = {
str(row.provider_id): {
"total": int(row.total or 0),
"active": int(row.active or 0),
}
for row in model_stats_rows
}
global_model_rows = (
db.query(Model.provider_id, Model.global_model_id)
.filter(
Model.provider_id.in_(provider_ids),
Model.is_active == True,
Model.global_model_id.isnot(None),
)
.distinct()
.all()
)
global_model_ids_by_provider: dict[str, list[Any]] = {}
for provider_id, global_model_id in global_model_rows:
global_model_ids_by_provider.setdefault(str(provider_id), []).append(global_model_id)
summaries: list[ProviderWithEndpointsSummary] = []
for provider in providers:
pid = str(provider.id)
key_stats = key_stats_by_provider.get(pid, {"total": 0, "active": 0})
model_stats = model_stats_by_provider.get(pid, {"total": 0, "active": 0})
summaries.append(
_compose_provider_summary(
provider=provider,
endpoints=endpoints_by_provider.get(pid, []),
all_keys=keys_by_provider.get(pid, []),
total_keys=key_stats["total"],
active_keys=key_stats["active"],
total_models=model_stats["total"],
active_models=model_stats["active"],
global_model_ids=global_model_ids_by_provider.get(pid, []),
)
)
return summaries
# -------- Adapters --------
@@ -493,6 +647,12 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
lookback_hours: int
per_endpoint_limit: int
@cache_result(
key_prefix="admin:providers:health-monitor",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=["provider_id", "lookback_hours", "per_endpoint_limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
@@ -501,6 +661,14 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
endpoints = (
db.query(ProviderEndpoint)
.options(
load_only(
ProviderEndpoint.id,
ProviderEndpoint.provider_id,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
.filter(ProviderEndpoint.provider_id == self.provider_id)
.all()
)
@@ -508,7 +676,7 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
now = datetime.now(timezone.utc)
since = now - timedelta(hours=self.lookback_hours)
endpoint_ids = [endpoint.id for endpoint in endpoints]
endpoint_ids = [str(endpoint.id) for endpoint in endpoints]
if not endpoint_ids:
response = ProviderEndpointHealthMonitorResponse(
provider_id=provider.id,
@@ -522,46 +690,65 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
endpoint_count=0,
lookback_hours=self.lookback_hours,
)
return response
return response.model_dump()
limit_rows = max(200, self.per_endpoint_limit * max(1, len(endpoint_ids)) * 2)
attempts_query = (
ranked_attempts_subq = (
db.query(RequestCandidate)
.with_entities(
RequestCandidate.endpoint_id.label("endpoint_id"),
RequestCandidate.status.label("status"),
RequestCandidate.status_code.label("status_code"),
RequestCandidate.latency_ms.label("latency_ms"),
RequestCandidate.error_type.label("error_type"),
RequestCandidate.error_message.label("error_message"),
func.coalesce(
RequestCandidate.finished_at,
RequestCandidate.started_at,
RequestCandidate.created_at,
).label("event_timestamp"),
func.row_number()
.over(
partition_by=RequestCandidate.endpoint_id,
order_by=RequestCandidate.created_at.desc(),
)
.label("rn"),
)
.filter(
RequestCandidate.endpoint_id.in_(endpoint_ids),
RequestCandidate.created_at >= since,
)
.order_by(RequestCandidate.created_at.desc())
.subquery()
)
attempt_rows = (
db.query(ranked_attempts_subq)
.filter(ranked_attempts_subq.c.rn <= self.per_endpoint_limit)
.order_by(
ranked_attempts_subq.c.endpoint_id.asc(),
ranked_attempts_subq.c.event_timestamp.asc(),
)
.all()
)
attempts = attempts_query.limit(limit_rows).all()
buffered_attempts: dict[str, list[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
counters: dict[str, int] = {eid: 0 for eid in endpoint_ids}
for attempt in attempts:
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
events_by_endpoint: dict[str, list[EndpointHealthEvent]] = {eid: [] for eid in endpoint_ids}
for row in attempt_rows:
endpoint_id = str(row.endpoint_id) if row.endpoint_id is not None else ""
if not endpoint_id or endpoint_id not in events_by_endpoint:
continue
if counters[attempt.endpoint_id] >= self.per_endpoint_limit:
continue
buffered_attempts[attempt.endpoint_id].append(attempt)
counters[attempt.endpoint_id] += 1
events_by_endpoint[endpoint_id].append(
EndpointHealthEvent(
timestamp=row.event_timestamp,
status=row.status,
status_code=row.status_code,
latency_ms=row.latency_ms,
error_type=row.error_type,
error_message=row.error_message,
)
)
endpoint_monitors: list[EndpointHealthMonitor] = []
for endpoint in endpoints:
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
events: list[EndpointHealthEvent] = []
for attempt in attempt_list:
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
events.append(
EndpointHealthEvent(
timestamp=event_timestamp,
status=attempt.status,
status_code=attempt.status_code,
latency_ms=attempt.latency_ms,
error_type=attempt.error_type,
error_message=attempt.error_message,
)
)
endpoint_id = str(endpoint.id)
events = events_by_endpoint.get(endpoint_id, [])
success_count = sum(1 for event in events if event.status == "success")
failed_count = sum(1 for event in events if event.status == "failed")
@@ -598,10 +785,15 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
lookback_hours=self.lookback_hours,
per_endpoint_limit=self.per_endpoint_limit,
)
return response
return response.model_dump()
class AdminProviderSummaryAdapter(AdminApiAdapter):
@cache_result(
key_prefix="admin:providers:summary",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
providers = (
@@ -609,7 +801,25 @@ class AdminProviderSummaryAdapter(AdminApiAdapter):
.order_by(Provider.provider_priority.asc(), Provider.created_at.asc())
.all()
)
return [_build_provider_summary(db, provider) for provider in providers]
return [item.model_dump() for item in _build_provider_summaries_batch(db, providers)]
@dataclass
class AdminProviderDetailAdapter(AdminApiAdapter):
provider_id: str
@cache_result(
key_prefix="admin:providers:summary:detail",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=["provider_id"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
raise NotFoundException(f"Provider {self.provider_id} not found")
return _build_provider_summary(db, provider).model_dump()
@dataclass

View File

@@ -10,9 +10,11 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.services.system.stats_aggregator import AggregatedStats, StatsFilter, query_stats_hybrid
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import pipeline
@@ -34,6 +36,18 @@ class AdminComparisonAdapter(AdminApiAdapter):
self.timezone_name = timezone_name
self.tz_offset_minutes = tz_offset_minutes
@cache_result(
key_prefix="admin:stats:comparison",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"current_start",
"current_end",
"comparison_type",
"timezone_name",
"tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if self.current_start > self.current_end:
raise HTTPException(status_code=400, detail="current_start must be <= current_end")

View File

@@ -11,10 +11,12 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.models.database import Usage
from src.services.system.stats_aggregator import query_time_series
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import (
_apply_admin_default_range,
@@ -40,11 +42,27 @@ class AdminCostForecastAdapter(AdminApiAdapter):
self.days = days
self.forecast_days = forecast_days
self.timezone_name = timezone_name
self.tz_offset_minutes = tz_offset_minutes
self.fallback_tz_offset_minutes = tz_offset_minutes
@cache_result(
key_prefix="admin:stats:cost:forecast",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"days",
"forecast_days",
"timezone_name",
"fallback_tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
time_range = self.time_range or _build_time_range_from_days(
self.days, self.timezone_name, self.tz_offset_minutes
self.days, self.timezone_name, self.fallback_tz_offset_minutes
)
time_range.granularity = "day"
try:
@@ -120,6 +138,20 @@ class AdminCostSavingsAdapter(AdminApiAdapter):
self.provider_name = provider_name
self.model = model
@cache_result(
key_prefix="admin:stats:cost:savings",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"provider_name",
"model",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return {

View File

@@ -11,9 +11,11 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.models.database import StatsDailyError, Usage
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import _apply_admin_default_range, _build_time_range_params, pipeline
@@ -24,6 +26,18 @@ class AdminErrorDistributionAdapter(AdminApiAdapter):
def __init__(self, time_range: TimeRangeParams | None) -> None:
self.time_range = _apply_admin_default_range(time_range)
@cache_result(
key_prefix="admin:stats:errors:distribution",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return {"distribution": [], "trend": []}

View File

@@ -10,10 +10,12 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.models.database import StatsDaily
from src.services.system.stats_aggregator import StatsAggregatorService
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import _apply_admin_default_range, _build_time_range_params, pipeline
@@ -24,6 +26,18 @@ class AdminPercentilesAdapter(AdminApiAdapter):
def __init__(self, time_range: TimeRangeParams | None) -> None:
self.time_range = _apply_admin_default_range(time_range)
@cache_result(
key_prefix="admin:stats:performance:percentiles",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return []

View File

@@ -10,9 +10,11 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.core.enums import ProviderBillingType
from src.database import get_db
from src.models.database import Provider
from src.utils.cache_decorator import cache_result
from .common import pipeline
@@ -20,6 +22,11 @@ router = APIRouter()
class AdminQuotaUsageAdapter(AdminApiAdapter):
@cache_result(
key_prefix="admin:stats:providers:quota_usage",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
providers = (

View File

@@ -10,9 +10,11 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.services.system.stats_aggregator import TimeSeriesFilter, query_time_series
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import _apply_admin_default_range, _build_time_range_params, pipeline
@@ -32,6 +34,22 @@ class AdminTimeSeriesAdapter(AdminApiAdapter):
self.model = model
self.provider_name = provider_name
@cache_result(
key_prefix="admin:stats:time_series",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"time_range.granularity",
"user_id",
"model",
"provider_name",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return []

View File

@@ -9,11 +9,13 @@ from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import ValidationError
from sqlalchemy import case, func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
from src.core.logger import logger
from src.database import get_db
@@ -21,6 +23,7 @@ from src.models.api import SystemSettingsRequest, SystemSettingsResponse
from src.models.database import ApiKey, Provider, Usage, User
from src.services.email.email_template import EmailTemplate
from src.services.system.config import SystemConfigService
from src.utils.cache_decorator import cache_result
router = APIRouter(prefix="/api/admin/system", tags=["Admin - System"])
@@ -749,14 +752,27 @@ class AdminDeleteSystemConfigAdapter(AdminApiAdapter):
class AdminSystemStatsAdapter(AdminApiAdapter):
@cache_result(
key_prefix="admin:system:stats",
ttl=CacheTTL.DASHBOARD_STATS,
user_specific=False,
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
total_users = db.query(User).count()
active_users = db.query(User).filter(User.is_active.is_(True)).count()
total_providers = db.query(Provider).count()
active_providers = db.query(Provider).filter(Provider.is_active.is_(True)).count()
total_api_keys = db.query(ApiKey).count()
total_requests = db.query(Usage).count()
user_stats = db.query(
func.count(User.id).label("total"),
func.sum(case((User.is_active.is_(True), 1), else_=0)).label("active"),
).first()
provider_stats = db.query(
func.count(Provider.id).label("total"),
func.sum(case((Provider.is_active.is_(True), 1), else_=0)).label("active"),
).first()
total_api_keys = int(db.query(func.count(ApiKey.id)).scalar() or 0)
total_requests = int(db.query(func.count(Usage.id)).scalar() or 0)
total_users = int(user_stats.total or 0) if user_stats else 0
active_users = int(user_stats.active or 0) if user_stats else 0
total_providers = int(provider_stats.total or 0) if provider_stats else 0
active_providers = int(provider_stats.active or 0) if provider_stats else 0
return {
"users": {"total": total_users, "active": active_users},
@@ -776,16 +792,18 @@ class AdminTriggerCleanupAdapter(AdminApiAdapter):
db = context.db
# 获取清理前的统计信息
total_before = db.query(Usage).count()
total_before = int(db.query(func.count(Usage.id)).scalar() or 0)
with_body_before = (
db.query(Usage)
db.query(func.count(Usage.id))
.filter((Usage.request_body.isnot(None)) | (Usage.response_body.isnot(None)))
.count()
.scalar()
or 0
)
with_headers_before = (
db.query(Usage)
db.query(func.count(Usage.id))
.filter((Usage.request_headers.isnot(None)) | (Usage.response_headers.isnot(None)))
.count()
.scalar()
or 0
)
# 触发清理
@@ -793,16 +811,18 @@ class AdminTriggerCleanupAdapter(AdminApiAdapter):
await maintenance_scheduler._perform_cleanup()
# 获取清理后的统计信息
total_after = db.query(Usage).count()
total_after = int(db.query(func.count(Usage.id)).scalar() or 0)
with_body_after = (
db.query(Usage)
db.query(func.count(Usage.id))
.filter((Usage.request_body.isnot(None)) | (Usage.response_body.isnot(None)))
.count()
.scalar()
or 0
)
with_headers_after = (
db.query(Usage)
db.query(func.count(Usage.id))
.filter((Usage.request_headers.isnot(None)) | (Usage.response_headers.isnot(None)))
.count()
.scalar()
or 0
)
return {
@@ -2402,11 +2422,11 @@ class AdminPurgeConfigAdapter(AdminApiAdapter):
db = context.db
# 统计
providers_count = db.query(Provider).count()
endpoints_count = db.query(ProviderEndpoint).count()
keys_count = db.query(ProviderAPIKey).count()
models_count = db.query(Model).count()
global_models_count = db.query(GlobalModel).count()
providers_count = int(db.query(func.count(Provider.id)).scalar() or 0)
endpoints_count = int(db.query(func.count(ProviderEndpoint.id)).scalar() or 0)
keys_count = int(db.query(func.count(ProviderAPIKey.id)).scalar() or 0)
models_count = int(db.query(func.count(Model.id)).scalar() or 0)
global_models_count = int(db.query(func.count(GlobalModel.id)).scalar() or 0)
# VideoTask 的 provider_id/endpoint_id/key_id 无 ondelete置 NULL 保留任务记录
db.query(VideoTask).filter(
@@ -2471,7 +2491,9 @@ class AdminPurgeUsersAdapter(AdminApiAdapter):
)
# 统计关联 API Keys 数量DB 级别 CASCADE 会随 User 自动删除)
keys_count = db.query(ApiKey).filter(ApiKey.user_id.in_(user_ids)).count()
keys_count = int(
db.query(func.count(ApiKey.id)).filter(ApiKey.user_id.in_(user_ids)).scalar() or 0
)
# 将使用记录的 user_id 置空(保留记录)
db.query(Usage).filter(Usage.user_id.in_(user_ids)).update(
@@ -2565,9 +2587,9 @@ class AdminPurgeUsageAdapter(AdminApiAdapter):
db = context.db
usage_count = db.query(Usage).count()
candidates_count = db.query(RequestCandidate).count()
usage_counts_count = db.query(UserModelUsageCount).count()
usage_count = int(db.query(func.count(Usage.id)).scalar() or 0)
candidates_count = int(db.query(func.count(RequestCandidate.id)).scalar() or 0)
usage_counts_count = int(db.query(func.count(UserModelUsageCount.id)).scalar() or 0)
# 清空使用记录
db.query(RequestCandidate).delete()
@@ -2594,7 +2616,7 @@ class AdminPurgeAuditLogsAdapter(AdminApiAdapter):
db = context.db
count = db.query(AuditLog).count()
count = int(db.query(func.count(AuditLog.id)).scalar() or 0)
db.query(AuditLog).delete()
db.commit()
@@ -2612,8 +2634,8 @@ class AdminPurgeRequestBodiesAdapter(AdminApiAdapter):
db = context.db
# 统计有 body 的记录数
with_body = (
db.query(Usage)
with_body = int(
db.query(func.count(Usage.id))
.filter(
(Usage.request_body.isnot(None))
| (Usage.response_body.isnot(None))
@@ -2624,7 +2646,8 @@ class AdminPurgeRequestBodiesAdapter(AdminApiAdapter):
| (Usage.provider_request_body_compressed.isnot(None))
| (Usage.client_response_body_compressed.isnot(None))
)
.count()
.scalar()
or 0
)
# 批量清空所有 body 字段

View File

@@ -2117,7 +2117,7 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
async def get_interval_timeline(
request: Request,
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
limit: int = Query(3000, ge=100, le=50000, description="最大返回数据点数量"),
user_id: str | None = Query(None, description="指定用户 ID"),
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
db: Session = Depends(get_db),
@@ -2156,6 +2156,12 @@ class IntervalTimelineAdapter(AdminApiAdapter):
self.user_id = user_id
self.include_user_info = include_user_info
@cache_result(
key_prefix="admin:usage:interval_timeline",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["hours", "limit", "user_id", "include_user_info"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db

View File

@@ -7,11 +7,13 @@ from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from pydantic import ValidationError
from sqlalchemy.orm import Session
from sqlalchemy import func
from sqlalchemy.orm import Session, load_only
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
from src.core.logger import logger
from src.database import get_db
@@ -21,6 +23,7 @@ from src.models.database import ApiKey, User, UserRole
from src.services.system.config import SystemConfigService
from src.services.user.apikey import ApiKeyService
from src.services.user.service import UserService
from src.utils.cache_decorator import cache_result
router = APIRouter(prefix="/api/admin/users", tags=["Admin - Users"])
pipeline = ApiRequestPipeline()
@@ -282,10 +285,43 @@ class AdminListUsersAdapter(AdminApiAdapter):
self.role = role
self.is_active = is_active
@cache_result(
key_prefix="admin:users:list",
ttl=CacheTTL.USER,
user_specific=False,
vary_by=["skip", "limit", "role", "is_active"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
role_enum = UserRole[self.role.upper()] if self.role else None
users = UserService.list_users(db, self.skip, self.limit, role_enum, self.is_active)
role_enum = None
if self.role:
try:
role_enum = UserRole[self.role.upper()]
except KeyError as exc:
raise InvalidRequestException("角色参数不合法") from exc
query = db.query(User).options(
load_only(
User.id,
User.email,
User.username,
User.role,
User.allowed_providers,
User.allowed_api_formats,
User.allowed_models,
User.quota_usd,
User.used_usd,
User.total_usd,
User.is_active,
User.created_at,
)
)
if role_enum:
query = query.filter(User.role == role_enum)
if self.is_active is not None:
query = query.filter(User.is_active == self.is_active)
users = query.order_by(User.created_at.desc()).offset(self.skip).limit(self.limit).all()
return [
{
"id": u.id,
@@ -416,7 +452,9 @@ class AdminDeleteUserAdapter(AdminApiAdapter):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在")
if user.role == UserRole.ADMIN:
admin_count = db.query(User).filter(User.role == UserRole.ADMIN).count()
admin_count = int(
db.query(func.count(User.id)).filter(User.role == UserRole.ADMIN).scalar() or 0
)
if admin_count <= 1:
raise InvalidRequestException("不能删除最后一个管理员账户")

View File

@@ -18,11 +18,13 @@ from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.api.dashboard.routes import DashboardAdapter
from src.clients.http_client import HTTPClientPool
from src.config.constants import CacheTTL
from src.core.crypto import crypto_service
from src.core.enums import UserRole
from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.utils.cache_decorator import cache_result
router = APIRouter(prefix="/api/admin/video-tasks", tags=["Admin - Video Tasks"])
pipeline = ApiRequestPipeline()
@@ -247,6 +249,12 @@ class VideoTaskListAdapter(DashboardAdapter):
page: int
page_size: int
@cache_result(
key_prefix="admin:video_tasks:list",
ttl=min(5, CacheTTL.ADMIN_USAGE_RECORDS),
user_specific=True,
vary_by=["status", "user_id", "model", "page", "page_size"],
)
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
user = context.user
@@ -270,8 +278,8 @@ class VideoTaskListAdapter(DashboardAdapter):
escaped = self.model.replace("%", "\\%").replace("_", "\\_")
query = query.filter(VideoTask.model.ilike(f"%{escaped}%"))
# 统计总数
total = query.count()
# 统计总数(避免 Query.count() 生成大子查询)
total = int(query.with_entities(func.count(VideoTask.id)).scalar() or 0)
# 分页
offset = (self.page - 1) * self.page_size
@@ -341,6 +349,11 @@ class VideoTaskListAdapter(DashboardAdapter):
class VideoTaskStatsAdapter(DashboardAdapter):
"""视频任务统计适配器"""
@cache_result(
key_prefix="admin:video_tasks:stats",
ttl=min(5, CacheTTL.ADMIN_USAGE_RECORDS),
user_specific=True,
)
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
user = context.user
@@ -350,8 +363,8 @@ class VideoTaskStatsAdapter(DashboardAdapter):
if not is_admin:
base_query = base_query.filter(VideoTask.user_id == user.id)
# 总数
total = base_query.count()
# 总数(避免 Query.count() 生成大子查询)
total = int(base_query.with_entities(func.count(VideoTask.id)).scalar() or 0)
# 按状态分组
status_stats = (
@@ -379,7 +392,12 @@ class VideoTaskStatsAdapter(DashboardAdapter):
# 今日任务数
today = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
today_count = base_query.filter(VideoTask.created_at >= today).count()
today_count = int(
base_query.filter(VideoTask.created_at >= today)
.with_entities(func.count(VideoTask.id))
.scalar()
or 0
)
# 管理员额外统计
result = {