mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +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",
|
||||
}
|
||||
@@ -16,6 +16,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import httpx
|
||||
@@ -45,6 +46,24 @@ def sanitize_error_message(message: str, max_length: int = 200) -> str:
|
||||
return sanitized[:max_length]
|
||||
|
||||
|
||||
def normalize_gemini_operation_id(operation_id: str) -> str:
|
||||
"""
|
||||
规范化 Gemini operation ID,确保以 "operations/" 开头
|
||||
|
||||
Gemini API 返回的任务 ID 格式可能是 "operations/xxx" 或 "xxx",
|
||||
此函数统一规范化为 "operations/xxx" 格式。
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID
|
||||
|
||||
Returns:
|
||||
规范化后的 operation ID
|
||||
"""
|
||||
if not operation_id.startswith("operations/"):
|
||||
return f"operations/{operation_id}"
|
||||
return operation_id
|
||||
|
||||
|
||||
class VideoHandlerBase(ABC):
|
||||
"""视频处理器基类"""
|
||||
|
||||
@@ -204,5 +223,28 @@ class VideoHandlerBase(ABC):
|
||||
extra={"model": task.model},
|
||||
)
|
||||
|
||||
def _build_billing_rule_snapshot(
|
||||
self, rule_lookup: BillingRuleLookupResult | None
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
构建 billing_rule 快照,用于冻结到视频任务的 request_metadata 中。
|
||||
|
||||
__all__ = ["VideoHandlerBase", "sanitize_error_message"]
|
||||
快照确保异步任务完成时使用创建时刻的计费规则,避免规则变更导致成本计算不一致。
|
||||
"""
|
||||
if not rule_lookup:
|
||||
return {"status": "no_rule"}
|
||||
|
||||
rule = rule_lookup.rule
|
||||
return {
|
||||
"status": "ok",
|
||||
"scope": rule_lookup.scope,
|
||||
"effective_task_type": rule_lookup.effective_task_type,
|
||||
"rule_id": rule.id,
|
||||
"rule_name": rule.name,
|
||||
"expression": rule.expression,
|
||||
"variables": rule.variables,
|
||||
"dimension_mappings": rule.dimension_mappings,
|
||||
}
|
||||
|
||||
|
||||
__all__ = ["VideoHandlerBase", "normalize_gemini_operation_id", "sanitize_error_message"]
|
||||
|
||||
@@ -14,8 +14,13 @@ from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
VideoHandlerBase,
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
|
||||
from src.core.api_format.conversion.internal_video import (
|
||||
InternalVideoRequest,
|
||||
@@ -27,6 +32,7 @@ from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
@@ -34,8 +40,6 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
FORMAT_ID = "GEMINI"
|
||||
|
||||
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
|
||||
POLL_INTERVAL_SECONDS = 10
|
||||
MAX_POLL_COUNT = 360
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -80,16 +84,39 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
candidate = await self._select_candidate(internal_request.model)
|
||||
candidate, candidate_keys, rule_lookup = await self._select_candidate(
|
||||
internal_request.model
|
||||
)
|
||||
if not candidate:
|
||||
raise HTTPException(
|
||||
status_code=503, detail="No available provider for video generation"
|
||||
detail = "No available provider for video generation"
|
||||
if config.billing_require_rule:
|
||||
detail = "No available provider with billing rule for video generation"
|
||||
raise HTTPException(status_code=503, detail=detail)
|
||||
|
||||
# 冻结 billing_rule 配置(用于异步任务的成本一致性)
|
||||
# 复用 _select_candidate 中已查询的结果;billing_require_rule=false 时需补查
|
||||
if rule_lookup is None:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=candidate.provider.id,
|
||||
model_name=internal_request.model,
|
||||
task_type="video",
|
||||
)
|
||||
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
|
||||
|
||||
upstream_key, endpoint, key, auth_info = await self._resolve_upstream_key(candidate)
|
||||
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
|
||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint, auth_info)
|
||||
|
||||
api_key_header = (
|
||||
headers.get("x-goog-api-key", "")[:10] + "..."
|
||||
if headers.get("x-goog-api-key")
|
||||
else "MISSING"
|
||||
)
|
||||
logger.info(
|
||||
f"[GeminiVeoHandler] Create task: endpoint_id={endpoint.id}, base_url={endpoint.base_url}, upstream_url={upstream_url}, api_key_prefix={api_key_header}"
|
||||
)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||
if response.status_code >= 400:
|
||||
@@ -99,18 +126,25 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
external_task_id = str(payload.get("name") or "")
|
||||
if not external_task_id:
|
||||
raise HTTPException(status_code=502, detail="Upstream returned empty task id")
|
||||
external_task_id = normalize_gemini_operation_id(external_task_id)
|
||||
|
||||
task = self._create_task_record(
|
||||
external_task_id=external_task_id,
|
||||
candidate=candidate,
|
||||
original_request_body=original_request_body,
|
||||
internal_request=internal_request,
|
||||
candidate_keys=candidate_keys,
|
||||
original_headers=original_headers,
|
||||
billing_rule_snapshot=billing_rule_snapshot,
|
||||
)
|
||||
try:
|
||||
self.db.add(task)
|
||||
self.db.flush() # 先 flush 检测冲突
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
logger.info(
|
||||
f"[GeminiVeoHandler] Task created: id={task.id}, external_task_id={task.external_task_id}, user_id={task.user_id}"
|
||||
)
|
||||
except IntegrityError:
|
||||
self.db.rollback()
|
||||
raise HTTPException(status_code=409, detail="Task already exists")
|
||||
@@ -136,6 +170,8 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
) -> JSONResponse:
|
||||
# Gemini 使用 operations/{id} 格式,需要按 external_task_id 查找
|
||||
task = self._get_task_by_external_id(task_id)
|
||||
|
||||
# 直接从数据库返回任务状态(后台轮询服务会持续更新状态)
|
||||
internal_task = self._task_to_internal(task)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task)
|
||||
return JSONResponse(response_body)
|
||||
@@ -272,7 +308,10 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _select_candidate(self, model_name: str) -> ProviderCandidate | None:
|
||||
async def _select_candidate(
|
||||
self, model_name: str
|
||||
) -> tuple[ProviderCandidate | None, list[dict[str, Any]], BillingRuleLookupResult | None]:
|
||||
"""选择候选 key,返回 (选中的候选, 所有候选列表, 选中候选的 billing rule lookup)"""
|
||||
scheduler = CacheAwareScheduler()
|
||||
candidates, _ = await scheduler.list_all_candidates(
|
||||
db=self.db,
|
||||
@@ -282,11 +321,47 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
user_api_key=self.api_key,
|
||||
max_candidates=10,
|
||||
)
|
||||
for candidate in candidates:
|
||||
# 记录所有候选 key 信息
|
||||
candidate_keys = []
|
||||
selected_candidate = None
|
||||
selected_index = -1
|
||||
selected_rule_lookup: BillingRuleLookupResult | None = None
|
||||
for idx, candidate in enumerate(candidates):
|
||||
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
|
||||
if auth_type in {"api_key", "vertex_ai"}:
|
||||
return candidate
|
||||
return None
|
||||
has_billing_rule = True
|
||||
rule_lookup: BillingRuleLookupResult | None = None
|
||||
if config.billing_require_rule:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=candidate.provider.id,
|
||||
model_name=model_name,
|
||||
task_type="video",
|
||||
)
|
||||
has_billing_rule = rule_lookup is not None
|
||||
candidate_info = {
|
||||
"index": idx,
|
||||
"provider_id": candidate.provider.id,
|
||||
"provider_name": candidate.provider.name,
|
||||
"endpoint_id": candidate.endpoint.id,
|
||||
"key_id": candidate.key.id,
|
||||
"key_name": candidate.key.name,
|
||||
"auth_type": auth_type,
|
||||
"has_billing_rule": has_billing_rule,
|
||||
"priority": getattr(candidate.key, "priority", 0) or 0,
|
||||
}
|
||||
candidate_keys.append(candidate_info)
|
||||
if (
|
||||
selected_candidate is None
|
||||
and auth_type in {"api_key", "vertex_ai"}
|
||||
and has_billing_rule
|
||||
):
|
||||
selected_candidate = candidate
|
||||
selected_index = idx
|
||||
selected_rule_lookup = rule_lookup
|
||||
# 标记选中的候选
|
||||
if selected_index >= 0:
|
||||
candidate_keys[selected_index]["selected"] = True
|
||||
return selected_candidate, candidate_keys, selected_rule_lookup
|
||||
|
||||
async def _resolve_upstream_key(
|
||||
self, candidate: ProviderCandidate
|
||||
@@ -351,8 +426,31 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
candidate: ProviderCandidate,
|
||||
original_request_body: dict[str, Any],
|
||||
internal_request: Any,
|
||||
candidate_keys: list[dict[str, Any]] | None = None,
|
||||
original_headers: dict[str, str] | None = None,
|
||||
billing_rule_snapshot: dict[str, Any] | None = None,
|
||||
) -> VideoTask:
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# 构建请求元数据(使用追踪信息)
|
||||
request_metadata = {
|
||||
"candidate_keys": candidate_keys or [],
|
||||
"selected_key_id": candidate.key.id,
|
||||
"selected_endpoint_id": candidate.endpoint.id,
|
||||
"client_ip": self.client_ip,
|
||||
"user_agent": self.user_agent,
|
||||
"request_id": self.request_id,
|
||||
"billing_rule_snapshot": billing_rule_snapshot,
|
||||
}
|
||||
# 记录请求头(脱敏处理)
|
||||
if original_headers:
|
||||
safe_headers = {
|
||||
k: v
|
||||
for k, v in original_headers.items()
|
||||
if k.lower() not in {"authorization", "x-api-key", "x-goog-api-key", "cookie"}
|
||||
}
|
||||
request_metadata["request_headers"] = safe_headers
|
||||
|
||||
return VideoTask(
|
||||
id=str(uuid4()),
|
||||
external_task_id=external_task_id,
|
||||
@@ -373,18 +471,21 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
aspect_ratio=internal_request.aspect_ratio,
|
||||
status=VideoStatus.SUBMITTED.value,
|
||||
progress_percent=0,
|
||||
poll_interval_seconds=self.POLL_INTERVAL_SECONDS,
|
||||
next_poll_at=now + timedelta(seconds=self.POLL_INTERVAL_SECONDS),
|
||||
poll_interval_seconds=config.video_poll_interval_seconds,
|
||||
next_poll_at=now + timedelta(seconds=config.video_poll_interval_seconds),
|
||||
poll_count=0,
|
||||
max_poll_count=self.MAX_POLL_COUNT,
|
||||
max_poll_count=config.video_max_poll_count,
|
||||
submitted_at=now,
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
|
||||
"""按 external_task_id 查找任务(Gemini 使用 operations/{id} 格式)"""
|
||||
normalized_id = external_id
|
||||
if not normalized_id.startswith("operations/"):
|
||||
normalized_id = f"operations/{normalized_id}"
|
||||
normalized_id = normalize_gemini_operation_id(external_id)
|
||||
|
||||
logger.info(
|
||||
f"[GeminiVeoHandler] Looking for task: normalized_id={normalized_id}, user_id={self.user.id}"
|
||||
)
|
||||
|
||||
task = (
|
||||
self.db.query(VideoTask)
|
||||
@@ -395,7 +496,13 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
.first()
|
||||
)
|
||||
if not task:
|
||||
logger.warning(
|
||||
f"[GeminiVeoHandler] Task not found: normalized_id={normalized_id}, user_id={self.user.id}"
|
||||
)
|
||||
raise HTTPException(status_code=404, detail="Video task not found")
|
||||
logger.info(
|
||||
f"[GeminiVeoHandler] Task found: id={task.id}, external_task_id={task.external_task_id}"
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
|
||||
from src.core.api_format.conversion.internal_video import (
|
||||
InternalVideoRequest,
|
||||
@@ -27,6 +28,7 @@ from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
@@ -34,8 +36,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
FORMAT_ID = "OPENAI"
|
||||
|
||||
DEFAULT_BASE_URL = "https://api.openai.com"
|
||||
POLL_INTERVAL_SECONDS = 10
|
||||
MAX_POLL_COUNT = 360
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -73,11 +73,25 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
internal_request = self._normalizer.video_request_to_internal(original_request_body)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
candidate = await self._select_candidate(internal_request.model)
|
||||
candidate, candidate_keys, rule_lookup = await self._select_candidate(
|
||||
internal_request.model
|
||||
)
|
||||
if not candidate:
|
||||
raise HTTPException(
|
||||
status_code=503, detail="No available provider for video generation"
|
||||
detail = "No available provider for video generation"
|
||||
if config.billing_require_rule:
|
||||
detail = "No available provider with billing rule for video generation"
|
||||
raise HTTPException(status_code=503, detail=detail)
|
||||
|
||||
# 冻结 billing_rule 配置(用于异步任务的成本一致性)
|
||||
# 复用 _select_candidate 中已查询的结果;billing_require_rule=false 时需补查
|
||||
if rule_lookup is None:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=candidate.provider.id,
|
||||
model_name=internal_request.model,
|
||||
task_type="video",
|
||||
)
|
||||
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
|
||||
|
||||
upstream_key, endpoint, provider_key = await self._resolve_upstream_key(candidate)
|
||||
upstream_url = self._build_upstream_url(endpoint.base_url)
|
||||
@@ -98,6 +112,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
candidate=candidate,
|
||||
original_request_body=original_request_body,
|
||||
internal_request=internal_request,
|
||||
candidate_keys=candidate_keys,
|
||||
original_headers=original_headers,
|
||||
billing_rule_snapshot=billing_rule_snapshot,
|
||||
)
|
||||
try:
|
||||
self.db.add(task)
|
||||
@@ -282,7 +299,10 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _select_candidate(self, model_name: str) -> ProviderCandidate | None:
|
||||
async def _select_candidate(
|
||||
self, model_name: str
|
||||
) -> tuple[ProviderCandidate | None, list[dict[str, Any]], BillingRuleLookupResult | None]:
|
||||
"""选择候选 key,返回 (选中的候选, 所有候选列表, 选中候选的 billing rule lookup)"""
|
||||
scheduler = CacheAwareScheduler()
|
||||
candidates, _ = await scheduler.list_all_candidates(
|
||||
db=self.db,
|
||||
@@ -292,11 +312,43 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
user_api_key=self.api_key,
|
||||
max_candidates=10,
|
||||
)
|
||||
for candidate in candidates:
|
||||
# 记录所有候选 key 信息
|
||||
candidate_keys = []
|
||||
selected_candidate = None
|
||||
selected_index = -1
|
||||
selected_rule_lookup: BillingRuleLookupResult | None = None
|
||||
for idx, candidate in enumerate(candidates):
|
||||
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
|
||||
if auth_type == "api_key":
|
||||
return candidate
|
||||
return None
|
||||
has_billing_rule = True
|
||||
rule_lookup: BillingRuleLookupResult | None = None
|
||||
if config.billing_require_rule:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=candidate.provider.id,
|
||||
model_name=model_name,
|
||||
task_type="video",
|
||||
)
|
||||
has_billing_rule = rule_lookup is not None
|
||||
candidate_info = {
|
||||
"index": idx,
|
||||
"provider_id": candidate.provider.id,
|
||||
"provider_name": candidate.provider.name,
|
||||
"endpoint_id": candidate.endpoint.id,
|
||||
"key_id": candidate.key.id,
|
||||
"key_name": candidate.key.name,
|
||||
"auth_type": auth_type,
|
||||
"has_billing_rule": has_billing_rule,
|
||||
"priority": getattr(candidate.key, "priority", 0) or 0,
|
||||
}
|
||||
candidate_keys.append(candidate_info)
|
||||
if selected_candidate is None and auth_type == "api_key" and has_billing_rule:
|
||||
selected_candidate = candidate
|
||||
selected_index = idx
|
||||
selected_rule_lookup = rule_lookup
|
||||
# 标记选中的候选
|
||||
if selected_index >= 0:
|
||||
candidate_keys[selected_index]["selected"] = True
|
||||
return selected_candidate, candidate_keys, selected_rule_lookup
|
||||
|
||||
async def _resolve_upstream_key(
|
||||
self, candidate: ProviderCandidate
|
||||
@@ -342,9 +394,32 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
candidate: ProviderCandidate,
|
||||
original_request_body: dict[str, Any],
|
||||
internal_request: InternalVideoRequest,
|
||||
candidate_keys: list[dict[str, Any]] | None = None,
|
||||
original_headers: dict[str, str] | None = None,
|
||||
billing_rule_snapshot: dict[str, Any] | None = None,
|
||||
) -> VideoTask:
|
||||
now = datetime.now(timezone.utc)
|
||||
size = internal_request.extra.get("original_size")
|
||||
|
||||
# 构建请求元数据(使用追踪信息)
|
||||
request_metadata = {
|
||||
"candidate_keys": candidate_keys or [],
|
||||
"selected_key_id": candidate.key.id,
|
||||
"selected_endpoint_id": candidate.endpoint.id,
|
||||
"client_ip": self.client_ip,
|
||||
"user_agent": self.user_agent,
|
||||
"request_id": self.request_id,
|
||||
"billing_rule_snapshot": billing_rule_snapshot,
|
||||
}
|
||||
# 记录请求头(脱敏处理)
|
||||
if original_headers:
|
||||
safe_headers = {
|
||||
k: v
|
||||
for k, v in original_headers.items()
|
||||
if k.lower() not in {"authorization", "x-api-key", "cookie"}
|
||||
}
|
||||
request_metadata["request_headers"] = safe_headers
|
||||
|
||||
return VideoTask(
|
||||
id=str(uuid4()),
|
||||
external_task_id=external_task_id,
|
||||
@@ -366,11 +441,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
size=size,
|
||||
status=VideoStatus.SUBMITTED.value,
|
||||
progress_percent=0,
|
||||
poll_interval_seconds=self.POLL_INTERVAL_SECONDS,
|
||||
next_poll_at=now + timedelta(seconds=self.POLL_INTERVAL_SECONDS),
|
||||
poll_interval_seconds=config.video_poll_interval_seconds,
|
||||
next_poll_at=now + timedelta(seconds=config.video_poll_interval_seconds),
|
||||
poll_count=0,
|
||||
max_poll_count=self.MAX_POLL_COUNT,
|
||||
max_poll_count=config.video_max_poll_count,
|
||||
submitted_at=now,
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
|
||||
|
||||
@@ -106,7 +106,7 @@ async def create_video_veo(model: str, http_request: Request, db: Session = Depe
|
||||
)
|
||||
|
||||
|
||||
@router.get("/v1beta/operations/{operation_id}")
|
||||
@router.get("/v1beta/operations/{operation_id:path}")
|
||||
async def get_video_veo(
|
||||
operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
||||
) -> Any:
|
||||
@@ -148,7 +148,7 @@ async def cancel_video_veo(
|
||||
)
|
||||
|
||||
|
||||
@router.get("/v1beta/operations/{operation_id}/content")
|
||||
@router.get("/v1beta/operations/{operation_id:path}/content")
|
||||
async def download_video_content_veo(
|
||||
operation_id: str, http_request: Request, db: Session = Depends(get_db)
|
||||
) -> Any:
|
||||
|
||||
Reference in New Issue
Block a user