mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +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
|
||||
),
|
||||
|
||||
Reference in New Issue
Block a user