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,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: