mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: 添加多维度计费系统和视频任务管理功能
计费系统: - 新增 BillingRule 和 DimensionCollector 数据模型 - 实现 FormulaEngine 安全表达式求值引擎 (AST 白名单) - 支持 dimension/matrix/tiered/constant 多种维度映射 - BillingRuleService 支持 Provider Model -> GlobalModel 规则回退 - CLI task_type 在计费域自动映射为 chat 视频任务增强: - 添加 request_metadata 字段记录候选 key 和计费规则快照 - 后台轮询支持并发控制 (Semaphore + 独立 session) - 任务终态自动写入 Usage 记录并计算成本 - 新增视频任务管理 API 和前端界面 其他改进: - UsageService 新增 record_usage_with_custom_cost 方法 - StandardizedUsage 支持 dimensions 字段 (兼容 extra) - 配置新增 BILLING_REQUIRE_RULE 和 BILLING_STRICT_MODE
This commit is contained in:
@@ -4,10 +4,11 @@ from fastapi import APIRouter
|
||||
|
||||
from .adaptive import router as adaptive_router
|
||||
from .api_keys import router as api_keys_router
|
||||
from .billing import router as billing_router
|
||||
from .endpoints import router as endpoints_router
|
||||
from .management_tokens import router as management_tokens_router
|
||||
from .modules import router as modules_router
|
||||
from .models import router as models_router
|
||||
from .modules import router as modules_router
|
||||
from .monitoring import router as monitoring_router
|
||||
from .provider_ops import router as provider_ops_router
|
||||
from .provider_query import router as provider_query_router
|
||||
@@ -17,12 +18,14 @@ from .security import router as security_router
|
||||
from .system import router as system_router
|
||||
from .usage import router as usage_router
|
||||
from .users import router as users_router
|
||||
from .video_tasks import router as video_tasks_router
|
||||
|
||||
router = APIRouter()
|
||||
router.include_router(system_router)
|
||||
router.include_router(users_router)
|
||||
router.include_router(providers_router)
|
||||
router.include_router(api_keys_router)
|
||||
router.include_router(billing_router)
|
||||
router.include_router(usage_router)
|
||||
router.include_router(monitoring_router)
|
||||
router.include_router(endpoints_router)
|
||||
@@ -34,6 +37,7 @@ router.include_router(provider_query_router)
|
||||
router.include_router(management_tokens_router)
|
||||
router.include_router(modules_router)
|
||||
router.include_router(provider_ops_router)
|
||||
router.include_router(video_tasks_router)
|
||||
|
||||
# 注意:ldap_router 已迁移到模块系统,由 ModuleRegistry 动态注册
|
||||
# 当 LDAP_AVAILABLE=true 时才会注册路由
|
||||
|
||||
5
src/api/admin/billing/__init__.py
Normal file
5
src/api/admin/billing/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
"""Billing 配置管理 API 模块(billing_rules / dimension_collectors)。"""
|
||||
|
||||
from .routes import router
|
||||
|
||||
__all__ = ["router"]
|
||||
527
src/api/admin/billing/routes.py
Normal file
527
src/api/admin/billing/routes.py
Normal file
@@ -0,0 +1,527 @@
|
||||
"""Billing 配置管理 API 路由。
|
||||
|
||||
包含:
|
||||
- billing_rules: 计费规则(公式/变量/维度映射)
|
||||
- dimension_collectors: 维度采集器(request/response/metadata/computed)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
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.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
|
||||
|
||||
router = APIRouter(prefix="/api/admin/billing", tags=["Admin - Billing"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
_expr_validator = SafeExpressionEvaluator()
|
||||
|
||||
|
||||
AllowedTaskType = Literal["chat", "video", "image", "audio"]
|
||||
AllowedCollectorSourceType = Literal["request", "response", "metadata", "computed"]
|
||||
AllowedValueType = Literal["float", "int", "string"]
|
||||
|
||||
|
||||
class BillingRuleUpsertRequest(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=100)
|
||||
task_type: AllowedTaskType = "chat"
|
||||
|
||||
global_model_id: str | None = None
|
||||
model_id: str | None = None
|
||||
|
||||
expression: str = Field(..., min_length=1)
|
||||
variables: dict[str, Any] = Field(default_factory=dict)
|
||||
dimension_mappings: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
is_enabled: bool = True
|
||||
|
||||
|
||||
class BillingRuleResponse(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
task_type: str
|
||||
global_model_id: str | None
|
||||
model_id: str | None
|
||||
expression: str
|
||||
variables: dict[str, Any]
|
||||
dimension_mappings: dict[str, Any]
|
||||
is_enabled: bool
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
@classmethod
|
||||
def from_orm_obj(cls, rule: BillingRule) -> "BillingRuleResponse":
|
||||
return cls(
|
||||
id=rule.id,
|
||||
name=rule.name,
|
||||
task_type=rule.task_type,
|
||||
global_model_id=rule.global_model_id,
|
||||
model_id=rule.model_id,
|
||||
expression=rule.expression,
|
||||
variables=rule.variables or {},
|
||||
dimension_mappings=rule.dimension_mappings or {},
|
||||
is_enabled=bool(rule.is_enabled),
|
||||
created_at=rule.created_at,
|
||||
updated_at=rule.updated_at,
|
||||
)
|
||||
|
||||
|
||||
class DimensionCollectorUpsertRequest(BaseModel):
|
||||
api_format: str = Field(..., min_length=1, max_length=50)
|
||||
task_type: str = Field(..., min_length=1, max_length=20)
|
||||
dimension_name: str = Field(..., min_length=1, max_length=100)
|
||||
|
||||
source_type: AllowedCollectorSourceType
|
||||
source_path: str | None = None
|
||||
value_type: AllowedValueType = "float"
|
||||
transform_expression: str | None = None
|
||||
default_value: str | None = None
|
||||
|
||||
priority: int = 0
|
||||
is_enabled: bool = True
|
||||
|
||||
|
||||
class DimensionCollectorResponse(BaseModel):
|
||||
id: str
|
||||
api_format: str
|
||||
task_type: str
|
||||
dimension_name: str
|
||||
source_type: str
|
||||
source_path: str | None
|
||||
value_type: str
|
||||
transform_expression: str | None
|
||||
default_value: str | None
|
||||
priority: int
|
||||
is_enabled: bool
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
@classmethod
|
||||
def from_orm_obj(cls, c: DimensionCollector) -> "DimensionCollectorResponse":
|
||||
return cls(
|
||||
id=c.id,
|
||||
api_format=c.api_format,
|
||||
task_type=c.task_type,
|
||||
dimension_name=c.dimension_name,
|
||||
source_type=c.source_type,
|
||||
source_path=c.source_path,
|
||||
value_type=c.value_type,
|
||||
transform_expression=c.transform_expression,
|
||||
default_value=c.default_value,
|
||||
priority=int(c.priority or 0),
|
||||
is_enabled=bool(c.is_enabled),
|
||||
created_at=c.created_at,
|
||||
updated_at=c.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/rules")
|
||||
async def list_billing_rules(
|
||||
request: Request,
|
||||
task_type: str | None = Query(None),
|
||||
is_enabled: bool | None = Query(None),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(50, ge=1, le=200),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
adapter = BillingRuleListAdapter(
|
||||
task_type=task_type,
|
||||
is_enabled=is_enabled,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/rules/{rule_id}")
|
||||
async def get_billing_rule(rule_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = BillingRuleDetailAdapter(rule_id=rule_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/rules")
|
||||
async def create_billing_rule(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = BillingRuleCreateAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.put("/rules/{rule_id}")
|
||||
async def update_billing_rule(rule_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = BillingRuleUpdateAdapter(rule_id=rule_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/collectors")
|
||||
async def list_dimension_collectors(
|
||||
request: Request,
|
||||
api_format: str | None = Query(None),
|
||||
task_type: str | None = Query(None),
|
||||
dimension_name: str | None = Query(None),
|
||||
is_enabled: bool | None = Query(None),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(50, ge=1, le=200),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
adapter = DimensionCollectorListAdapter(
|
||||
api_format=api_format,
|
||||
task_type=task_type,
|
||||
dimension_name=dimension_name,
|
||||
is_enabled=is_enabled,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/collectors/{collector_id}")
|
||||
async def get_dimension_collector(
|
||||
collector_id: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> Any:
|
||||
adapter = DimensionCollectorDetailAdapter(collector_id=collector_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/collectors")
|
||||
async def create_dimension_collector(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = DimensionCollectorCreateAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.put("/collectors/{collector_id}")
|
||||
async def update_dimension_collector(
|
||||
collector_id: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> Any:
|
||||
adapter = DimensionCollectorUpdateAdapter(collector_id=collector_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Adapters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class BillingRuleListAdapter(AdminApiAdapter):
|
||||
page: int
|
||||
page_size: int
|
||||
task_type: str | None = None
|
||||
is_enabled: bool | None = None
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
q = context.db.query(BillingRule)
|
||||
if self.task_type:
|
||||
q = q.filter(BillingRule.task_type == self.task_type.lower())
|
||||
if self.is_enabled is not None:
|
||||
q = q.filter(BillingRule.is_enabled == self.is_enabled)
|
||||
|
||||
total = q.count()
|
||||
items = (
|
||||
q.order_by(BillingRule.updated_at.desc())
|
||||
.offset((self.page - 1) * self.page_size)
|
||||
.limit(self.page_size)
|
||||
.all()
|
||||
)
|
||||
return {
|
||||
"items": [BillingRuleResponse.from_orm_obj(r).model_dump() for r in items],
|
||||
"total": total,
|
||||
"page": self.page,
|
||||
"page_size": self.page_size,
|
||||
"pages": (total + self.page_size - 1) // self.page_size,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class BillingRuleDetailAdapter(AdminApiAdapter):
|
||||
rule_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
rule = context.db.query(BillingRule).filter(BillingRule.id == self.rule_id).first()
|
||||
if not rule:
|
||||
raise NotFoundException("Billing rule not found")
|
||||
return BillingRuleResponse.from_orm_obj(rule).model_dump()
|
||||
|
||||
|
||||
class BillingRuleCreateAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = BillingRuleUpsertRequest.model_validate(payload)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(f"Invalid request body: {exc}")
|
||||
|
||||
_validate_billing_rule_request(req)
|
||||
|
||||
rule = BillingRule(
|
||||
name=req.name,
|
||||
task_type=req.task_type,
|
||||
global_model_id=req.global_model_id,
|
||||
model_id=req.model_id,
|
||||
expression=req.expression,
|
||||
variables=req.variables,
|
||||
dimension_mappings=req.dimension_mappings,
|
||||
is_enabled=req.is_enabled,
|
||||
)
|
||||
context.db.add(rule)
|
||||
try:
|
||||
context.db.commit()
|
||||
except IntegrityError as exc:
|
||||
context.db.rollback()
|
||||
raise InvalidRequestException(f"Integrity error: {exc}")
|
||||
|
||||
context.db.refresh(rule)
|
||||
return BillingRuleResponse.from_orm_obj(rule).model_dump()
|
||||
|
||||
|
||||
@dataclass
|
||||
class BillingRuleUpdateAdapter(AdminApiAdapter):
|
||||
rule_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
rule = context.db.query(BillingRule).filter(BillingRule.id == self.rule_id).first()
|
||||
if not rule:
|
||||
raise NotFoundException("Billing rule not found")
|
||||
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = BillingRuleUpsertRequest.model_validate(payload)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(f"Invalid request body: {exc}")
|
||||
|
||||
_validate_billing_rule_request(req)
|
||||
|
||||
rule.name = req.name
|
||||
rule.task_type = req.task_type
|
||||
rule.global_model_id = req.global_model_id
|
||||
rule.model_id = req.model_id
|
||||
rule.expression = req.expression
|
||||
rule.variables = req.variables
|
||||
rule.dimension_mappings = req.dimension_mappings
|
||||
rule.is_enabled = req.is_enabled
|
||||
|
||||
try:
|
||||
context.db.commit()
|
||||
except IntegrityError as exc:
|
||||
context.db.rollback()
|
||||
raise InvalidRequestException(f"Integrity error: {exc}")
|
||||
|
||||
context.db.refresh(rule)
|
||||
return BillingRuleResponse.from_orm_obj(rule).model_dump()
|
||||
|
||||
|
||||
@dataclass
|
||||
class DimensionCollectorListAdapter(AdminApiAdapter):
|
||||
page: int
|
||||
page_size: int
|
||||
api_format: str | None = None
|
||||
task_type: str | None = None
|
||||
dimension_name: str | None = None
|
||||
is_enabled: bool | None = None
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
q = context.db.query(DimensionCollector)
|
||||
if self.api_format:
|
||||
q = q.filter(DimensionCollector.api_format == self.api_format.upper())
|
||||
if self.task_type:
|
||||
q = q.filter(DimensionCollector.task_type == self.task_type.lower())
|
||||
if self.dimension_name:
|
||||
q = q.filter(DimensionCollector.dimension_name == self.dimension_name)
|
||||
if self.is_enabled is not None:
|
||||
q = q.filter(DimensionCollector.is_enabled == self.is_enabled)
|
||||
|
||||
total = q.count()
|
||||
items = (
|
||||
q.order_by(DimensionCollector.updated_at.desc())
|
||||
.offset((self.page - 1) * self.page_size)
|
||||
.limit(self.page_size)
|
||||
.all()
|
||||
)
|
||||
return {
|
||||
"items": [DimensionCollectorResponse.from_orm_obj(c).model_dump() for c in items],
|
||||
"total": total,
|
||||
"page": self.page,
|
||||
"page_size": self.page_size,
|
||||
"pages": (total + self.page_size - 1) // self.page_size,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class DimensionCollectorDetailAdapter(AdminApiAdapter):
|
||||
collector_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
c = (
|
||||
context.db.query(DimensionCollector)
|
||||
.filter(DimensionCollector.id == self.collector_id)
|
||||
.first()
|
||||
)
|
||||
if not c:
|
||||
raise NotFoundException("Dimension collector not found")
|
||||
return DimensionCollectorResponse.from_orm_obj(c).model_dump()
|
||||
|
||||
|
||||
class DimensionCollectorCreateAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = DimensionCollectorUpsertRequest.model_validate(payload)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(f"Invalid request body: {exc}")
|
||||
|
||||
_validate_dimension_collector_request(context.db, req, existing_id=None)
|
||||
|
||||
c = DimensionCollector(
|
||||
api_format=req.api_format.upper(),
|
||||
task_type=req.task_type.lower(),
|
||||
dimension_name=req.dimension_name,
|
||||
source_type=req.source_type,
|
||||
source_path=req.source_path,
|
||||
value_type=req.value_type,
|
||||
transform_expression=req.transform_expression,
|
||||
default_value=req.default_value,
|
||||
priority=req.priority,
|
||||
is_enabled=req.is_enabled,
|
||||
)
|
||||
context.db.add(c)
|
||||
try:
|
||||
context.db.commit()
|
||||
except IntegrityError as exc:
|
||||
context.db.rollback()
|
||||
raise InvalidRequestException(f"Integrity error: {exc}")
|
||||
|
||||
context.db.refresh(c)
|
||||
return DimensionCollectorResponse.from_orm_obj(c).model_dump()
|
||||
|
||||
|
||||
@dataclass
|
||||
class DimensionCollectorUpdateAdapter(AdminApiAdapter):
|
||||
collector_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
|
||||
c = (
|
||||
context.db.query(DimensionCollector)
|
||||
.filter(DimensionCollector.id == self.collector_id)
|
||||
.first()
|
||||
)
|
||||
if not c:
|
||||
raise NotFoundException("Dimension collector not found")
|
||||
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = DimensionCollectorUpsertRequest.model_validate(payload)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(f"Invalid request body: {exc}")
|
||||
|
||||
_validate_dimension_collector_request(context.db, req, existing_id=self.collector_id)
|
||||
|
||||
c.api_format = req.api_format.upper()
|
||||
c.task_type = req.task_type.lower()
|
||||
c.dimension_name = req.dimension_name
|
||||
c.source_type = req.source_type
|
||||
c.source_path = req.source_path
|
||||
c.value_type = req.value_type
|
||||
c.transform_expression = req.transform_expression
|
||||
c.default_value = req.default_value
|
||||
c.priority = req.priority
|
||||
c.is_enabled = req.is_enabled
|
||||
|
||||
try:
|
||||
context.db.commit()
|
||||
except IntegrityError as exc:
|
||||
context.db.rollback()
|
||||
raise InvalidRequestException(f"Integrity error: {exc}")
|
||||
|
||||
context.db.refresh(c)
|
||||
return DimensionCollectorResponse.from_orm_obj(c).model_dump()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Validation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _validate_billing_rule_request(req: BillingRuleUpsertRequest) -> None:
|
||||
# model/global_model 二选一
|
||||
if bool(req.global_model_id) == bool(req.model_id):
|
||||
raise InvalidRequestException("Exactly one of global_model_id or model_id must be provided")
|
||||
|
||||
# task_type 校验:Pydantic Literal 已限制为 "chat", "video", "image", "audio"
|
||||
# 注:CLI 在计费域等同于 chat,billing_rules 不存储 "cli"
|
||||
|
||||
# expression 安全校验
|
||||
try:
|
||||
_expr_validator.validate(req.expression)
|
||||
except UnsafeExpressionError as exc:
|
||||
raise InvalidRequestException(f"Invalid expression: {exc}")
|
||||
|
||||
# variables 必须为数值(JSON 可包含 int/float)
|
||||
if not isinstance(req.variables, dict):
|
||||
raise InvalidRequestException("variables must be a JSON object")
|
||||
for k, v in req.variables.items():
|
||||
if not isinstance(k, str) or not k:
|
||||
raise InvalidRequestException("variables keys must be non-empty strings")
|
||||
if isinstance(v, bool) or not isinstance(v, (int, float)):
|
||||
raise InvalidRequestException(f"variables['{k}'] must be a number")
|
||||
|
||||
# dimension_mappings 结构做轻量校验(详细 schema 由业务侧保障)
|
||||
if not isinstance(req.dimension_mappings, dict):
|
||||
raise InvalidRequestException("dimension_mappings must be a JSON object")
|
||||
for var_name, mapping in req.dimension_mappings.items():
|
||||
if not isinstance(var_name, str) or not var_name:
|
||||
raise InvalidRequestException("dimension_mappings keys must be non-empty strings")
|
||||
if not isinstance(mapping, dict):
|
||||
raise InvalidRequestException(f"dimension_mappings['{var_name}'] must be an object")
|
||||
if "source" not in mapping:
|
||||
raise InvalidRequestException(f"dimension_mappings['{var_name}'].source is required")
|
||||
|
||||
|
||||
def _validate_dimension_collector_request(
|
||||
db: Session,
|
||||
req: DimensionCollectorUpsertRequest,
|
||||
*,
|
||||
existing_id: str | None,
|
||||
) -> None:
|
||||
src = req.source_type
|
||||
if src == "computed":
|
||||
if req.source_path is not None:
|
||||
raise InvalidRequestException("computed collector must have source_path=null")
|
||||
if not req.transform_expression:
|
||||
raise InvalidRequestException("computed collector must have transform_expression")
|
||||
else:
|
||||
if not req.source_path:
|
||||
raise InvalidRequestException("non-computed collector must have source_path")
|
||||
|
||||
# transform_expression 安全校验(如配置)
|
||||
if req.transform_expression:
|
||||
try:
|
||||
_expr_validator.validate(req.transform_expression)
|
||||
except UnsafeExpressionError as exc:
|
||||
raise InvalidRequestException(f"Invalid transform_expression: {exc}")
|
||||
|
||||
# default_value 仅允许同一维度一条(enabled=true)
|
||||
if req.default_value is not None and req.is_enabled:
|
||||
q = db.query(DimensionCollector).filter(
|
||||
DimensionCollector.api_format == req.api_format.upper(),
|
||||
DimensionCollector.task_type == req.task_type.lower(),
|
||||
DimensionCollector.dimension_name == req.dimension_name,
|
||||
DimensionCollector.is_enabled.is_(True),
|
||||
DimensionCollector.default_value.isnot(None),
|
||||
)
|
||||
if existing_id:
|
||||
q = q.filter(DimensionCollector.id != existing_id)
|
||||
exists = db.query(q.exists()).scalar()
|
||||
if exists:
|
||||
raise InvalidRequestException(
|
||||
"default_value already exists for this (api_format, task_type, dimension_name)"
|
||||
)
|
||||
5
src/api/admin/video_tasks/__init__.py
Normal file
5
src/api/admin/video_tasks/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
"""视频任务管理 API 模块。"""
|
||||
|
||||
from .routes import router
|
||||
|
||||
__all__ = ["router"]
|
||||
424
src/api/admin/video_tasks/routes.py
Normal file
424
src/api/admin/video_tasks/routes.py
Normal file
@@ -0,0 +1,424 @@
|
||||
"""视频任务管理 API 路由。
|
||||
|
||||
管理员可以查看所有视频任务,用户只能查看自己的任务。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
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.core.enums import UserRole
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, ProviderEndpoint, User, VideoTask
|
||||
|
||||
router = APIRouter(prefix="/api/admin/video-tasks", tags=["Admin - Video Tasks"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def list_video_tasks(
|
||||
request: Request,
|
||||
status: str | None = Query(None, description="Filter by status"),
|
||||
user_id: str | None = Query(None, description="Filter by user ID (admin only)"),
|
||||
model: str | None = Query(None, description="Filter by model"),
|
||||
page: int = Query(1, ge=1, description="Page number"),
|
||||
page_size: int = Query(20, ge=1, le=100, description="Items per page"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
获取视频任务列表
|
||||
|
||||
管理员可以查看所有用户的任务,普通用户只能查看自己的任务。
|
||||
|
||||
**查询参数**:
|
||||
- `status`: 按状态筛选(pending/submitted/processing/completed/failed/cancelled)
|
||||
- `user_id`: 按用户 ID 筛选(仅管理员)
|
||||
- `model`: 按模型筛选
|
||||
- `page`: 页码,默认 1
|
||||
- `page_size`: 每页数量,默认 20,最大 100
|
||||
|
||||
**返回字段**:
|
||||
- `items`: 任务列表
|
||||
- `total`: 总数
|
||||
- `page`: 当前页码
|
||||
- `page_size`: 每页数量
|
||||
- `pages`: 总页数
|
||||
"""
|
||||
adapter = VideoTaskListAdapter(
|
||||
status=status,
|
||||
user_id=user_id,
|
||||
model=model,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/stats")
|
||||
async def get_video_task_stats(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
获取视频任务统计
|
||||
|
||||
**返回字段**:
|
||||
- `total`: 总任务数
|
||||
- `by_status`: 按状态分组的数量
|
||||
- `by_model`: 按模型分组的数量(前 10)
|
||||
- `today_count`: 今日任务数
|
||||
"""
|
||||
adapter = VideoTaskStatsAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{task_id}")
|
||||
async def get_video_task_detail(
|
||||
task_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
获取视频任务详情
|
||||
|
||||
**路径参数**:
|
||||
- `task_id`: 任务 ID
|
||||
|
||||
**返回字段**:
|
||||
- 任务的完整信息,包括请求体、响应、状态等
|
||||
"""
|
||||
adapter = VideoTaskDetailAdapter(task_id=task_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{task_id}/cancel")
|
||||
async def cancel_video_task(
|
||||
task_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
取消视频任务
|
||||
|
||||
**路径参数**:
|
||||
- `task_id`: 任务 ID
|
||||
|
||||
**返回**:
|
||||
- 更新后的任务信息
|
||||
"""
|
||||
adapter = VideoTaskCancelAdapter(task_id=task_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
# ==================== Adapters ====================
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoTaskListAdapter(DashboardAdapter):
|
||||
"""视频任务列表适配器"""
|
||||
|
||||
status: str | None
|
||||
user_id: str | None
|
||||
model: str | None
|
||||
page: int
|
||||
page_size: int
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
db = context.db
|
||||
user = context.user
|
||||
is_admin = user.role == UserRole.ADMIN
|
||||
|
||||
query = db.query(VideoTask)
|
||||
|
||||
# 权限过滤:普通用户只能看自己的任务
|
||||
if not is_admin:
|
||||
query = query.filter(VideoTask.user_id == user.id)
|
||||
elif self.user_id:
|
||||
# 管理员可以按用户筛选
|
||||
query = query.filter(VideoTask.user_id == self.user_id)
|
||||
|
||||
# 状态筛选
|
||||
if self.status:
|
||||
query = query.filter(VideoTask.status == self.status)
|
||||
|
||||
# 模型筛选
|
||||
if self.model:
|
||||
escaped = self.model.replace("%", "\\%").replace("_", "\\_")
|
||||
query = query.filter(VideoTask.model.ilike(f"%{escaped}%"))
|
||||
|
||||
# 统计总数
|
||||
total = query.count()
|
||||
|
||||
# 分页
|
||||
offset = (self.page - 1) * self.page_size
|
||||
tasks = (
|
||||
query.order_by(VideoTask.created_at.desc()).offset(offset).limit(self.page_size).all()
|
||||
)
|
||||
|
||||
# 获取用户信息映射
|
||||
user_ids = list(set(t.user_id for t in tasks if t.user_id))
|
||||
users_map = {}
|
||||
if user_ids:
|
||||
users = db.query(User).filter(User.id.in_(user_ids)).all()
|
||||
users_map = {u.id: u.username for u in users}
|
||||
|
||||
# 获取 Provider 信息映射
|
||||
provider_ids = list(set(t.provider_id for t in tasks if t.provider_id))
|
||||
providers_map = {}
|
||||
if provider_ids:
|
||||
providers = db.query(Provider).filter(Provider.id.in_(provider_ids)).all()
|
||||
providers_map = {p.id: p.name for p in providers}
|
||||
|
||||
items = []
|
||||
for task in tasks:
|
||||
items.append(
|
||||
{
|
||||
"id": task.id,
|
||||
"external_task_id": task.external_task_id,
|
||||
"user_id": task.user_id,
|
||||
"username": users_map.get(task.user_id, "Unknown"),
|
||||
"model": task.model,
|
||||
"prompt": (
|
||||
task.prompt[:100] + "..."
|
||||
if task.prompt and len(task.prompt) > 100
|
||||
else task.prompt
|
||||
),
|
||||
"status": task.status,
|
||||
"progress_percent": task.progress_percent,
|
||||
"progress_message": task.progress_message,
|
||||
"provider_id": task.provider_id,
|
||||
"provider_name": providers_map.get(task.provider_id, "Unknown"),
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"video_url": task.video_url,
|
||||
"error_code": task.error_code,
|
||||
"error_message": task.error_message,
|
||||
"poll_count": task.poll_count,
|
||||
"max_poll_count": task.max_poll_count,
|
||||
"created_at": task.created_at.isoformat() if task.created_at else None,
|
||||
"completed_at": task.completed_at.isoformat() if task.completed_at else None,
|
||||
"submitted_at": task.submitted_at.isoformat() if task.submitted_at else None,
|
||||
}
|
||||
)
|
||||
|
||||
pages = (total + self.page_size - 1) // self.page_size
|
||||
|
||||
return {
|
||||
"items": items,
|
||||
"total": total,
|
||||
"page": self.page,
|
||||
"page_size": self.page_size,
|
||||
"pages": pages,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoTaskStatsAdapter(DashboardAdapter):
|
||||
"""视频任务统计适配器"""
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
db = context.db
|
||||
user = context.user
|
||||
is_admin = user.role == UserRole.ADMIN
|
||||
|
||||
base_query = db.query(VideoTask)
|
||||
if not is_admin:
|
||||
base_query = base_query.filter(VideoTask.user_id == user.id)
|
||||
|
||||
# 总数
|
||||
total = base_query.count()
|
||||
|
||||
# 按状态分组
|
||||
status_stats = (
|
||||
base_query.with_entities(
|
||||
VideoTask.status,
|
||||
func.count(VideoTask.id).label("count"),
|
||||
)
|
||||
.group_by(VideoTask.status)
|
||||
.all()
|
||||
)
|
||||
by_status = {stat.status: stat.count for stat in status_stats}
|
||||
|
||||
# 按模型分组(前 10)
|
||||
model_stats = (
|
||||
base_query.with_entities(
|
||||
VideoTask.model,
|
||||
func.count(VideoTask.id).label("count"),
|
||||
)
|
||||
.group_by(VideoTask.model)
|
||||
.order_by(func.count(VideoTask.id).desc())
|
||||
.limit(10)
|
||||
.all()
|
||||
)
|
||||
by_model = {stat.model: stat.count for stat in model_stats}
|
||||
|
||||
# 今日任务数
|
||||
today = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
today_count = base_query.filter(VideoTask.created_at >= today).count()
|
||||
|
||||
# 管理员额外统计
|
||||
result = {
|
||||
"total": total,
|
||||
"by_status": by_status,
|
||||
"by_model": by_model,
|
||||
"today_count": today_count,
|
||||
}
|
||||
|
||||
if is_admin:
|
||||
# 活跃用户数(有视频任务的用户)
|
||||
active_users = db.query(func.count(func.distinct(VideoTask.user_id))).scalar() or 0
|
||||
result["active_users"] = active_users
|
||||
|
||||
# 处理中的任务数
|
||||
processing_count = (
|
||||
db.query(func.count(VideoTask.id))
|
||||
.filter(VideoTask.status.in_(["submitted", "queued", "processing"]))
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
result["processing_count"] = processing_count
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoTaskDetailAdapter(DashboardAdapter):
|
||||
"""视频任务详情适配器"""
|
||||
|
||||
task_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
db = context.db
|
||||
user = context.user
|
||||
is_admin = user.role == UserRole.ADMIN
|
||||
|
||||
query = db.query(VideoTask).filter(VideoTask.id == self.task_id)
|
||||
if not is_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")
|
||||
|
||||
# 获取用户信息
|
||||
task_user = db.query(User).filter(User.id == task.user_id).first()
|
||||
username = task_user.username if task_user else "Unknown"
|
||||
|
||||
# 获取 Provider 信息
|
||||
provider = db.query(Provider).filter(Provider.id == task.provider_id).first()
|
||||
provider_name = provider.name if provider else "Unknown"
|
||||
|
||||
# 获取 Endpoint 信息
|
||||
endpoint = (
|
||||
db.query(ProviderEndpoint).filter(ProviderEndpoint.id == task.endpoint_id).first()
|
||||
)
|
||||
endpoint_info = None
|
||||
if endpoint:
|
||||
endpoint_info = {
|
||||
"id": endpoint.id,
|
||||
"base_url": endpoint.base_url,
|
||||
"api_format": str(endpoint.api_format),
|
||||
}
|
||||
|
||||
return {
|
||||
"id": task.id,
|
||||
"external_task_id": task.external_task_id,
|
||||
"user_id": task.user_id,
|
||||
"username": username,
|
||||
"api_key_id": task.api_key_id,
|
||||
"provider_id": task.provider_id,
|
||||
"provider_name": provider_name,
|
||||
"endpoint_id": task.endpoint_id,
|
||||
"endpoint": endpoint_info,
|
||||
"key_id": task.key_id,
|
||||
"client_api_format": task.client_api_format,
|
||||
"provider_api_format": task.provider_api_format,
|
||||
"format_converted": task.format_converted,
|
||||
"model": task.model,
|
||||
"prompt": task.prompt,
|
||||
"original_request_body": task.original_request_body,
|
||||
"converted_request_body": task.converted_request_body,
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size,
|
||||
"status": task.status,
|
||||
"progress_percent": task.progress_percent,
|
||||
"progress_message": task.progress_message,
|
||||
"video_url": task.video_url,
|
||||
"video_urls": task.video_urls,
|
||||
"thumbnail_url": task.thumbnail_url,
|
||||
"video_size_bytes": task.video_size_bytes,
|
||||
"video_expires_at": (
|
||||
task.video_expires_at.isoformat() if task.video_expires_at else None
|
||||
),
|
||||
"stored_video_path": task.stored_video_path,
|
||||
"storage_provider": task.storage_provider,
|
||||
"error_code": task.error_code,
|
||||
"error_message": task.error_message,
|
||||
"retry_count": task.retry_count,
|
||||
"max_retries": task.max_retries,
|
||||
"poll_interval_seconds": task.poll_interval_seconds,
|
||||
"next_poll_at": task.next_poll_at.isoformat() if task.next_poll_at else None,
|
||||
"poll_count": task.poll_count,
|
||||
"max_poll_count": task.max_poll_count,
|
||||
"created_at": task.created_at.isoformat() if task.created_at else None,
|
||||
"updated_at": task.updated_at.isoformat() if task.updated_at else None,
|
||||
"submitted_at": task.submitted_at.isoformat() if task.submitted_at else None,
|
||||
"completed_at": task.completed_at.isoformat() if task.completed_at else None,
|
||||
"request_metadata": task.request_metadata,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoTaskCancelAdapter(DashboardAdapter):
|
||||
"""视频任务取消适配器"""
|
||||
|
||||
task_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
from src.core.api_format.conversion.internal_video import VideoStatus
|
||||
|
||||
db = context.db
|
||||
user = context.user
|
||||
is_admin = user.role == UserRole.ADMIN
|
||||
|
||||
query = db.query(VideoTask).filter(VideoTask.id == self.task_id)
|
||||
if not is_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 task.status in [
|
||||
VideoStatus.COMPLETED.value,
|
||||
VideoStatus.FAILED.value,
|
||||
VideoStatus.CANCELLED.value,
|
||||
]:
|
||||
raise HTTPException(
|
||||
status_code=400, detail=f"Cannot cancel task with status: {task.status}"
|
||||
)
|
||||
|
||||
# 更新状态
|
||||
task.status = VideoStatus.CANCELLED.value
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
|
||||
return {
|
||||
"id": task.id,
|
||||
"status": task.status,
|
||||
"message": "Task cancelled successfully",
|
||||
}
|
||||
Reference in New Issue
Block a user