mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +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,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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user