mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: 视频计费增强与影子计费系统
This commit is contained in:
@@ -23,6 +23,7 @@ from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.database import BillingRule, DimensionCollector
|
||||
from src.services.billing.formula_engine import SafeExpressionEvaluator, UnsafeExpressionError
|
||||
from src.services.billing.presets import BillingPresetService, PresetApplyMode, list_preset_packs
|
||||
|
||||
router = APIRouter(prefix="/api/admin/billing", tags=["Admin - Billing"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -127,6 +128,30 @@ class DimensionCollectorResponse(BaseModel):
|
||||
)
|
||||
|
||||
|
||||
class BillingPresetInfoResponse(BaseModel):
|
||||
name: str
|
||||
version: str
|
||||
description: str
|
||||
collector_count: int
|
||||
|
||||
|
||||
class ApplyBillingPresetRequest(BaseModel):
|
||||
preset: str = Field(..., min_length=1, max_length=100)
|
||||
mode: PresetApplyMode = "merge"
|
||||
|
||||
|
||||
@router.get("/presets")
|
||||
async def list_billing_presets(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = BillingPresetListAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/presets/apply")
|
||||
async def apply_billing_preset(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = BillingPresetApplyAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/rules")
|
||||
async def list_billing_rules(
|
||||
request: Request,
|
||||
@@ -525,3 +550,38 @@ def _validate_dimension_collector_request(
|
||||
raise InvalidRequestException(
|
||||
"default_value already exists for this (api_format, task_type, dimension_name)"
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BillingPresetListAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
items = []
|
||||
for p in list_preset_packs():
|
||||
items.append(
|
||||
BillingPresetInfoResponse(
|
||||
name=p.name,
|
||||
version=p.version,
|
||||
description=p.description,
|
||||
collector_count=len(p.collectors or []),
|
||||
).model_dump()
|
||||
)
|
||||
return {"items": items}
|
||||
|
||||
|
||||
class BillingPresetApplyAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = ApplyBillingPresetRequest.model_validate(payload)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(f"Invalid request body: {exc}")
|
||||
|
||||
result = BillingPresetService.apply_preset(
|
||||
context.db,
|
||||
preset_name=req.preset,
|
||||
mode=req.mode,
|
||||
)
|
||||
if result.errors:
|
||||
# still return counts; caller can display partial results
|
||||
return {"ok": False, **result.to_dict()}
|
||||
return {"ok": True, **result.to_dict()}
|
||||
|
||||
@@ -726,9 +726,9 @@ class AdminGetApiFormatsAdapter(AdminApiAdapter):
|
||||
def _label_for(sig: str) -> str:
|
||||
fam, kind = (sig.split(":", 1) + [""])[:2]
|
||||
fam_title = {"claude": "Claude", "openai": "OpenAI", "gemini": "Gemini"}.get(fam, fam)
|
||||
if kind == "chat":
|
||||
return fam_title
|
||||
kind_title = {"cli": "CLI", "video": "Video", "image": "Image"}.get(kind, kind)
|
||||
kind_title = {"chat": "Chat", "cli": "CLI", "video": "Video", "image": "Image"}.get(
|
||||
kind, kind
|
||||
)
|
||||
return f"{fam_title} {kind_title}".strip()
|
||||
|
||||
endpoint_defs = list_endpoint_definitions()
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
@@ -14,6 +14,8 @@ 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.config.settings import config
|
||||
from src.database import get_db
|
||||
from src.models.database import (
|
||||
ApiKey,
|
||||
@@ -25,11 +27,31 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.usage.service import UsageService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/admin/usage", tags=["Admin - Usage"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def _apply_admin_default_range(
|
||||
start_date: datetime | None, end_date: datetime | None
|
||||
) -> tuple[datetime | None, datetime | None]:
|
||||
"""
|
||||
Apply a default time range for admin usage endpoints to protect DB from unbounded scans.
|
||||
|
||||
Enabled by setting ADMIN_USAGE_DEFAULT_DAYS>0.
|
||||
"""
|
||||
if start_date is not None or end_date is not None:
|
||||
return start_date, end_date
|
||||
|
||||
days = int(getattr(config, "admin_usage_default_days", 0) or 0)
|
||||
if days <= 0:
|
||||
return start_date, end_date
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
return now - timedelta(days=days), now
|
||||
|
||||
|
||||
# ==================== RESTful Routes ====================
|
||||
|
||||
|
||||
@@ -267,9 +289,14 @@ async def get_usage_detail(
|
||||
|
||||
class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:usage:stats",
|
||||
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
|
||||
user_specific=False,
|
||||
vary_by=["start_date", "end_date"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# Perf: use a single aggregate query (avoid 3 full scans).
|
||||
from sqlalchemy import case
|
||||
@@ -347,10 +374,15 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
|
||||
|
||||
class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
|
||||
self.limit = limit
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:usage:agg:model",
|
||||
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
|
||||
user_specific=False,
|
||||
vary_by=["start_date", "end_date", "limit"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
query = db.query(
|
||||
@@ -394,10 +426,15 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
|
||||
class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
|
||||
self.limit = limit
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:usage:agg:user",
|
||||
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
|
||||
user_specific=False,
|
||||
vary_by=["start_date", "end_date", "limit"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
query = (
|
||||
@@ -444,10 +481,15 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
|
||||
class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
|
||||
self.limit = limit
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:usage:agg:provider",
|
||||
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
|
||||
user_specific=False,
|
||||
vary_by=["start_date", "end_date", "limit"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
@@ -558,10 +600,15 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
|
||||
self.limit = limit
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:usage:agg:api_format",
|
||||
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
|
||||
user_specific=False,
|
||||
vary_by=["start_date", "end_date", "limit"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
query = db.query(
|
||||
@@ -624,8 +671,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
limit: int,
|
||||
offset: int,
|
||||
):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
|
||||
self.search = search
|
||||
self.user_id = user_id
|
||||
self.username = username
|
||||
@@ -635,6 +681,23 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
self.limit = limit
|
||||
self.offset = offset
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:usage:records",
|
||||
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
|
||||
user_specific=False,
|
||||
vary_by=[
|
||||
"start_date",
|
||||
"end_date",
|
||||
"search",
|
||||
"user_id",
|
||||
"username",
|
||||
"model",
|
||||
"provider",
|
||||
"status",
|
||||
"limit",
|
||||
"offset",
|
||||
],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy.orm import load_only
|
||||
@@ -955,7 +1018,11 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
# 先通过主键 id 查找,如果找不到再尝试通过 request_id 查找
|
||||
usage_record = db.query(Usage).filter(Usage.id == self.usage_id).first()
|
||||
if not usage_record:
|
||||
# 兼容通过 request_id 查找(用于异步任务等场景)
|
||||
usage_record = db.query(Usage).filter(Usage.request_id == self.usage_id).first()
|
||||
if not usage_record:
|
||||
raise HTTPException(status_code=404, detail="Usage record not found")
|
||||
|
||||
@@ -970,6 +1037,9 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
usage_id=self.usage_id,
|
||||
)
|
||||
|
||||
# 提取视频/图像/音频计费信息
|
||||
video_billing_info = self._extract_video_billing_info(usage_record)
|
||||
|
||||
return {
|
||||
"id": usage_record.id,
|
||||
"request_id": usage_record.request_id,
|
||||
@@ -1022,6 +1092,7 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
"response_body": usage_record.get_response_body(),
|
||||
"metadata": usage_record.request_metadata,
|
||||
"tiered_pricing": tiered_pricing_info,
|
||||
"video_billing": video_billing_info,
|
||||
}
|
||||
|
||||
async def _get_tiered_pricing_info(self, db: Session, usage_record: Any) -> dict | None:
|
||||
@@ -1077,6 +1148,75 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
"source": pricing_source, # 定价来源: 'provider' 或 'global'
|
||||
}
|
||||
|
||||
def _extract_video_billing_info(self, usage_record: Any) -> dict | None:
|
||||
"""
|
||||
从 request_metadata.billing_snapshot 和 dimensions 中提取视频/图像/音频计费信息。
|
||||
|
||||
返回结构:
|
||||
{
|
||||
"task_type": "video" | "image" | "audio",
|
||||
"duration_seconds": 10.5, # 视频时长(秒)
|
||||
"resolution": "1080p", # 分辨率
|
||||
"video_price_per_second": 0.1, # 每秒单价
|
||||
"video_cost": 1.05, # 视频费用
|
||||
"rule_name": "...", # 计费规则名称
|
||||
"expression": "...", # 计费公式
|
||||
"status": "complete", # 计费状态
|
||||
}
|
||||
"""
|
||||
request_type = getattr(usage_record, "request_type", None)
|
||||
if request_type not in {"video", "image", "audio"}:
|
||||
return None
|
||||
|
||||
metadata = getattr(usage_record, "request_metadata", None)
|
||||
if not metadata:
|
||||
return None
|
||||
|
||||
billing_snapshot = metadata.get("billing_snapshot") if isinstance(metadata, dict) else None
|
||||
dimensions = metadata.get("dimensions") if isinstance(metadata, dict) else None
|
||||
|
||||
result: dict = {
|
||||
"task_type": request_type,
|
||||
}
|
||||
|
||||
# 从 billing_snapshot 中提取计费规则信息
|
||||
if billing_snapshot and isinstance(billing_snapshot, dict):
|
||||
result["rule_name"] = billing_snapshot.get("rule_name")
|
||||
result["expression"] = billing_snapshot.get("expression")
|
||||
result["status"] = billing_snapshot.get("status")
|
||||
result["cost"] = billing_snapshot.get("cost")
|
||||
|
||||
# 从 dimensions_used 中提取维度
|
||||
dims_used = billing_snapshot.get("dimensions_used")
|
||||
if dims_used and isinstance(dims_used, dict):
|
||||
if "duration_seconds" in dims_used:
|
||||
result["duration_seconds"] = dims_used["duration_seconds"]
|
||||
if "video_resolution_key" in dims_used:
|
||||
result["resolution"] = dims_used["video_resolution_key"]
|
||||
if "video_price_per_second" in dims_used:
|
||||
result["video_price_per_second"] = dims_used["video_price_per_second"]
|
||||
if "video_cost" in dims_used:
|
||||
result["video_cost"] = dims_used["video_cost"]
|
||||
|
||||
# 补充从 dimensions 中提取(备用)
|
||||
if dimensions and isinstance(dimensions, dict):
|
||||
if "duration_seconds" not in result and "duration_seconds" in dimensions:
|
||||
result["duration_seconds"] = dimensions["duration_seconds"]
|
||||
if "resolution" not in result and "video_resolution_key" in dimensions:
|
||||
result["resolution"] = dimensions["video_resolution_key"]
|
||||
|
||||
# 如果没有有意义的视频计费信息,返回 None
|
||||
has_video_info = (
|
||||
result.get("duration_seconds")
|
||||
or result.get("resolution")
|
||||
or result.get("video_cost")
|
||||
or result.get("cost")
|
||||
)
|
||||
if not has_video_info:
|
||||
return None
|
||||
|
||||
return result
|
||||
|
||||
|
||||
# ==================== 缓存亲和性分析 ====================
|
||||
|
||||
|
||||
@@ -7,18 +7,22 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
from typing import Any, AsyncIterator
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
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.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, ProviderEndpoint, User, VideoTask
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
|
||||
router = APIRouter(prefix="/api/admin/video-tasks", tags=["Admin - Video Tasks"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -119,6 +123,117 @@ async def cancel_video_task(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{task_id}/video")
|
||||
async def proxy_video_stream(
|
||||
task_id: str,
|
||||
request: Request,
|
||||
token: str | None = Query(None, description="JWT access token"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> StreamingResponse:
|
||||
"""
|
||||
代理视频流(用于需要认证的视频链接)
|
||||
|
||||
**路径参数**:
|
||||
- `task_id`: 任务 ID
|
||||
|
||||
**查询参数**:
|
||||
- `token`: JWT access token(用于 video 标签请求)
|
||||
|
||||
**返回**:
|
||||
- 视频流
|
||||
"""
|
||||
from src.services.auth.service import AuthService
|
||||
|
||||
# 尝试从多个来源获取 token:query param > cookie > header
|
||||
auth_token = token
|
||||
if not auth_token:
|
||||
auth_token = request.cookies.get("access_token")
|
||||
if not auth_token:
|
||||
auth_header = request.headers.get("Authorization")
|
||||
if auth_header and auth_header.startswith("Bearer "):
|
||||
auth_token = auth_header[7:]
|
||||
|
||||
if not auth_token:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
|
||||
try:
|
||||
# 验证 token 并获取 payload
|
||||
payload = await AuthService.verify_token(auth_token, token_type="access")
|
||||
user_id = payload.get("user_id") or payload.get("sub")
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
# 查询用户
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
|
||||
# 查询任务
|
||||
query = db.query(VideoTask).filter(VideoTask.id == task_id)
|
||||
if user.role != UserRole.ADMIN:
|
||||
query = query.filter(VideoTask.user_id == user.id)
|
||||
|
||||
task = query.first()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="Video task not found")
|
||||
|
||||
if not task.video_url:
|
||||
raise HTTPException(status_code=404, detail="Video not available")
|
||||
|
||||
# 检查是否需要代理(Google API 链接需要认证)
|
||||
video_url = task.video_url
|
||||
needs_proxy = "generativelanguage.googleapis.com" in video_url
|
||||
|
||||
if not needs_proxy:
|
||||
# 不需要代理,重定向到原始 URL
|
||||
from fastapi.responses import RedirectResponse
|
||||
|
||||
return RedirectResponse(url=video_url)
|
||||
|
||||
# 需要代理:获取 provider key 进行认证
|
||||
if not task.key_id:
|
||||
raise HTTPException(status_code=500, detail="Missing provider key")
|
||||
|
||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
|
||||
if not key or not key.api_key:
|
||||
raise HTTPException(status_code=500, detail="Provider key not found")
|
||||
|
||||
try:
|
||||
api_key = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
|
||||
|
||||
# 构建认证头
|
||||
headers = {"x-goog-api-key": api_key}
|
||||
|
||||
async def stream_video() -> AsyncIterator[bytes]:
|
||||
"""流式下载并返回视频"""
|
||||
try:
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
async with client.stream("GET", video_url, headers=headers) as response:
|
||||
if response.status_code >= 400:
|
||||
logger.warning(
|
||||
"Video proxy failed: task={} status={}", task_id, response.status_code
|
||||
)
|
||||
return
|
||||
async for chunk in response.aiter_bytes(chunk_size=65536):
|
||||
yield chunk
|
||||
except Exception as e:
|
||||
logger.exception("Video proxy error: task={} error={}", task_id, str(e))
|
||||
|
||||
return StreamingResponse(
|
||||
stream_video(),
|
||||
media_type="video/mp4",
|
||||
headers={
|
||||
"Content-Disposition": f'inline; filename="video_{task_id}.mp4"',
|
||||
"Cache-Control": "private, max-age=3600",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ==================== Adapters ====================
|
||||
|
||||
|
||||
@@ -359,6 +474,7 @@ class VideoTaskDetailAdapter(DashboardAdapter):
|
||||
"video_urls": task.video_urls,
|
||||
"thumbnail_url": task.thumbnail_url,
|
||||
"video_size_bytes": task.video_size_bytes,
|
||||
"video_duration_seconds": task.video_duration_seconds,
|
||||
"video_expires_at": (
|
||||
task.video_expires_at.isoformat() if task.video_expires_at else None
|
||||
),
|
||||
|
||||
@@ -211,8 +211,10 @@ class AdminDashboardStatsAdapter(AdminApiAdapter):
|
||||
month_start = month_start_local.astimezone(timezone.utc)
|
||||
|
||||
# ==================== 使用预聚合数据 ====================
|
||||
# 今日实时数据只查询一次,避免重复扫描 Usage 表
|
||||
today_stats = StatsAggregatorService.get_today_realtime_stats(db)
|
||||
# 从 stats_summary + 今日实时数据获取全局统计
|
||||
combined_stats = StatsAggregatorService.get_combined_stats(db)
|
||||
combined_stats = StatsAggregatorService.get_combined_stats(db, today_stats=today_stats)
|
||||
|
||||
all_time_requests = combined_stats["total_requests"]
|
||||
all_time_success_requests = combined_stats["success_requests"]
|
||||
@@ -237,7 +239,6 @@ class AdminDashboardStatsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# ==================== 今日实时统计 ====================
|
||||
today_stats = StatsAggregatorService.get_today_realtime_stats(db)
|
||||
requests_today = today_stats["total_requests"]
|
||||
cost_today = today_stats["total_cost"]
|
||||
actual_cost_today = today_stats["actual_total_cost"]
|
||||
@@ -951,26 +952,9 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
|
||||
today_stats = StatsAggregatorService.get_today_realtime_stats(db)
|
||||
today_str = today_local.date().isoformat()
|
||||
if today_stats["total_requests"] > 0:
|
||||
# 今日平均响应时间需要单独查询
|
||||
today_avg_rt = (
|
||||
db.query(func.avg(Usage.response_time_ms))
|
||||
.filter(Usage.created_at >= today, Usage.response_time_ms.isnot(None))
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
# 今日 unique_models 和 unique_providers
|
||||
today_unique_models = (
|
||||
db.query(func.count(func.distinct(Usage.model)))
|
||||
.filter(Usage.created_at >= today)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
today_unique_providers = (
|
||||
db.query(func.count(func.distinct(Usage.provider_name)))
|
||||
.filter(Usage.created_at >= today)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
today_avg_rt_ms = float(today_stats.get("avg_response_time_ms") or 0.0)
|
||||
today_unique_models = int(today_stats.get("unique_models") or 0)
|
||||
today_unique_providers = int(today_stats.get("unique_providers") or 0)
|
||||
# 今日 fallback_count
|
||||
today_fallback_count = (
|
||||
db.query(func.count())
|
||||
@@ -996,7 +980,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
|
||||
+ today_stats["cache_read_tokens"]
|
||||
),
|
||||
"cost": today_stats["total_cost"],
|
||||
"avg_response_time": float(today_avg_rt) / 1000.0 if today_avg_rt else 0,
|
||||
"avg_response_time": today_avg_rt_ms / 1000.0 if today_avg_rt_ms else 0,
|
||||
"unique_models": today_unique_models,
|
||||
"unique_providers": today_unique_providers,
|
||||
"fallback_count": today_fallback_count,
|
||||
|
||||
@@ -100,10 +100,7 @@ class StreamTelemetryRecorder:
|
||||
return
|
||||
actual_request_body = ctx.provider_request_body or original_request_body
|
||||
response_body = None
|
||||
if (
|
||||
not isinstance(writer, QueueTelemetryWriter)
|
||||
or config.usage_queue_include_bodies
|
||||
):
|
||||
if not isinstance(writer, QueueTelemetryWriter) or writer.include_bodies:
|
||||
response_body = ctx.build_response_body(response_time_ms)
|
||||
|
||||
try:
|
||||
@@ -403,10 +400,25 @@ class StreamTelemetryRecorder:
|
||||
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
|
||||
) -> TelemetryWriter | None:
|
||||
if config.usage_queue_enabled and self.user_id and self.api_key_id:
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
# Queue payload detail follows system config request_record_level.
|
||||
log_level = SystemConfigService.get_request_record_level(bg_db).value
|
||||
sensitive_headers = SystemConfigService.get_sensitive_headers(bg_db) or []
|
||||
max_request_body_size = int(
|
||||
SystemConfigService.get_config(bg_db, "max_request_body_size", 5242880) or 0
|
||||
)
|
||||
max_response_body_size = int(
|
||||
SystemConfigService.get_config(bg_db, "max_response_body_size", 5242880) or 0
|
||||
)
|
||||
return QueueTelemetryWriter(
|
||||
request_id=self.request_id,
|
||||
user_id=self.user_id,
|
||||
api_key_id=self.api_key_id,
|
||||
log_level=log_level,
|
||||
sensitive_headers=sensitive_headers,
|
||||
max_request_body_size=max_request_body_size,
|
||||
max_response_body_size=max_response_body_size,
|
||||
)
|
||||
db_writer = self._build_db_writer(bg_db)
|
||||
if db_writer is None:
|
||||
|
||||
@@ -75,6 +75,17 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
)
|
||||
self._normalizer = GeminiNormalizer()
|
||||
|
||||
@staticmethod
|
||||
def _get_request_base_url(http_request: Request) -> str:
|
||||
"""从 HTTP 请求中获取基础 URL(协议 + 主机)"""
|
||||
# 优先使用 X-Forwarded-Proto 和 X-Forwarded-Host(代理场景)
|
||||
proto = http_request.headers.get("x-forwarded-proto") or http_request.url.scheme
|
||||
host = http_request.headers.get("x-forwarded-host") or http_request.headers.get("host")
|
||||
if host:
|
||||
return f"{proto}://{host}"
|
||||
# 回退到 request.url
|
||||
return f"{http_request.url.scheme}://{http_request.url.netloc}"
|
||||
|
||||
async def handle_create_task(
|
||||
self,
|
||||
*,
|
||||
@@ -259,7 +270,8 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
created_at=task.created_at,
|
||||
original_request=internal_request,
|
||||
)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task)
|
||||
base_url = self._get_request_base_url(http_request)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
|
||||
|
||||
# 提交成功后立即结算 Usage(费用暂时为 0,轮询完成后更新)
|
||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||
@@ -313,7 +325,8 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
|
||||
# 直接从数据库返回任务状态(后台轮询服务会持续更新状态)
|
||||
internal_task = self._task_to_internal(task)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task)
|
||||
base_url = self._get_request_base_url(http_request)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
|
||||
return JSONResponse(response_body)
|
||||
|
||||
async def handle_list_tasks(
|
||||
@@ -331,8 +344,10 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
.limit(100)
|
||||
.all()
|
||||
)
|
||||
base_url = self._get_request_base_url(http_request)
|
||||
items = [
|
||||
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
|
||||
self._normalizer.video_task_from_internal(self._task_to_internal(t), base_url=base_url)
|
||||
for t in tasks
|
||||
]
|
||||
return JSONResponse({"operations": items})
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, AsyncIterator
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
@@ -650,8 +651,10 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
try:
|
||||
# 使用 httpx 的 stream 方法并正确管理上下文
|
||||
# 视频下载可能较大,设置 5 分钟超时
|
||||
request = client.build_request("GET", upstream_url, headers=headers)
|
||||
response = await client.send(request, stream=True, timeout=300.0)
|
||||
request = client.build_request(
|
||||
"GET", upstream_url, headers=headers, timeout=httpx.Timeout(300.0)
|
||||
)
|
||||
response = await client.send(request, stream=True)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Upstream connection failed task={} url={}: {}",
|
||||
@@ -723,8 +726,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
"""代理直接的视频 URL(如 CDN URL),保持与官方 API 一致的流式返回行为"""
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
try:
|
||||
request = client.build_request("GET", url)
|
||||
response = await client.send(request, stream=True, timeout=300.0)
|
||||
request = client.build_request("GET", url, timeout=httpx.Timeout(300.0))
|
||||
response = await client.send(request, stream=True)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Direct URL connection failed task={} url={}: {}",
|
||||
|
||||
@@ -73,11 +73,11 @@ def _get_formats_for_api(api_format: str) -> list[str]:
|
||||
return _OPENAI_FORMATS
|
||||
|
||||
|
||||
def _is_format_conversion_enabled() -> bool:
|
||||
"""检查全局格式转换开关(从环境变量读取,默认开启)"""
|
||||
from src.config.settings import config
|
||||
def _is_format_conversion_enabled(db: Session) -> bool:
|
||||
"""检查全局格式转换开关(从数据库配置读取,默认开启)"""
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
return config.format_conversion_enabled
|
||||
return SystemConfigService.is_format_conversion_enabled(db)
|
||||
|
||||
|
||||
def _get_convertible_formats(client_format: str, global_conversion_enabled: bool) -> list[str]:
|
||||
@@ -500,7 +500,7 @@ async def list_models(
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
|
||||
|
||||
# 获取可用格式(包括可转换的格式)
|
||||
global_conversion_enabled = _is_format_conversion_enabled()
|
||||
global_conversion_enabled = _is_format_conversion_enabled(db)
|
||||
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
|
||||
candidate_formats, empty_response = _filter_formats_by_restrictions(
|
||||
candidate_formats, restrictions, api_format
|
||||
@@ -604,7 +604,7 @@ async def retrieve_model(
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
|
||||
|
||||
# 获取可用格式(包括可转换的格式)
|
||||
global_conversion_enabled = _is_format_conversion_enabled()
|
||||
global_conversion_enabled = _is_format_conversion_enabled(db)
|
||||
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
|
||||
candidate_formats, _ = _filter_formats_by_restrictions(
|
||||
candidate_formats, restrictions, api_format
|
||||
@@ -687,7 +687,7 @@ async def list_models_gemini(
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
|
||||
|
||||
# 获取可用格式(包括可转换的格式)
|
||||
global_conversion_enabled = _is_format_conversion_enabled()
|
||||
global_conversion_enabled = _is_format_conversion_enabled(db)
|
||||
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
|
||||
candidate_formats, empty_response = _filter_formats_by_restrictions(
|
||||
candidate_formats, restrictions, api_format
|
||||
@@ -766,7 +766,7 @@ async def get_model_gemini(
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
|
||||
|
||||
# 获取可用格式(包括可转换的格式)
|
||||
global_conversion_enabled = _is_format_conversion_enabled()
|
||||
global_conversion_enabled = _is_format_conversion_enabled(db)
|
||||
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
|
||||
candidate_formats, _ = _filter_formats_by_restrictions(
|
||||
candidate_formats, restrictions, api_format
|
||||
|
||||
@@ -1110,7 +1110,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
||||
from sqlalchemy import or_
|
||||
|
||||
from src.api.base.models_service import AccessRestrictions
|
||||
from src.config.settings import config as app_config
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
db = context.db
|
||||
user = context.user
|
||||
@@ -1118,8 +1118,8 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
||||
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)
|
||||
|
||||
# 检查全局格式转换开关
|
||||
global_conversion_enabled = app_config.format_conversion_enabled
|
||||
# 检查全局格式转换开关(从数据库配置读取)
|
||||
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
|
||||
|
||||
# 获取所有可用的 Provider ID(考虑格式转换)
|
||||
available_provider_ids = self._get_all_available_provider_ids(db, global_conversion_enabled)
|
||||
|
||||
Reference in New Issue
Block a user