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 = {

View File

@@ -2,6 +2,7 @@ from collections.abc import Sequence
from dataclasses import asdict, dataclass
from typing import Any, TypeVar
from sqlalchemy import func
from sqlalchemy.orm import Query
T = TypeVar("T")
@@ -22,7 +23,7 @@ def paginate_query(query: Query, limit: int, offset: int) -> tuple[int, list[T]]
"""
对 SQLAlchemy 查询应用 limit/offset并返回总数与结果列表。
"""
total = query.order_by(None).count()
total = int(query.order_by(None).with_entities(func.count()).scalar() or 0)
records = query.offset(offset).limit(limit).all()
return total, records

View File

@@ -7,7 +7,7 @@ from datetime import date, datetime, timedelta, timezone
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import and_, func
from sqlalchemy import and_, case, func
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter, ApiMode
@@ -661,95 +661,96 @@ class UserDashboardStatsAdapter(DashboardAdapter):
month_start_local = today_local.replace(day=1)
month_start = month_start_local.astimezone(timezone.utc)
user_api_keys = db.query(func.count(ApiKey.id)).filter(ApiKey.user_id == user.id).scalar()
active_keys = (
db.query(func.count(ApiKey.id))
.filter(and_(ApiKey.user_id == user.id, ApiKey.is_active.is_(True)))
.scalar()
)
# 全局 Token 统计
all_time_token_stats = (
api_key_stats = (
db.query(
func.sum(Usage.input_tokens).label("input_tokens"),
func.sum(Usage.output_tokens).label("output_tokens"),
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
func.count(ApiKey.id).label("total"),
func.sum(case((ApiKey.is_active.is_(True), 1), else_=0)).label("active"),
)
.filter(ApiKey.user_id == user.id)
.first()
)
user_api_keys = int(api_key_stats.total or 0) if api_key_stats else 0
active_keys = int(api_key_stats.active or 0) if api_key_stats else 0
# 使用单次聚合查询返回全量 + 本月 + 今日 + 昨日统计
usage_stats = (
db.query(
# 全量 Token 统计
func.sum(Usage.input_tokens).label("all_time_input_tokens"),
func.sum(Usage.output_tokens).label("all_time_output_tokens"),
func.sum(Usage.cache_creation_input_tokens).label("all_time_cache_creation_tokens"),
func.sum(Usage.cache_read_input_tokens).label("all_time_cache_read_tokens"),
# 本月
func.sum(case((Usage.created_at >= month_start, 1), else_=0)).label(
"monthly_requests"
),
func.sum(
case((Usage.created_at >= month_start, Usage.total_cost_usd), else_=0.0)
).label("monthly_cost"),
func.sum(
case(
(Usage.created_at >= month_start, Usage.cache_creation_input_tokens),
else_=0,
)
).label("monthly_cache_creation_tokens"),
func.sum(
case((Usage.created_at >= month_start, Usage.cache_read_input_tokens), else_=0)
).label("monthly_cache_read_tokens"),
func.sum(
case((Usage.created_at >= month_start, Usage.input_tokens), else_=0)
).label("monthly_input_tokens"),
# 今日
func.sum(case((Usage.created_at >= today, 1), else_=0)).label("today_requests"),
func.sum(case((Usage.created_at >= today, Usage.total_cost_usd), else_=0.0)).label(
"today_cost"
),
func.sum(case((Usage.created_at >= today, Usage.total_tokens), else_=0)).label(
"today_tokens"
),
func.sum(
case((Usage.created_at >= today, Usage.cache_creation_input_tokens), else_=0)
).label("today_cache_creation_tokens"),
func.sum(
case((Usage.created_at >= today, Usage.cache_read_input_tokens), else_=0)
).label("today_cache_read_tokens"),
# 昨日(用于变化趋势)
func.sum(
case(
(
and_(
Usage.created_at >= yesterday,
Usage.created_at < today,
),
1,
),
else_=0,
)
).label("yesterday_requests"),
)
.filter(Usage.user_id == user.id)
.first()
)
all_time_input_tokens = (
int(all_time_token_stats.input_tokens or 0) if all_time_token_stats else 0
)
all_time_output_tokens = (
int(all_time_token_stats.output_tokens or 0) if all_time_token_stats else 0
)
all_time_input_tokens = int(usage_stats.all_time_input_tokens or 0) if usage_stats else 0
all_time_output_tokens = int(usage_stats.all_time_output_tokens or 0) if usage_stats else 0
all_time_cache_creation = (
int(all_time_token_stats.cache_creation_tokens or 0) if all_time_token_stats else 0
)
all_time_cache_read = (
int(all_time_token_stats.cache_read_tokens or 0) if all_time_token_stats else 0
int(usage_stats.all_time_cache_creation_tokens or 0) if usage_stats else 0
)
all_time_cache_read = int(usage_stats.all_time_cache_read_tokens or 0) if usage_stats else 0
# 本月请求统计
user_requests = (
db.query(func.count(Usage.id))
.filter(and_(Usage.user_id == user.id, Usage.created_at >= month_start))
.scalar()
)
user_cost = (
db.query(func.sum(Usage.total_cost_usd))
.filter(and_(Usage.user_id == user.id, Usage.created_at >= month_start))
.scalar()
or 0
)
user_requests = int(usage_stats.monthly_requests or 0) if usage_stats else 0
user_cost = float(usage_stats.monthly_cost or 0.0) if usage_stats else 0.0
# 今日统计
requests_today = (
db.query(func.count(Usage.id))
.filter(and_(Usage.user_id == user.id, Usage.created_at >= today))
.scalar()
)
cost_today = (
db.query(func.sum(Usage.total_cost_usd))
.filter(and_(Usage.user_id == user.id, Usage.created_at >= today))
.scalar()
or 0
)
tokens_today = (
db.query(func.sum(Usage.total_tokens))
.filter(and_(Usage.user_id == user.id, Usage.created_at >= today))
.scalar()
or 0
)
requests_today = int(usage_stats.today_requests or 0) if usage_stats else 0
cost_today = float(usage_stats.today_cost or 0.0) if usage_stats else 0.0
tokens_today = int(usage_stats.today_tokens or 0) if usage_stats else 0
requests_yesterday = int(usage_stats.yesterday_requests or 0) if usage_stats else 0
# 昨日统计(用于计算变化)
requests_yesterday = (
db.query(func.count(Usage.id))
.filter(
and_(
Usage.user_id == user.id,
Usage.created_at >= yesterday,
Usage.created_at < today,
)
)
.scalar()
cache_creation_tokens = (
int(usage_stats.monthly_cache_creation_tokens or 0) if usage_stats else 0
)
# 缓存统计(本月)
cache_stats = (
db.query(
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
func.sum(Usage.input_tokens).label("total_input_tokens"),
)
.filter(and_(Usage.user_id == user.id, Usage.created_at >= month_start))
.first()
)
cache_creation_tokens = int(cache_stats.cache_creation_tokens or 0) if cache_stats else 0
cache_read_tokens = int(cache_stats.cache_read_tokens or 0) if cache_stats else 0
monthly_input_tokens = int(cache_stats.total_input_tokens or 0) if cache_stats else 0
cache_read_tokens = int(usage_stats.monthly_cache_read_tokens or 0) if usage_stats else 0
monthly_input_tokens = int(usage_stats.monthly_input_tokens or 0) if usage_stats else 0
# 计算本月缓存命中率cache_read / (input_tokens + cache_read)
# input_tokens 是实际发送给模型的输入不含缓存读取cache_read 是从缓存读取的
@@ -762,19 +763,11 @@ class UserDashboardStatsAdapter(DashboardAdapter):
)
# 今日缓存统计
cache_stats_today = (
db.query(
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
)
.filter(and_(Usage.user_id == user.id, Usage.created_at >= today))
.first()
)
cache_creation_tokens_today = (
int(cache_stats_today.cache_creation_tokens or 0) if cache_stats_today else 0
int(usage_stats.today_cache_creation_tokens or 0) if usage_stats else 0
)
cache_read_tokens_today = (
int(cache_stats_today.cache_read_tokens or 0) if cache_stats_today else 0
int(usage_stats.today_cache_read_tokens or 0) if usage_stats else 0
)
# 配额状态
@@ -859,6 +852,12 @@ class UserDashboardStatsAdapter(DashboardAdapter):
class DashboardRecentRequestsAdapter(DashboardAdapter):
limit: int
@cache_result(
key_prefix="dashboard:recent:requests",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=True,
vary_by=["limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
@@ -1003,28 +1002,53 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
# 补充 unique_models / unique_providers
# query_time_series 使用小时粒度数据,不含这些维度统计
# 直接从 Usage 表按本地日 UTC 范围查询,避免 StatsDaily 历史数据未回填的问题
# 使用 CASE 一次性分桶,避免按天循环查询造成 N 次 SQL。
granularity = (self.time_range.granularity or "day").lower()
if formatted and granularity == "day":
local_days = self.time_range.get_local_day_hours()
enrichment: dict[str, dict] = {}
local_days = self.time_range.get_local_day_hours()
range_start = local_days[0][1] if local_days else None
range_end = local_days[-1][2] if local_days else None
day_bucket = (
case(
*[
(
and_(Usage.created_at >= day_start, Usage.created_at < day_end),
local_date.isoformat(),
)
for local_date, day_start, day_end in local_days
],
else_=None,
).label("local_day")
if local_days
else None
)
for local_date, day_start_utc, day_end_utc in local_days:
q = db.query(
func.count(func.distinct(Usage.model)).label("um"),
func.count(func.distinct(Usage.provider_name)).label("up"),
).filter(
Usage.created_at >= day_start_utc,
Usage.created_at < day_end_utc,
)
if not is_admin:
q = q.filter(Usage.user_id == user.id)
row = q.first()
if row:
enrichment[local_date.isoformat()] = {
"unique_models": row.um or 0,
"unique_providers": row.up or 0,
}
if (
formatted
and granularity == "day"
and day_bucket is not None
and range_start
and range_end
):
enrichment: dict[str, dict[str, int]] = {}
enrich_query = db.query(
day_bucket,
func.count(func.distinct(Usage.model)).label("um"),
func.count(func.distinct(Usage.provider_name)).label("up"),
).filter(
Usage.created_at >= range_start,
Usage.created_at < range_end,
)
if not is_admin:
enrich_query = enrich_query.filter(Usage.user_id == user.id)
enrich_rows = enrich_query.group_by(day_bucket).all()
for local_day, unique_models, unique_providers in enrich_rows:
if not local_day:
continue
enrichment[str(local_day)] = {
"unique_models": int(unique_models or 0),
"unique_providers": int(unique_providers or 0),
}
for item in formatted:
date_key = item["date"][:10] # YYYY-MM-DD
@@ -1062,26 +1086,35 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
# Daily model breakdown (aligned to local days)
breakdown_map: dict[str, list[dict]] = {}
for local_date, day_start, day_end in self.time_range.get_local_day_hours():
day_query = db.query(
if granularity == "day" and day_bucket is not None and range_start and range_end:
breakdown_query = db.query(
day_bucket,
Usage.model,
func.count(Usage.id).label("requests"),
func.sum(Usage.total_tokens).label("tokens"),
func.sum(Usage.total_cost_usd).label("cost"),
).filter(Usage.created_at >= day_start, Usage.created_at < day_end)
).filter(
Usage.created_at >= range_start,
Usage.created_at < range_end,
)
if not is_admin:
day_query = day_query.filter(Usage.user_id == user.id)
day_stats = day_query.group_by(Usage.model).all()
breakdown_map[local_date.isoformat()] = [
{
"model": stat.model,
"requests": stat.requests or 0,
"tokens": int(stat.tokens or 0),
"cost": float(stat.cost or 0),
}
for stat in day_stats
if stat.model
]
breakdown_query = breakdown_query.filter(Usage.user_id == user.id)
breakdown_rows = (
breakdown_query.group_by(day_bucket, Usage.model)
.order_by(day_bucket.asc(), func.sum(Usage.total_cost_usd).desc())
.all()
)
for local_day, model_name, requests, tokens, cost in breakdown_rows:
if not local_day or not model_name:
continue
breakdown_map.setdefault(str(local_day), []).append(
{
"model": model_name,
"requests": int(requests or 0),
"tokens": int(tokens or 0),
"cost": float(cost or 0.0),
}
)
for item in formatted:
item["model_breakdown"] = breakdown_map.get(item["date"], [])

View File

@@ -11,12 +11,13 @@ from datetime import datetime, timedelta, timezone
from typing import Any
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import and_, or_
from sqlalchemy.orm import Session, joinedload
from sqlalchemy import and_, func, or_
from sqlalchemy.orm import Session, joinedload, load_only
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
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.api import (
@@ -40,6 +41,7 @@ from src.models.endpoint_models import (
)
from src.services.health.endpoint import EndpointHealthService
from src.services.system.config import SystemConfigService
from src.utils.cache_decorator import cache_result
router = APIRouter(prefix="/api/public", tags=["System Catalog"])
pipeline = ApiRequestPipeline()
@@ -300,33 +302,84 @@ class PublicProvidersAdapter(PublicApiAdapter):
skip: int
limit: int
@cache_result(
key_prefix="public:catalog:providers",
ttl=CacheTTL.PROVIDER,
user_specific=False,
vary_by=["is_active", "skip", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求提供商列表")
query = db.query(Provider)
query = db.query(Provider).options(
load_only(
Provider.id,
Provider.name,
Provider.description,
Provider.is_active,
Provider.provider_priority,
)
)
if self.is_active is not None:
query = query.filter(Provider.is_active == self.is_active)
else:
query = query.filter(Provider.is_active.is_(True))
providers = query.offset(self.skip).limit(self.limit).all()
provider_ids = [provider.id for provider in providers]
models_count_map: dict[str, int] = {}
active_models_count_map: dict[str, int] = {}
endpoints_count_map: dict[str, int] = {}
active_endpoints_count_map: dict[str, int] = {}
if provider_ids:
model_counts = (
db.query(Model.provider_id, func.count(Model.id))
.filter(Model.provider_id.in_(provider_ids))
.group_by(Model.provider_id)
.all()
)
models_count_map = {provider_id: int(count) for provider_id, count in model_counts}
active_model_counts = (
db.query(Model.provider_id, func.count(Model.id))
.filter(Model.provider_id.in_(provider_ids), Model.is_active.is_(True))
.group_by(Model.provider_id)
.all()
)
active_models_count_map = {
provider_id: int(count) for provider_id, count in active_model_counts
}
endpoint_counts = (
db.query(ProviderEndpoint.provider_id, func.count(ProviderEndpoint.id))
.filter(ProviderEndpoint.provider_id.in_(provider_ids))
.group_by(ProviderEndpoint.provider_id)
.all()
)
endpoints_count_map = {
provider_id: int(count) for provider_id, count in endpoint_counts
}
active_endpoint_counts = (
db.query(ProviderEndpoint.provider_id, func.count(ProviderEndpoint.id))
.filter(
ProviderEndpoint.provider_id.in_(provider_ids),
ProviderEndpoint.is_active.is_(True),
)
.group_by(ProviderEndpoint.provider_id)
.all()
)
active_endpoints_count_map = {
provider_id: int(count) for provider_id, count in active_endpoint_counts
}
result = []
for provider in providers:
models_count = db.query(Model).filter(Model.provider_id == provider.id).count()
active_models_count = (
db.query(Model)
.filter(
and_(
Model.provider_id == provider.id,
Model.is_active.is_(True),
)
)
.count()
)
endpoints_count = len(provider.endpoints) if provider.endpoints else 0
active_endpoints_count = (
sum(1 for ep in provider.endpoints if ep.is_active) if provider.endpoints else 0
)
models_count = models_count_map.get(provider.id, 0)
active_models_count = active_models_count_map.get(provider.id, 0)
endpoints_count = endpoints_count_map.get(provider.id, 0)
active_endpoints_count = active_endpoints_count_map.get(provider.id, 0)
provider_data = PublicProviderResponse(
id=provider.id,
name=provider.name,
@@ -351,6 +404,12 @@ class PublicModelsAdapter(PublicApiAdapter):
skip: int
limit: int
@cache_result(
key_prefix="public:catalog:models",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["provider_id", "is_active", "skip", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求模型列表")
@@ -407,12 +466,19 @@ class PublicModelsAdapter(PublicApiAdapter):
class PublicStatsAdapter(PublicApiAdapter):
@cache_result(
key_prefix="public:catalog:stats",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求系统统计信息")
active_providers = db.query(Provider).filter(Provider.is_active.is_(True)).count()
active_models = (
db.query(Model)
active_providers = int(
db.query(func.count(Provider.id)).filter(Provider.is_active.is_(True)).scalar() or 0
)
active_models = int(
db.query(func.count(Model.id))
.join(Provider)
.filter(
and_(
@@ -420,12 +486,21 @@ class PublicStatsAdapter(PublicApiAdapter):
Provider.is_active.is_(True),
)
)
.count()
.scalar()
or 0
)
formats = (
db.query(Provider.api_format).filter(Provider.is_active.is_(True)).distinct().all()
db.query(ProviderEndpoint.api_format)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
ProviderEndpoint.is_active.is_(True),
Provider.is_active.is_(True),
ProviderEndpoint.api_format.isnot(None),
)
.distinct()
.all()
)
supported_formats = [f.api_format for f in formats if f.api_format]
supported_formats = [row[0] for row in formats if row[0]]
stats = ProviderStatsResponse(
total_providers=active_providers,
active_providers=active_providers,
@@ -443,6 +518,12 @@ class PublicSearchModelsAdapter(PublicApiAdapter):
provider_id: int | None
limit: int
@cache_result(
key_prefix="public:catalog:search_models",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["query", "provider_id", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug(f"公共API搜索模型: {self.query}")
@@ -512,6 +593,12 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
lookback_hours: int
per_format_limit: int
@cache_result(
key_prefix="public:catalog:health_api_formats",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=["lookback_hours", "per_format_limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
now = datetime.now(timezone.utc)
@@ -657,7 +744,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
)
logger.debug(f"公开健康监控: 返回 {len(monitors)} 个 API 格式的健康数据")
return response
return response.model_dump()
@dataclass
@@ -669,6 +756,12 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
is_active: bool | None
search: str | None
@cache_result(
key_prefix="public:catalog:global_models",
ttl=CacheTTL.MODEL,
user_specific=False,
vary_by=["skip", "limit", "is_active", "search"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求 GlobalModel 列表")
@@ -691,8 +784,8 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
)
)
# 统计总数
total = query.count()
# 统计总数(避免 Query.count() 生成大子查询)
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()
@@ -714,4 +807,4 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
)
logger.debug(f"返回 {len(model_responses)} 个 GlobalModel")
return PublicGlobalModelListResponse(models=model_responses, total=total)
return PublicGlobalModelListResponse(models=model_responses, total=total).model_dump()

View File

@@ -12,15 +12,18 @@ from typing import Any
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session, selectinload
from sqlalchemy.orm import Session, load_only, selectinload
from src.api.handlers.base.request_builder import PassthroughRequestBuilder
from src.api.handlers.base.request_builder import build_test_request_body, get_provider_auth
from src.api.handlers.base.request_builder import (
PassthroughRequestBuilder,
build_test_request_body,
get_provider_auth,
)
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
from src.database import get_db
from src.database.database import get_pool_status
from src.models.database import Model, Provider, ProviderAPIKey, ProviderEndpoint
from src.models.database import GlobalModel, Model, Provider, ProviderAPIKey, ProviderEndpoint
from src.services.provider.transport import build_provider_url
from src.utils.ssl_utils import get_ssl_context
@@ -85,7 +88,7 @@ def _serialize_provider(
def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
"""选择 Provider按 provider_priority 优先级选择)"""
query = db.query(Provider).filter(Provider.is_active == True)
query = db.query(Provider).filter(Provider.is_active.is_(True))
if provider_name:
provider = query.filter(Provider.name == provider_name).first()
if provider:
@@ -102,9 +105,9 @@ def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
async def service_health(db: Session = Depends(get_db)) -> Any:
"""返回服务健康状态与依赖信息"""
active_providers = (
db.query(func.count(Provider.id)).filter(Provider.is_active == True).scalar() or 0
db.query(func.count(Provider.id)).filter(Provider.is_active.is_(True)).scalar() or 0
)
active_models = db.query(func.count(Model.id)).filter(Model.is_active == True).scalar() or 0
active_models = db.query(func.count(Model.id)).filter(Model.is_active.is_(True)).scalar() or 0
redis_info: dict[str, Any] = {"status": "unknown"}
try:
@@ -163,11 +166,14 @@ async def root(db: Session = Depends(get_db)) -> Any:
# 按优先级选择最高优先级的提供商
top_provider = (
db.query(Provider)
.filter(Provider.is_active == True)
.options(load_only(Provider.id, Provider.name, Provider.provider_priority))
.filter(Provider.is_active.is_(True))
.order_by(Provider.provider_priority.asc())
.first()
)
active_providers = db.query(Provider).filter(Provider.is_active == True).count()
active_providers = (
db.query(func.count(Provider.id)).filter(Provider.is_active.is_(True)).scalar() or 0
)
return {
"message": "AI Proxy with Modular Architecture v4.0.0",
@@ -193,17 +199,37 @@ async def list_providers(
active_only: bool = Query(True),
) -> Any:
"""列出所有 Provider"""
load_options = []
load_options = [
load_only(Provider.id, Provider.name, Provider.is_active, Provider.provider_priority)
]
if include_models:
load_options.append(selectinload(Provider.models).selectinload(Model.global_model))
load_options.append(
selectinload(Provider.models)
.load_only(
Model.id,
Model.provider_model_name,
Model.is_active,
Model.supports_streaming,
Model.global_model_id,
)
.selectinload(Model.global_model)
.load_only(GlobalModel.id, GlobalModel.name, GlobalModel.display_name)
)
if include_endpoints:
load_options.append(selectinload(Provider.endpoints))
load_options.append(
selectinload(Provider.endpoints).load_only(
ProviderEndpoint.id,
ProviderEndpoint.base_url,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
base_query = db.query(Provider)
if load_options:
base_query = base_query.options(*load_options)
if active_only:
base_query = base_query.filter(Provider.is_active == True)
base_query = base_query.filter(Provider.is_active.is_(True))
base_query = base_query.order_by(Provider.provider_priority.asc(), Provider.name.asc())
providers = base_query.all()
@@ -223,11 +249,31 @@ async def provider_detail(
include_endpoints: bool = Query(False),
) -> Any:
"""获取单个 Provider 详情"""
load_options = []
load_options = [
load_only(Provider.id, Provider.name, Provider.is_active, Provider.provider_priority)
]
if include_models:
load_options.append(selectinload(Provider.models).selectinload(Model.global_model))
load_options.append(
selectinload(Provider.models)
.load_only(
Model.id,
Model.provider_model_name,
Model.is_active,
Model.supports_streaming,
Model.global_model_id,
)
.selectinload(Model.global_model)
.load_only(GlobalModel.id, GlobalModel.name, GlobalModel.display_name)
)
if include_endpoints:
load_options.append(selectinload(Provider.endpoints))
load_options.append(
selectinload(Provider.endpoints).load_only(
ProviderEndpoint.id,
ProviderEndpoint.base_url,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
base_query = db.query(Provider)
if load_options:

View File

@@ -14,6 +14,7 @@ from sqlalchemy.orm import Session
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.core.crypto import crypto_service
from src.core.exceptions import (
ForbiddenException,
@@ -45,6 +46,7 @@ from src.services.system.time_range import TimeRangeParams
from src.services.usage.service import UsageService
from src.services.user.apikey import ApiKeyService
from src.services.user.preference import PreferenceService
from src.utils.cache_decorator import cache_result
router = APIRouter(prefix="/api/users/me", tags=["User Profile"])
pipeline = ApiRequestPipeline()
@@ -259,7 +261,7 @@ async def get_my_active_requests(
async def get_my_interval_timeline(
request: Request,
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
limit: int = Query(5000, ge=100, le=20000, description="最大返回数据点数量"),
limit: int = Query(2000, ge=100, le=20000, description="最大返回数据点数量"),
db: Session = Depends(get_db),
) -> Any:
"""
@@ -767,6 +769,21 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
limit: int = 100
offset: int = 0
@cache_result(
key_prefix="user:usage:records",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=True,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"search",
"limit",
"offset",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import or_
from sqlalchemy.orm import load_only
@@ -954,17 +971,18 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
)
avg_resp_query = db.query(func.avg(Usage.response_time_ms)).filter(
Usage.user_id == user.id,
Usage.status_code == 200,
Usage.response_time_ms.isnot(None),
# 复用 summary 聚合中的成功请求响应时间,避免额外 AVG SQL
total_success_response_time_ms = sum(
float(item.get("success_response_time_sum_ms", 0.0) or 0.0) for item in summary_list
)
total_success_response_count = sum(
int(item.get("success_response_time_count", 0) or 0) for item in summary_list
)
avg_response_time = (
total_success_response_time_ms / total_success_response_count / 1000.0
if total_success_response_count > 0
else 0.0
)
if start_utc and end_utc:
avg_resp_query = avg_resp_query.filter(
Usage.created_at >= start_utc, Usage.created_at < end_utc
)
avg_response_ms = avg_resp_query.scalar() or 0
avg_response_time = float(avg_response_ms) / 1000.0 if avg_response_ms else 0
# 构建响应数据
response_data = {
@@ -1100,6 +1118,12 @@ class GetMyIntervalTimelineAdapter(AuthenticatedApiAdapter):
hours: int
limit: int
@cache_result(
key_prefix="user:usage:interval_timeline",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=True,
vary_by=["hours", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user