mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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 ==========
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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": []}
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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 字段
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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("不能删除最后一个管理员账户")
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user