mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
refactor: 重构异步任务系统和计费服务架构
- 重构任务系统:新增 lifecycle (TaskStatus/BillingStatus)、context、application 模块 - 将 video tasks 泛化为 async tasks,支持更通用的异步任务管理 - 新增 Gemini Files 管理模块和管理界面 - 重构 billing 服务:拆分 schema.py 和 service.py - 新增 candidate 服务模块用于请求候选管理 - 数据库迁移:添加 billing_status、request_id、gemini_file_mappings 表和索引 - 移除废弃的 video_telemetry、task orchestrator 等模块
This commit is contained in:
@@ -804,8 +804,33 @@ class OAuthService:
|
||||
|
||||
cfg = OAuthService._get_provider_config(db, provider_type)
|
||||
|
||||
# Read all required fields first, then release DB connection before any awaits.
|
||||
# This prevents holding a pooled connection while doing network I/O.
|
||||
auth_url = provider.get_effective_authorization_url(cfg)
|
||||
token_url = provider.get_effective_token_url(cfg)
|
||||
redirect_uri = cfg.redirect_uri
|
||||
client_id = cfg.client_id
|
||||
has_secret = bool(cfg.client_secret_encrypted)
|
||||
client_secret = cfg.get_client_secret() if has_secret else None
|
||||
|
||||
# Release DB connection (safe only when session has no pending changes).
|
||||
try:
|
||||
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
|
||||
except Exception:
|
||||
has_pending_changes = False
|
||||
if not has_pending_changes:
|
||||
original_expire_on_commit = getattr(db, "expire_on_commit", True)
|
||||
db.expire_on_commit = False
|
||||
try:
|
||||
if db.in_transaction():
|
||||
db.commit()
|
||||
except Exception:
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
db.expire_on_commit = original_expire_on_commit
|
||||
|
||||
async def _reachable(url: str) -> bool:
|
||||
try:
|
||||
@@ -823,7 +848,7 @@ class OAuthService:
|
||||
secret_status = "unknown"
|
||||
details = ""
|
||||
|
||||
if cfg.client_secret_encrypted:
|
||||
if has_secret and client_secret:
|
||||
# 使用无效 code 做一次 token 请求(仅做粗略判定)
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
@@ -834,9 +859,9 @@ class OAuthService:
|
||||
data={
|
||||
"grant_type": "authorization_code",
|
||||
"code": "invalid",
|
||||
"redirect_uri": cfg.redirect_uri,
|
||||
"client_id": cfg.client_id,
|
||||
"client_secret": cfg.get_client_secret(),
|
||||
"redirect_uri": redirect_uri,
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
},
|
||||
)
|
||||
try:
|
||||
|
||||
@@ -29,6 +29,8 @@ from src.services.billing.models import (
|
||||
CostBreakdown,
|
||||
StandardizedUsage,
|
||||
)
|
||||
from src.services.billing.schema import BillingSnapshot, CostResult
|
||||
from src.services.billing.service import BillingService
|
||||
from src.services.billing.templates import BILLING_TEMPLATE_REGISTRY, BillingTemplates
|
||||
from src.services.billing.usage_mapper import UsageMapper, map_usage, map_usage_from_response
|
||||
|
||||
@@ -44,6 +46,10 @@ __all__ = [
|
||||
# 计算器
|
||||
"BillingCalculator",
|
||||
"calculate_request_cost",
|
||||
# 统一入口(Phase2)
|
||||
"BillingService",
|
||||
"BillingSnapshot",
|
||||
"CostResult",
|
||||
# 映射器
|
||||
"UsageMapper",
|
||||
"map_usage",
|
||||
|
||||
64
src/services/billing/schema.py
Normal file
64
src/services/billing/schema.py
Normal file
@@ -0,0 +1,64 @@
|
||||
"""
|
||||
Billing schema (stable contracts)
|
||||
|
||||
These dataclasses are meant to be stored in `Usage.request_metadata` / `Task.request_metadata`
|
||||
for auditability. They are internal-only and MUST NOT be exposed to end users without sanitizing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
|
||||
BILLING_SNAPSHOT_SCHEMA_VERSION = "1.0"
|
||||
|
||||
BillingSnapshotStatus = Literal["complete", "incomplete", "no_rule", "legacy"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BillingSnapshot:
|
||||
"""Stable billing snapshot for audit."""
|
||||
|
||||
schema_version: str = BILLING_SNAPSHOT_SCHEMA_VERSION
|
||||
|
||||
# Rule info (optional for legacy/no_rule)
|
||||
rule_id: str | None = None
|
||||
rule_name: str | None = None
|
||||
scope: str | None = None
|
||||
|
||||
# Rule expression (internal, do not expose to clients)
|
||||
expression: str | None = None
|
||||
|
||||
# Dimensions
|
||||
dimensions_used: dict[str, Any] = field(default_factory=dict)
|
||||
missing_required: list[str] = field(default_factory=list)
|
||||
|
||||
# Result
|
||||
cost: float = 0.0
|
||||
status: BillingSnapshotStatus = "no_rule"
|
||||
|
||||
# Audit
|
||||
calculated_at: str = "" # ISO 8601
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"schema_version": self.schema_version,
|
||||
"rule_id": self.rule_id,
|
||||
"rule_name": self.rule_name,
|
||||
"scope": self.scope,
|
||||
"expression": self.expression,
|
||||
"dimensions_used": self.dimensions_used,
|
||||
"missing_required": self.missing_required,
|
||||
"cost": self.cost,
|
||||
"status": self.status,
|
||||
"calculated_at": self.calculated_at,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CostResult:
|
||||
"""Billing calculation output."""
|
||||
|
||||
cost: float
|
||||
status: BillingSnapshotStatus
|
||||
snapshot: BillingSnapshot
|
||||
145
src/services/billing/service.py
Normal file
145
src/services/billing/service.py
Normal file
@@ -0,0 +1,145 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.services.billing.dimension_collector_service import DimensionCollectorService
|
||||
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
|
||||
from src.services.billing.rule_service import BillingRuleService
|
||||
from src.services.model.cost import ModelCostService
|
||||
|
||||
from .schema import BILLING_SNAPSHOT_SCHEMA_VERSION, BillingSnapshot, CostResult
|
||||
|
||||
|
||||
class BillingService:
|
||||
"""
|
||||
BillingService (pure-ish application helper for billing domain).
|
||||
|
||||
Notes:
|
||||
- This service **does not** write Usage rows.
|
||||
- It may read billing rules & collectors from DB.
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
self._formula_engine = FormulaEngine()
|
||||
self._dimension_collector = DimensionCollectorService(db)
|
||||
|
||||
def collect_dimensions(
|
||||
self,
|
||||
*,
|
||||
api_format: str | None,
|
||||
task_type: str | None,
|
||||
request: dict[str, Any] | None = None,
|
||||
response: dict[str, Any] | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
base_dimensions: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return self._dimension_collector.collect_dimensions(
|
||||
api_format=api_format,
|
||||
task_type=task_type,
|
||||
request=request,
|
||||
response=response,
|
||||
metadata=metadata,
|
||||
base_dimensions=base_dimensions,
|
||||
)
|
||||
|
||||
def calculate(
|
||||
self,
|
||||
*,
|
||||
task_type: str,
|
||||
model: str,
|
||||
provider_id: str,
|
||||
dimensions: dict[str, Any],
|
||||
strict_mode: bool | None = None,
|
||||
) -> CostResult:
|
||||
"""
|
||||
Calculate cost for a task.
|
||||
|
||||
Returns:
|
||||
CostResult (includes BillingSnapshot)
|
||||
|
||||
Raises:
|
||||
BillingIncompleteError: when strict_mode=True and required dims missing.
|
||||
"""
|
||||
strict = config.billing_strict_mode if strict_mode is None else bool(strict_mode)
|
||||
|
||||
lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=provider_id,
|
||||
model_name=model,
|
||||
task_type=task_type,
|
||||
)
|
||||
|
||||
if lookup and lookup.rule and lookup.rule.expression:
|
||||
rule = lookup.rule
|
||||
result = self._formula_engine.evaluate(
|
||||
expression=rule.expression,
|
||||
variables=rule.variables or {},
|
||||
dimensions=dimensions,
|
||||
dimension_mappings=rule.dimension_mappings or {},
|
||||
strict_mode=strict,
|
||||
)
|
||||
cost = float(result.cost) if result.status == "complete" else 0.0
|
||||
snapshot = BillingSnapshot(
|
||||
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
|
||||
rule_id=str(rule.id),
|
||||
rule_name=str(rule.name),
|
||||
scope=str(getattr(lookup, "scope", None) or ""),
|
||||
expression=str(rule.expression),
|
||||
dimensions_used=dimensions,
|
||||
missing_required=result.missing_required,
|
||||
cost=cost,
|
||||
status=result.status,
|
||||
calculated_at=datetime.now(timezone.utc).isoformat(),
|
||||
)
|
||||
return CostResult(cost=cost, status=result.status, snapshot=snapshot)
|
||||
|
||||
# No rule fallback
|
||||
if task_type in ("chat", "cli"):
|
||||
input_tokens = int(dimensions.get("input_tokens") or 0)
|
||||
output_tokens = int(dimensions.get("output_tokens") or 0)
|
||||
cost = float(
|
||||
ModelCostService.calculate_cost(
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
)
|
||||
snapshot = BillingSnapshot(
|
||||
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
|
||||
rule_id=None,
|
||||
rule_name=None,
|
||||
scope=None,
|
||||
expression=None,
|
||||
dimensions_used=dimensions,
|
||||
missing_required=[],
|
||||
cost=cost,
|
||||
status="legacy",
|
||||
calculated_at=datetime.now(timezone.utc).isoformat(),
|
||||
)
|
||||
return CostResult(cost=cost, status="legacy", snapshot=snapshot)
|
||||
|
||||
logger.warning(
|
||||
"No billing rule for task (task_type=%s, model=%s, provider_id=%s)",
|
||||
task_type,
|
||||
model,
|
||||
provider_id,
|
||||
)
|
||||
snapshot = BillingSnapshot(
|
||||
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
|
||||
rule_id=None,
|
||||
rule_name=None,
|
||||
scope=None,
|
||||
expression=None,
|
||||
dimensions_used=dimensions,
|
||||
missing_required=[],
|
||||
cost=0.0,
|
||||
status="no_rule",
|
||||
calculated_at=datetime.now(timezone.utc).isoformat(),
|
||||
)
|
||||
return CostResult(cost=0.0, status="no_rule", snapshot=snapshot)
|
||||
72
src/services/cache/aware_scheduler.py
vendored
72
src/services/cache/aware_scheduler.py
vendored
@@ -182,6 +182,43 @@ class CacheAwareScheduler:
|
||||
"last_reservation_result": None,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _release_db_connection_before_await(db: Session) -> None:
|
||||
"""
|
||||
Best-effort: end a read-only transaction before awaiting async I/O.
|
||||
|
||||
This scheduler does a lot of async work (cache/Redis) mixed with sync SQLAlchemy reads.
|
||||
If a SELECT has already started a transaction, the pooled connection can remain checked
|
||||
out while we await, causing pool pressure under concurrency.
|
||||
|
||||
Safety:
|
||||
- Only commits when the Session has no ORM pending changes.
|
||||
- Temporarily disables expire_on_commit to keep already-loaded ORM objects usable.
|
||||
"""
|
||||
try:
|
||||
if db is None:
|
||||
return
|
||||
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
|
||||
if has_pending_changes:
|
||||
return
|
||||
if not db.in_transaction():
|
||||
return
|
||||
|
||||
original_expire_on_commit = getattr(db, "expire_on_commit", True)
|
||||
db.expire_on_commit = False
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
db.expire_on_commit = original_expire_on_commit
|
||||
except Exception:
|
||||
# Never let this optimization break scheduling
|
||||
return
|
||||
|
||||
async def _ensure_initialized(self) -> None:
|
||||
"""确保所有异步组件已初始化"""
|
||||
if self._affinity_manager is None:
|
||||
@@ -577,6 +614,8 @@ class CacheAwareScheduler:
|
||||
Returns:
|
||||
(候选列表, global_model_id) - global_model_id 用于缓存亲和性
|
||||
"""
|
||||
# If the caller already touched the DB, release the connection before we do async work.
|
||||
self._release_db_connection_before_await(db)
|
||||
await self._ensure_initialized()
|
||||
|
||||
target_format = normalize_endpoint_signature(api_format)
|
||||
@@ -648,6 +687,9 @@ class CacheAwareScheduler:
|
||||
provider_limit=provider_limit,
|
||||
)
|
||||
|
||||
# Provider query starts a transaction; release connection before entering async candidate build.
|
||||
self._release_db_connection_before_await(db)
|
||||
|
||||
logger.debug(
|
||||
"[Scheduler] Found %d active providers",
|
||||
len(providers),
|
||||
@@ -680,8 +722,13 @@ class CacheAwareScheduler:
|
||||
|
||||
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤)
|
||||
from src.config.settings import config
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
global_conversion_enabled = config.format_conversion_enabled
|
||||
# 全局格式转换开关:优先使用数据库配置,回退到环境变量
|
||||
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
|
||||
# 如果环境变量明确禁用,则禁用(环境变量可作为强制禁用开关)
|
||||
if not config.format_conversion_enabled:
|
||||
global_conversion_enabled = False
|
||||
candidates = await self._build_candidates(
|
||||
db=db,
|
||||
providers=providers,
|
||||
@@ -801,6 +848,9 @@ class CacheAwareScheduler:
|
||||
- supported_capabilities: 模型支持的能力列表
|
||||
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
|
||||
"""
|
||||
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
|
||||
self._release_db_connection_before_await(db)
|
||||
|
||||
# 使用 ModelCacheService 解析模型名称(支持映射名)
|
||||
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
||||
db, model_name
|
||||
@@ -1120,18 +1170,34 @@ class CacheAwareScheduler:
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
|
||||
# 计算格式转换的有效开关状态(三层优先级)
|
||||
# 全局 ON → 强制允许(跳过端点检查)
|
||||
# 全局 OFF → 提供商 ON → 强制允许(跳过端点检查)
|
||||
# 全局 OFF → 提供商 OFF → 看端点配置
|
||||
provider_allows_conversion = getattr(provider, "enable_format_conversion", True)
|
||||
effective_conversion_enabled = (
|
||||
global_conversion_enabled or provider_allows_conversion
|
||||
)
|
||||
# 如果全局或提供商开关为 ON,跳过端点配置检查
|
||||
skip_endpoint_check = global_conversion_enabled or provider_allows_conversion
|
||||
|
||||
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
|
||||
client_format_str,
|
||||
endpoint_format_str,
|
||||
getattr(endpoint, "format_acceptance_config", None),
|
||||
is_stream,
|
||||
global_conversion_enabled,
|
||||
effective_conversion_enabled,
|
||||
skip_endpoint_check=skip_endpoint_check,
|
||||
)
|
||||
logger.debug(
|
||||
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, reason=%s",
|
||||
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, "
|
||||
"global=%s, provider=%s, skip_endpoint=%s, reason=%s",
|
||||
client_format_str,
|
||||
endpoint_format_str,
|
||||
is_compatible,
|
||||
global_conversion_enabled,
|
||||
provider_allows_conversion,
|
||||
skip_endpoint_check,
|
||||
_compat_reason,
|
||||
)
|
||||
if not is_compatible:
|
||||
|
||||
28
src/services/candidate/__init__.py
Normal file
28
src/services/candidate/__init__.py
Normal file
@@ -0,0 +1,28 @@
|
||||
"""
|
||||
Candidate domain (Phase2)
|
||||
|
||||
This package centralizes:
|
||||
- candidate resolving (Provider/Endpoint/Key combinations)
|
||||
- request_candidates recording & audit
|
||||
- failover execution policies
|
||||
"""
|
||||
|
||||
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
|
||||
from src.services.candidate.schema import (
|
||||
CANDIDATE_KEY_SCHEMA_VERSION,
|
||||
CandidateKey,
|
||||
CandidateResult,
|
||||
)
|
||||
from src.services.candidate.service import CandidateService
|
||||
|
||||
__all__ = [
|
||||
"CandidateService",
|
||||
# schema
|
||||
"CANDIDATE_KEY_SCHEMA_VERSION",
|
||||
"CandidateKey",
|
||||
"CandidateResult",
|
||||
# policies
|
||||
"RetryMode",
|
||||
"RetryPolicy",
|
||||
"SkipPolicy",
|
||||
]
|
||||
57
src/services/candidate/failover.py
Normal file
57
src/services/candidate/failover.py
Normal file
@@ -0,0 +1,57 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Protocol
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
|
||||
from .policy import RetryPolicy, SkipPolicy
|
||||
from .schema import CandidateKey, CandidateResult
|
||||
|
||||
|
||||
class AttemptFunc(Protocol):
|
||||
async def __call__(self, candidate: ProviderCandidate) -> Any: ...
|
||||
|
||||
|
||||
class FailoverEngine:
|
||||
"""
|
||||
FailoverEngine executes candidate attempts under policies.
|
||||
|
||||
Phase2 scaffolding: implementation will gradually replace legacy orchestrators.
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
*,
|
||||
candidates: list[ProviderCandidate],
|
||||
attempt_func: AttemptFunc,
|
||||
retry_policy: RetryPolicy,
|
||||
skip_policy: SkipPolicy,
|
||||
request_id: str | None = None,
|
||||
max_candidates: int | None = None,
|
||||
) -> CandidateResult:
|
||||
# NOTE: intentionally minimal for now; legacy orchestrators still in use.
|
||||
# This will be implemented when migrating video/chat flows to CandidateService.
|
||||
_ = (retry_policy, skip_policy, request_id, max_candidates)
|
||||
candidate_keys: list[CandidateKey] = []
|
||||
for idx, cand in enumerate(candidates):
|
||||
candidate_keys.append(
|
||||
CandidateKey(
|
||||
candidate_index=idx,
|
||||
provider_id=str(cand.provider.id),
|
||||
provider_name=str(cand.provider.name),
|
||||
endpoint_id=str(cand.endpoint.id),
|
||||
key_id=str(cand.key.id),
|
||||
key_name=str(getattr(cand.key, "name", "") or ""),
|
||||
auth_type=str(getattr(cand.key, "auth_type", "") or ""),
|
||||
priority=int(getattr(cand.key, "priority", 0) or 0),
|
||||
is_cached=bool(getattr(cand, "is_cached", False)),
|
||||
status="available",
|
||||
)
|
||||
)
|
||||
|
||||
raise NotImplementedError("FailoverEngine.execute is not implemented yet")
|
||||
41
src/services/candidate/policy.py
Normal file
41
src/services/candidate/policy.py
Normal file
@@ -0,0 +1,41 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class RetryMode(str, Enum):
|
||||
"""Retry mode for candidate attempts."""
|
||||
|
||||
PRE_EXPAND = "pre_expand" # pre-create retry slots (sync)
|
||||
ON_DEMAND = "on_demand" # create retry record only when retry happens
|
||||
DISABLED = "disabled" # no retry
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RetryPolicy:
|
||||
"""Unified retry policy."""
|
||||
|
||||
mode: RetryMode = RetryMode.DISABLED
|
||||
max_retries: int = 1
|
||||
retry_on_cached_only: bool = True
|
||||
|
||||
@classmethod
|
||||
def for_sync_task(cls) -> "RetryPolicy":
|
||||
return cls(mode=RetryMode.PRE_EXPAND, max_retries=2)
|
||||
|
||||
@classmethod
|
||||
def for_async_task(cls) -> "RetryPolicy":
|
||||
return cls(mode=RetryMode.DISABLED, max_retries=1)
|
||||
|
||||
@classmethod
|
||||
def for_async_submit_with_retry(cls) -> "RetryPolicy":
|
||||
return cls(mode=RetryMode.ON_DEMAND, max_retries=2)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SkipPolicy:
|
||||
"""Rules for skipping unsupported candidates."""
|
||||
|
||||
allow_format_conversion: bool = True
|
||||
supported_auth_types: set[str] | None = None
|
||||
59
src/services/candidate/recorder.py
Normal file
59
src/services/candidate/recorder.py
Normal file
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import RequestCandidate
|
||||
|
||||
from .schema import CandidateKey
|
||||
|
||||
|
||||
class CandidateRecorder:
|
||||
"""Read helpers for RequestCandidate audit data."""
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
|
||||
def get_candidate_keys(self, request_id: str) -> list[CandidateKey]:
|
||||
rows: list[RequestCandidate] = (
|
||||
self.db.query(RequestCandidate)
|
||||
.filter(RequestCandidate.request_id == request_id)
|
||||
.order_by(RequestCandidate.candidate_index.asc(), RequestCandidate.retry_index.asc())
|
||||
.all()
|
||||
)
|
||||
|
||||
result: list[CandidateKey] = []
|
||||
for row in rows:
|
||||
provider_name = None
|
||||
if getattr(row, "provider", None) is not None:
|
||||
provider_name = getattr(row.provider, "name", None)
|
||||
|
||||
key_name = None
|
||||
auth_type = None
|
||||
priority = None
|
||||
if getattr(row, "key", None) is not None:
|
||||
key_name = getattr(row.key, "name", None)
|
||||
auth_type = getattr(row.key, "auth_type", None)
|
||||
priority = getattr(row.key, "priority", None)
|
||||
|
||||
result.append(
|
||||
CandidateKey(
|
||||
candidate_index=int(row.candidate_index or 0),
|
||||
retry_index=int(row.retry_index or 0),
|
||||
provider_id=str(row.provider_id) if row.provider_id else None,
|
||||
provider_name=str(provider_name) if provider_name else None,
|
||||
endpoint_id=str(row.endpoint_id) if row.endpoint_id else None,
|
||||
key_id=str(row.key_id) if row.key_id else None,
|
||||
key_name=str(key_name) if key_name else None,
|
||||
auth_type=str(auth_type) if auth_type else None,
|
||||
priority=int(priority) if priority is not None else None,
|
||||
is_cached=bool(getattr(row, "is_cached", False)),
|
||||
status=str(getattr(row, "status", "") or "pending"),
|
||||
skip_reason=getattr(row, "skip_reason", None),
|
||||
error_type=getattr(row, "error_type", None),
|
||||
error_message=getattr(row, "error_message", None),
|
||||
status_code=getattr(row, "status_code", None),
|
||||
latency_ms=getattr(row, "latency_ms", None),
|
||||
)
|
||||
)
|
||||
|
||||
return result
|
||||
10
src/services/candidate/resolver.py
Normal file
10
src/services/candidate/resolver.py
Normal file
@@ -0,0 +1,10 @@
|
||||
"""
|
||||
CandidateResolver facade import.
|
||||
|
||||
Phase2 keeps the implementation in `services/orchestration/` for compatibility,
|
||||
and gradually migrates it into this package.
|
||||
"""
|
||||
|
||||
from src.services.orchestration.candidate_resolver import CandidateResolver
|
||||
|
||||
__all__ = ["CandidateResolver"]
|
||||
71
src/services/candidate/schema.py
Normal file
71
src/services/candidate/schema.py
Normal file
@@ -0,0 +1,71 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
|
||||
CANDIDATE_KEY_SCHEMA_VERSION = "1.0"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CandidateKey:
|
||||
"""Stable candidate key snapshot for audit."""
|
||||
|
||||
schema_version: str = CANDIDATE_KEY_SCHEMA_VERSION
|
||||
|
||||
candidate_index: int = 0
|
||||
retry_index: int = 0
|
||||
|
||||
provider_id: str | None = None
|
||||
provider_name: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
key_name: str | None = None
|
||||
auth_type: str | None = None
|
||||
priority: int | None = None
|
||||
is_cached: bool = False
|
||||
|
||||
status: str = "pending" # pending/success/failed/skipped/available/...
|
||||
skip_reason: str | None = None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
status_code: int | None = None
|
||||
latency_ms: int | None = None
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
data: dict[str, Any] = {
|
||||
"schema_version": self.schema_version,
|
||||
"candidate_index": self.candidate_index,
|
||||
"retry_index": self.retry_index,
|
||||
"provider_id": self.provider_id,
|
||||
"provider_name": self.provider_name,
|
||||
"endpoint_id": self.endpoint_id,
|
||||
"key_id": self.key_id,
|
||||
"key_name": self.key_name,
|
||||
"auth_type": self.auth_type,
|
||||
"priority": self.priority,
|
||||
"is_cached": self.is_cached,
|
||||
"status": self.status,
|
||||
"skip_reason": self.skip_reason,
|
||||
"error_type": self.error_type,
|
||||
"error_message": self.error_message,
|
||||
"status_code": self.status_code,
|
||||
"latency_ms": self.latency_ms,
|
||||
}
|
||||
# drop Nones for compact audit payload
|
||||
return {k: v for k, v in data.items() if v is not None}
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class CandidateResult:
|
||||
"""Failover execution result."""
|
||||
|
||||
success: bool
|
||||
selected: ProviderCandidate | None
|
||||
selected_index: int | None
|
||||
candidate_keys: list[CandidateKey]
|
||||
|
||||
external_task_id: str | None = None
|
||||
error: Exception | None = None
|
||||
last_status_code: int | None = None
|
||||
515
src/services/candidate/service.py
Normal file
515
src/services/candidate/service.py
Normal file
@@ -0,0 +1,515 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import update
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, RequestCandidate
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
||||
from src.services.candidate.submit import (
|
||||
AllCandidatesFailedError,
|
||||
SubmitOutcome,
|
||||
UpstreamClientRequestError,
|
||||
)
|
||||
from src.services.orchestration.error_classifier import ErrorClassifier
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
from .recorder import CandidateRecorder
|
||||
from .resolver import CandidateResolver
|
||||
from .schema import CandidateKey
|
||||
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _sanitize(message: str, max_length: int = 200) -> str:
|
||||
if not message:
|
||||
return "request_failed"
|
||||
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
|
||||
|
||||
|
||||
class CandidateService:
|
||||
"""
|
||||
CandidateService (Facade).
|
||||
|
||||
Phase2 note: this is introduced as a new domain entrypoint. Legacy orchestrators
|
||||
still exist and will be migrated gradually to use this service.
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._cache_scheduler = None
|
||||
self._resolver: CandidateResolver | None = None
|
||||
self._error_classifier: ErrorClassifier | None = None
|
||||
self._recorder = CandidateRecorder(db)
|
||||
|
||||
async def _ensure_initialized(self) -> None:
|
||||
if self._cache_scheduler is not None:
|
||||
return
|
||||
|
||||
priority_mode = SystemConfigService.get_config(
|
||||
self.db,
|
||||
"provider_priority_mode",
|
||||
"provider",
|
||||
)
|
||||
scheduling_mode = SystemConfigService.get_config(
|
||||
self.db,
|
||||
"scheduling_mode",
|
||||
"cache_affinity",
|
||||
)
|
||||
self._cache_scheduler = await get_cache_aware_scheduler(
|
||||
self.redis,
|
||||
priority_mode=priority_mode,
|
||||
scheduling_mode=scheduling_mode,
|
||||
)
|
||||
self._resolver = CandidateResolver(db=self.db, cache_scheduler=self._cache_scheduler)
|
||||
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
|
||||
|
||||
async def resolve(
|
||||
self,
|
||||
*,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey | None = None,
|
||||
request_id: str | None = None,
|
||||
is_stream: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
preferred_key_ids: list[str] | None = None,
|
||||
) -> tuple[list[ProviderCandidate], str]:
|
||||
await self._ensure_initialized()
|
||||
assert self._resolver is not None
|
||||
return await self._resolver.fetch_candidates(
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
user_api_key=user_api_key,
|
||||
request_id=request_id,
|
||||
is_stream=is_stream,
|
||||
capability_requirements=capability_requirements,
|
||||
preferred_key_ids=preferred_key_ids,
|
||||
)
|
||||
|
||||
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
|
||||
"""
|
||||
Decide whether an upstream HTTP error is a client error (no failover).
|
||||
|
||||
Rules:
|
||||
- 401/403/429 are usually key/permission/ratelimit issues -> allow failover
|
||||
- other 4xx: stop only if ErrorClassifier says it's a client error
|
||||
"""
|
||||
if status_code in (401, 403, 429):
|
||||
return False
|
||||
if 400 <= status_code < 500:
|
||||
assert self._error_classifier is not None
|
||||
return self._error_classifier.is_client_error(error_text)
|
||||
return False
|
||||
|
||||
async def submit_with_failover(
|
||||
self,
|
||||
*,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey,
|
||||
request_id: str | None,
|
||||
task_type: str,
|
||||
submit_func: Any,
|
||||
extract_external_task_id: Any,
|
||||
supported_auth_types: set[str] | None = None,
|
||||
allow_format_conversion: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
max_candidates: int | None = None,
|
||||
) -> SubmitOutcome:
|
||||
"""
|
||||
Submit async task with failover, returning the selected candidate + external_task_id.
|
||||
|
||||
Phase2 submit entrypoint (replaces legacy submit orchestrator).
|
||||
"""
|
||||
# IMPORTANT:
|
||||
# This method awaits upstream HTTP calls. If we have an open DB transaction before awaiting,
|
||||
# the connection can be held for a long time (pool exhaustion under concurrency).
|
||||
#
|
||||
# Also note SQLAlchemy's default expire_on_commit=True would expire ORM objects and may
|
||||
# trigger unexpected lazy DB loads after we commit (potentially during the await).
|
||||
# We disable it temporarily to keep candidate/provider/key objects in-memory.
|
||||
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
|
||||
self.db.expire_on_commit = False
|
||||
await self._ensure_initialized()
|
||||
assert self._resolver is not None
|
||||
try:
|
||||
candidates, _global_model_id = await self._resolver.fetch_candidates(
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
user_api_key=user_api_key,
|
||||
request_id=request_id,
|
||||
is_stream=False,
|
||||
capability_requirements=capability_requirements,
|
||||
)
|
||||
|
||||
if not candidates:
|
||||
raise ProviderNotAvailableException("No candidates available")
|
||||
|
||||
if max_candidates is not None and max_candidates > 0:
|
||||
candidates = candidates[:max_candidates]
|
||||
|
||||
# Pre-create RequestCandidate records (no retry expand for async submit stage)
|
||||
record_map: dict[tuple[int, int], str] = {}
|
||||
if request_id:
|
||||
try:
|
||||
record_map = self.create_candidate_records(
|
||||
candidates=candidates,
|
||||
request_id=request_id,
|
||||
user_api_key=user_api_key,
|
||||
required_capabilities=capability_requirements,
|
||||
expand_retries=False,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[CandidateService] Failed to create candidate records: %s",
|
||||
_sanitize(str(exc)),
|
||||
)
|
||||
record_map = {}
|
||||
|
||||
candidate_keys: list[dict[str, Any]] = []
|
||||
eligible_count = 0
|
||||
last_status_code: int | None = None
|
||||
|
||||
for idx, cand in enumerate(candidates):
|
||||
now = datetime.now(timezone.utc)
|
||||
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
|
||||
|
||||
candidate_info: dict[str, Any] = {
|
||||
"index": idx,
|
||||
"provider_id": cand.provider.id,
|
||||
"provider_name": cand.provider.name,
|
||||
"endpoint_id": cand.endpoint.id,
|
||||
"key_id": cand.key.id,
|
||||
"key_name": getattr(cand.key, "name", None),
|
||||
"auth_type": auth_type,
|
||||
"priority": getattr(cand.key, "priority", 0) or 0,
|
||||
"is_cached": bool(getattr(cand, "is_cached", False)),
|
||||
}
|
||||
candidate_keys.append(candidate_info)
|
||||
|
||||
record_id = record_map.get((idx, 0))
|
||||
|
||||
# Scheduler marked skip
|
||||
if getattr(cand, "is_skipped", False):
|
||||
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
|
||||
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||
if record_id:
|
||||
# record is usually already skipped, but keep it consistent
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="skipped", skip_reason=skip_reason)
|
||||
)
|
||||
continue
|
||||
|
||||
# Format conversion checks
|
||||
# 优先级:全局开关 ON 强制允许,全局开关 OFF 看提供商开关
|
||||
needs_conversion = bool(getattr(cand, "needs_conversion", False))
|
||||
if needs_conversion:
|
||||
# 1. Check handler-level switch (handler 不支持则直接跳过)
|
||||
if not allow_format_conversion:
|
||||
skip_reason = "format_conversion_not_supported"
|
||||
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="skipped", skip_reason=skip_reason)
|
||||
)
|
||||
continue
|
||||
|
||||
# 2. Check global + provider switches
|
||||
# 全局 ON → 允许;全局 OFF → 看提供商
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
global_enabled = SystemConfigService.is_format_conversion_enabled(self.db)
|
||||
provider_enabled = getattr(cand.provider, "enable_format_conversion", True)
|
||||
effective_enabled = global_enabled or provider_enabled
|
||||
|
||||
if not effective_enabled:
|
||||
skip_reason = "format_conversion_disabled"
|
||||
candidate_info.update(
|
||||
{
|
||||
"skipped": True,
|
||||
"skip_reason": skip_reason,
|
||||
"global_conversion_enabled": global_enabled,
|
||||
"provider_conversion_enabled": provider_enabled,
|
||||
}
|
||||
)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="skipped", skip_reason=skip_reason)
|
||||
)
|
||||
continue
|
||||
|
||||
# auth_type filter
|
||||
if supported_auth_types is not None and auth_type not in supported_auth_types:
|
||||
skip_reason = f"unsupported_auth_type:{auth_type}"
|
||||
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="skipped", skip_reason=skip_reason)
|
||||
)
|
||||
continue
|
||||
|
||||
# billing rule filter
|
||||
rule_lookup: BillingRuleLookupResult | None = None
|
||||
has_billing_rule = True
|
||||
if config.billing_require_rule:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=cand.provider.id,
|
||||
model_name=model_name,
|
||||
task_type=task_type,
|
||||
)
|
||||
has_billing_rule = rule_lookup is not None
|
||||
if not has_billing_rule:
|
||||
skip_reason = "billing_rule_missing"
|
||||
candidate_info.update(
|
||||
{"has_billing_rule": False, "skipped": True, "skip_reason": skip_reason}
|
||||
)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="skipped", skip_reason=skip_reason)
|
||||
)
|
||||
continue
|
||||
candidate_info["has_billing_rule"] = has_billing_rule
|
||||
|
||||
eligible_count += 1
|
||||
|
||||
# Mark pending
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(status="pending", started_at=now)
|
||||
)
|
||||
|
||||
# Flush/commit BEFORE awaiting upstream submit to avoid holding DB connections
|
||||
# during potentially slow network operations.
|
||||
if self.db.in_transaction():
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
|
||||
# Attempt submit (upstream HTTP)
|
||||
try:
|
||||
response: httpx.Response = await submit_func(cand)
|
||||
except Exception as exc:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
error_msg = _sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "exception",
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(
|
||||
status="failed",
|
||||
error_type=type(exc).__name__,
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
last_status_code = int(getattr(response, "status_code", 0) or 0)
|
||||
|
||||
if response.status_code >= 400:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
try:
|
||||
error_text = response.text or ""
|
||||
except Exception:
|
||||
error_text = ""
|
||||
error_msg = _sanitize(error_text)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "http_error",
|
||||
"status_code": response.status_code,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="http_error",
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
)
|
||||
|
||||
if self._should_stop_on_http_error(
|
||||
status_code=response.status_code, error_text=error_text
|
||||
):
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise UpstreamClientRequestError(
|
||||
response=response,
|
||||
candidate_keys=candidate_keys,
|
||||
)
|
||||
continue
|
||||
|
||||
# Parse JSON
|
||||
payload: dict[str, Any] | None = None
|
||||
try:
|
||||
data = response.json()
|
||||
if isinstance(data, dict):
|
||||
payload = data
|
||||
except Exception as exc:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
error_msg = _sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "invalid_json",
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="invalid_json",
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
external_task_id = extract_external_task_id(payload or {})
|
||||
if not external_task_id:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "empty_task_id",
|
||||
"error_message": "Upstream returned empty task id",
|
||||
}
|
||||
)
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="empty_task_id",
|
||||
error_message="Upstream returned empty task id",
|
||||
finished_at=finished_at,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Success
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
candidate_info.update({"attempt_status": "success", "selected": True})
|
||||
if record_id:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.values(
|
||||
status="success",
|
||||
status_code=response.status_code,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
)
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
|
||||
return SubmitOutcome(
|
||||
candidate=cand,
|
||||
candidate_keys=candidate_keys,
|
||||
external_task_id=str(external_task_id),
|
||||
rule_lookup=rule_lookup,
|
||||
upstream_payload=payload,
|
||||
upstream_headers=dict(response.headers),
|
||||
upstream_status_code=response.status_code,
|
||||
)
|
||||
|
||||
# Persist candidate records before raising
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
|
||||
if eligible_count == 0:
|
||||
reason = "no_eligible_candidates"
|
||||
if config.billing_require_rule:
|
||||
reason = "no_candidate_with_billing_rule"
|
||||
raise AllCandidatesFailedError(
|
||||
reason=reason,
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
|
||||
raise AllCandidatesFailedError(
|
||||
reason="all_candidates_failed",
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
finally:
|
||||
# Restore Session behavior for the rest of the request lifecycle.
|
||||
self.db.expire_on_commit = original_expire_on_commit
|
||||
|
||||
def create_candidate_records(
|
||||
self,
|
||||
*,
|
||||
candidates: list[ProviderCandidate],
|
||||
request_id: str,
|
||||
user_api_key: ApiKey,
|
||||
required_capabilities: dict[str, bool] | None = None,
|
||||
expand_retries: bool = True,
|
||||
) -> dict[tuple[int, int], str]:
|
||||
# CandidateResolver.create_candidate_records is currently the canonical implementation
|
||||
assert self._resolver is not None, "Call resolve() once before create_candidate_records()"
|
||||
return self._resolver.create_candidate_records(
|
||||
all_candidates=candidates,
|
||||
request_id=request_id,
|
||||
user_id=str(user_api_key.user_id),
|
||||
user_api_key=user_api_key,
|
||||
required_capabilities=required_capabilities,
|
||||
expand_retries=expand_retries,
|
||||
)
|
||||
|
||||
def get_candidate_keys(self, request_id: str) -> list["CandidateKey"]:
|
||||
return self._recorder.get_candidate_keys(request_id)
|
||||
66
src/services/candidate/submit.py
Normal file
66
src/services/candidate/submit.py
Normal file
@@ -0,0 +1,66 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SubmitFunc(Protocol):
|
||||
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ExtractExternalTaskIdFunc(Protocol):
|
||||
def __call__(self, payload: dict[str, Any]) -> str | None: ...
|
||||
|
||||
|
||||
class UpstreamClientRequestError(RuntimeError):
|
||||
"""可判定为客户端请求问题(不应 failover)的上游错误。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
response: httpx.Response,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
) -> None:
|
||||
self.response = response
|
||||
self.candidate_keys = candidate_keys
|
||||
super().__init__(f"Upstream client error: HTTP {response.status_code}")
|
||||
|
||||
|
||||
class AllCandidatesFailedError(RuntimeError):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
reason: str,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
last_status_code: int | None = None,
|
||||
) -> None:
|
||||
self.reason = reason
|
||||
self.candidate_keys = candidate_keys
|
||||
self.last_status_code = last_status_code
|
||||
super().__init__(f"All candidates failed: {reason}")
|
||||
|
||||
|
||||
class CandidateUnsupportedError(RuntimeError):
|
||||
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
|
||||
|
||||
|
||||
class CandidateSubmissionError(RuntimeError):
|
||||
"""候选提交异常(网络/解密/解析等)。"""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SubmitOutcome:
|
||||
candidate: ProviderCandidate
|
||||
candidate_keys: list[dict[str, Any]]
|
||||
external_task_id: str
|
||||
rule_lookup: BillingRuleLookupResult | None
|
||||
upstream_payload: dict[str, Any] | None = None
|
||||
upstream_headers: dict[str, str] | None = None
|
||||
upstream_status_code: int | None = None
|
||||
@@ -1,21 +1,41 @@
|
||||
"""
|
||||
Gemini Files API - 文件与 Key 绑定缓存
|
||||
Gemini Files API - 文件与 Key 绑定映射服务
|
||||
|
||||
用于在上传文件后记录 file_id -> provider_key_id,
|
||||
并在后续 generateContent 请求中优先使用同一 Key。
|
||||
|
||||
存储策略:
|
||||
- 数据库(持久化):主存储,支持服务重启后恢复
|
||||
- Redis(缓存):加速读取,TTL=48小时
|
||||
|
||||
读取策略:
|
||||
1. 先查 Redis 缓存
|
||||
2. 缓存未命中时回查数据库
|
||||
3. 从数据库读取后回填缓存
|
||||
|
||||
清理策略:
|
||||
- 数据库中 expires_at 过期的记录由定时任务清理
|
||||
- Redis 缓存由 TTL 自动过期
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, Optional, Set
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import delete
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.cache_service import CacheService
|
||||
from src.core.logger import logger
|
||||
|
||||
FILE_MAPPING_TTL_SECONDS = 60 * 60 * 48 # 48小时
|
||||
FILE_MAPPING_CACHE_PREFIX = "gemini_files:key"
|
||||
|
||||
|
||||
def _normalize_file_name(file_name: str) -> str:
|
||||
"""规范化文件名,确保以 files/ 开头"""
|
||||
name = (file_name or "").strip()
|
||||
if not name:
|
||||
return ""
|
||||
@@ -23,31 +43,294 @@ def _normalize_file_name(file_name: str) -> str:
|
||||
|
||||
|
||||
def build_file_mapping_key(file_name: str) -> str:
|
||||
"""构建 Redis 缓存键"""
|
||||
normalized = _normalize_file_name(file_name)
|
||||
return f"{FILE_MAPPING_CACHE_PREFIX}:{normalized}" if normalized else ""
|
||||
|
||||
|
||||
async def store_file_key_mapping(file_name: str, key_id: str) -> None:
|
||||
cache_key = build_file_mapping_key(file_name)
|
||||
if not cache_key or not key_id:
|
||||
# =============================================================================
|
||||
# 异步接口(用于请求处理流程)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def store_file_key_mapping(
|
||||
file_name: str,
|
||||
key_id: str,
|
||||
user_id: str | None = None,
|
||||
display_name: str | None = None,
|
||||
mime_type: str | None = None,
|
||||
source_hash: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
存储文件→Key 映射(同时写入 Redis 和数据库)
|
||||
|
||||
Args:
|
||||
file_name: 文件名(如 files/abc123)
|
||||
key_id: Provider Key ID
|
||||
user_id: 用户 ID(可选,用于权限验证)
|
||||
display_name: 文件显示名(可选)
|
||||
mime_type: 文件 MIME 类型(可选)
|
||||
source_hash: 源文件哈希(可选,用于关联相同源文件的不同上传)
|
||||
"""
|
||||
normalized_name = _normalize_file_name(file_name)
|
||||
if not normalized_name or not key_id:
|
||||
return
|
||||
|
||||
# 1. 写入 Redis 缓存
|
||||
cache_key = build_file_mapping_key(normalized_name)
|
||||
await CacheService.set(cache_key, str(key_id), ttl_seconds=FILE_MAPPING_TTL_SECONDS)
|
||||
|
||||
# 2. 写入数据库(异步执行,不阻塞主流程)
|
||||
try:
|
||||
await _store_to_database(
|
||||
file_name=normalized_name,
|
||||
key_id=key_id,
|
||||
user_id=user_id,
|
||||
display_name=display_name,
|
||||
mime_type=mime_type,
|
||||
source_hash=source_hash,
|
||||
)
|
||||
except Exception as e:
|
||||
# 数据库写入失败只记录警告,不影响主流程
|
||||
logger.warning(f"Failed to persist Gemini file mapping to database: {e}")
|
||||
|
||||
|
||||
async def _store_to_database(
|
||||
file_name: str,
|
||||
key_id: str,
|
||||
user_id: str | None = None,
|
||||
display_name: str | None = None,
|
||||
mime_type: str | None = None,
|
||||
source_hash: str | None = None,
|
||||
) -> None:
|
||||
"""将映射写入数据库"""
|
||||
from src.database import get_db_context
|
||||
from src.models.database import GeminiFileMapping
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
expires_at = now + timedelta(hours=48)
|
||||
|
||||
with get_db_context() as db:
|
||||
# 使用 upsert 逻辑:存在则更新,不存在则插入
|
||||
existing = (
|
||||
db.query(GeminiFileMapping).filter(GeminiFileMapping.file_name == file_name).first()
|
||||
)
|
||||
|
||||
if existing:
|
||||
# 更新现有记录
|
||||
existing.key_id = key_id
|
||||
existing.user_id = user_id
|
||||
existing.display_name = display_name
|
||||
existing.mime_type = mime_type
|
||||
existing.source_hash = source_hash
|
||||
existing.expires_at = expires_at
|
||||
else:
|
||||
# 插入新记录
|
||||
mapping = GeminiFileMapping(
|
||||
id=str(uuid.uuid4()),
|
||||
file_name=file_name,
|
||||
key_id=key_id,
|
||||
user_id=user_id,
|
||||
display_name=display_name,
|
||||
mime_type=mime_type,
|
||||
source_hash=source_hash,
|
||||
created_at=now,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
db.add(mapping)
|
||||
|
||||
db.commit()
|
||||
|
||||
|
||||
async def get_file_key_mapping(file_name: str) -> str | None:
|
||||
cache_key = build_file_mapping_key(file_name)
|
||||
if not cache_key:
|
||||
"""
|
||||
获取文件→Key 映射
|
||||
|
||||
读取策略:
|
||||
1. 先查 Redis 缓存
|
||||
2. 缓存未命中时回查数据库
|
||||
3. 从数据库读取后回填缓存
|
||||
|
||||
Args:
|
||||
file_name: 文件名(如 files/abc123)
|
||||
|
||||
Returns:
|
||||
Provider Key ID,如果不存在或已过期则返回 None
|
||||
"""
|
||||
normalized_name = _normalize_file_name(file_name)
|
||||
if not normalized_name:
|
||||
return None
|
||||
value = await CacheService.get(cache_key)
|
||||
if value:
|
||||
return str(value)
|
||||
|
||||
cache_key = build_file_mapping_key(normalized_name)
|
||||
|
||||
# 1. 先查 Redis 缓存
|
||||
cached_value = await CacheService.get(cache_key)
|
||||
if cached_value:
|
||||
return str(cached_value)
|
||||
|
||||
# 2. 缓存未命中,回查数据库
|
||||
key_id = await _get_from_database(normalized_name)
|
||||
|
||||
if key_id:
|
||||
# 3. 回填缓存(使用剩余有效期或默认 TTL)
|
||||
await CacheService.set(cache_key, key_id, ttl_seconds=FILE_MAPPING_TTL_SECONDS)
|
||||
logger.debug(f"Gemini file mapping cache refilled from database: {normalized_name}")
|
||||
|
||||
return key_id
|
||||
|
||||
|
||||
async def _get_from_database(file_name: str) -> str | None:
|
||||
"""从数据库查询映射"""
|
||||
from src.database import get_db_context
|
||||
from src.models.database import GeminiFileMapping
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
with get_db_context() as db:
|
||||
mapping = (
|
||||
db.query(GeminiFileMapping)
|
||||
.filter(
|
||||
GeminiFileMapping.file_name == file_name,
|
||||
GeminiFileMapping.expires_at > now, # 只返回未过期的
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if mapping:
|
||||
return str(mapping.key_id)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to query Gemini file mapping from database: {e}")
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def get_all_key_ids_for_file(file_name: str) -> list[str]:
|
||||
"""
|
||||
获取支持指定文件的所有 Key ID 列表
|
||||
|
||||
当同一个源文件被上传到多个 Key 时,返回所有可用的 Key ID。
|
||||
这允许系统在首选 Key 不可用时选择其他 Key。
|
||||
|
||||
Args:
|
||||
file_name: 文件名(如 files/abc123)
|
||||
|
||||
Returns:
|
||||
所有支持该文件的 Key ID 列表(包括原始映射和具有相同 source_hash 的映射)
|
||||
"""
|
||||
from src.database import get_db_context
|
||||
from src.models.database import GeminiFileMapping
|
||||
|
||||
normalized_name = _normalize_file_name(file_name)
|
||||
if not normalized_name:
|
||||
return []
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
with get_db_context() as db:
|
||||
# 首先获取原始映射
|
||||
original_mapping = (
|
||||
db.query(GeminiFileMapping)
|
||||
.filter(
|
||||
GeminiFileMapping.file_name == normalized_name,
|
||||
GeminiFileMapping.expires_at > now,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if not original_mapping:
|
||||
return []
|
||||
|
||||
key_ids = [str(original_mapping.key_id)]
|
||||
|
||||
# 如果有 source_hash,查找所有具有相同 source_hash 的映射
|
||||
if original_mapping.source_hash:
|
||||
related_mappings = (
|
||||
db.query(GeminiFileMapping)
|
||||
.filter(
|
||||
GeminiFileMapping.source_hash == original_mapping.source_hash,
|
||||
GeminiFileMapping.expires_at > now,
|
||||
GeminiFileMapping.file_name != normalized_name, # 排除原始映射
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
for mapping in related_mappings:
|
||||
kid = str(mapping.key_id)
|
||||
if kid not in key_ids:
|
||||
key_ids.append(kid)
|
||||
|
||||
return key_ids
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to query related Gemini file mappings: {e}")
|
||||
return []
|
||||
|
||||
|
||||
async def delete_file_key_mapping(file_name: str) -> None:
|
||||
cache_key = build_file_mapping_key(file_name)
|
||||
if cache_key:
|
||||
await CacheService.delete(cache_key)
|
||||
"""
|
||||
删除文件→Key 映射(同时从 Redis 和数据库删除)
|
||||
|
||||
Args:
|
||||
file_name: 文件名(如 files/abc123)
|
||||
"""
|
||||
normalized_name = _normalize_file_name(file_name)
|
||||
if not normalized_name:
|
||||
return
|
||||
|
||||
# 1. 从 Redis 删除
|
||||
cache_key = build_file_mapping_key(normalized_name)
|
||||
await CacheService.delete(cache_key)
|
||||
|
||||
# 2. 从数据库删除
|
||||
try:
|
||||
await _delete_from_database(normalized_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to delete Gemini file mapping from database: {e}")
|
||||
|
||||
|
||||
async def _delete_from_database(file_name: str) -> None:
|
||||
"""从数据库删除映射"""
|
||||
from src.database import get_db_context
|
||||
from src.models.database import GeminiFileMapping
|
||||
|
||||
with get_db_context() as db:
|
||||
db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.file_name == file_name))
|
||||
db.commit()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# 同步接口(用于定时任务等场景)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def cleanup_expired_mappings(db: Session) -> int:
|
||||
"""
|
||||
清理过期的文件映射记录(同步方法,供定时任务调用)
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
|
||||
Returns:
|
||||
删除的记录数
|
||||
"""
|
||||
from src.models.database import GeminiFileMapping
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
|
||||
db.commit()
|
||||
|
||||
deleted_count = result.rowcount
|
||||
if deleted_count > 0:
|
||||
logger.info(f"Cleaned up {deleted_count} expired Gemini file mappings")
|
||||
|
||||
return deleted_count
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# 请求解析工具函数
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def _extract_file_name_from_uri(file_uri: str) -> str | None:
|
||||
|
||||
@@ -39,6 +39,10 @@ MAX_CONCURRENT_REQUESTS = 5
|
||||
# 单个 Key 处理的超时时间(秒)
|
||||
KEY_FETCH_TIMEOUT_SECONDS = 120
|
||||
|
||||
# 模型获取 HTTP 请求超时时间(秒)
|
||||
# 使用较短的超时(10秒),避免不支持 /models 端点的提供商长时间阻塞
|
||||
MODEL_FETCH_HTTP_TIMEOUT = 10.0
|
||||
|
||||
# 上游模型缓存 TTL(与定时任务间隔保持一致)
|
||||
UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60
|
||||
|
||||
@@ -260,7 +264,48 @@ class ModelFetchScheduler:
|
||||
logger.exception(f"更新 Key {key_id} 错误信息失败")
|
||||
|
||||
async def _fetch_models_for_key_by_id(self, key_id: str) -> str:
|
||||
"""根据 Key ID 获取模型并更新,返回结果状态"""
|
||||
"""
|
||||
根据 Key ID 获取模型并更新,返回结果状态
|
||||
|
||||
优化:分两个阶段处理,HTTP 请求期间不持有数据库连接,避免阻塞其他请求
|
||||
"""
|
||||
# ========== 阶段 1:准备数据(短暂持有连接)==========
|
||||
fetch_context = self._prepare_fetch_context(key_id)
|
||||
if fetch_context is None:
|
||||
return "skip"
|
||||
if isinstance(fetch_context, str):
|
||||
return fetch_context # "error" or "skip"
|
||||
|
||||
key_id, provider_id, provider_name, api_key_value, endpoint_configs = fetch_context
|
||||
|
||||
# ========== 阶段 2:HTTP 请求(不持有数据库连接)==========
|
||||
# 使用较短的超时时间(10秒),避免长时间阻塞
|
||||
all_models, errors, has_success = await fetch_models_from_endpoints(
|
||||
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
|
||||
)
|
||||
|
||||
# ========== 阶段 3:更新数据库(获取新连接)==========
|
||||
return await self._update_key_after_fetch(
|
||||
key_id=key_id,
|
||||
provider_id=provider_id,
|
||||
provider_name=provider_name,
|
||||
all_models=all_models,
|
||||
errors=errors,
|
||||
has_success=has_success,
|
||||
)
|
||||
|
||||
def _prepare_fetch_context(
|
||||
self, key_id: str
|
||||
) -> tuple[str, str, str, str, list[dict]] | str | None:
|
||||
"""
|
||||
准备获取模型所需的上下文数据
|
||||
|
||||
Returns:
|
||||
- tuple: (key_id, provider_id, provider_name, api_key_value, endpoint_configs)
|
||||
- "skip": 跳过该 Key
|
||||
- "error": 出错
|
||||
- None: Key 不存在
|
||||
"""
|
||||
with create_session() as db:
|
||||
key = (
|
||||
db.query(ProviderAPIKey)
|
||||
@@ -271,142 +316,154 @@ class ModelFetchScheduler:
|
||||
|
||||
if not key:
|
||||
logger.warning(f"Key {key_id} 不存在,跳过")
|
||||
return "skip"
|
||||
return None
|
||||
|
||||
if not key.is_active or not key.auto_fetch_models:
|
||||
logger.debug(f"Key {key_id} 已禁用或关闭自动获取,跳过")
|
||||
return "skip"
|
||||
|
||||
try:
|
||||
result = await self._fetch_models_for_key(db, key)
|
||||
now = datetime.now(timezone.utc)
|
||||
provider_id = key.provider_id
|
||||
|
||||
# 获取 Provider 和 Endpoints
|
||||
provider = (
|
||||
db.query(Provider)
|
||||
.options(joinedload(Provider.endpoints))
|
||||
.filter(Provider.id == provider_id)
|
||||
.first()
|
||||
)
|
||||
|
||||
if not provider:
|
||||
logger.warning(f"Provider {provider_id} 不存在,跳过 Key {key.id}")
|
||||
key.last_models_fetch_error = "Provider not found"
|
||||
key.last_models_fetch_at = now
|
||||
db.commit()
|
||||
return result
|
||||
return "error"
|
||||
|
||||
# Vertex AI 类型不支持自动获取模型
|
||||
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
|
||||
if auth_type == "vertex_ai":
|
||||
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key"
|
||||
key.last_models_fetch_at = now
|
||||
db.commit()
|
||||
logger.info(f"Key {key.id} 为 Vertex AI 类型,跳过自动获取模型")
|
||||
return "skip"
|
||||
|
||||
# 解密 API Key
|
||||
if not key.api_key:
|
||||
logger.warning(f"Key {key.id} 没有 API Key,跳过")
|
||||
key.last_models_fetch_error = "No API key configured"
|
||||
key.last_models_fetch_at = now
|
||||
db.commit()
|
||||
return "error"
|
||||
|
||||
try:
|
||||
api_key_value = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.error(f"解密 Key {key.id} 失败")
|
||||
key.last_models_fetch_error = "Decrypt error"
|
||||
key.last_models_fetch_at = now
|
||||
db.commit()
|
||||
return "error"
|
||||
|
||||
async def _fetch_models_for_key(
|
||||
# 构建 api_format -> endpoint 映射
|
||||
format_to_endpoint: dict[str, object] = {}
|
||||
for endpoint in provider.endpoints: # type: ignore[attr-defined]
|
||||
if endpoint.is_active:
|
||||
format_to_endpoint[endpoint.api_format] = endpoint
|
||||
|
||||
if not format_to_endpoint:
|
||||
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
|
||||
key.last_models_fetch_error = "No active endpoints"
|
||||
key.last_models_fetch_at = now
|
||||
db.commit()
|
||||
return "error"
|
||||
|
||||
# 使用公共函数构建所有格式的端点配置
|
||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
|
||||
|
||||
return (key_id, provider_id, provider.name, api_key_value, endpoint_configs)
|
||||
|
||||
async def _update_key_after_fetch(
|
||||
self,
|
||||
db: Session,
|
||||
key: ProviderAPIKey,
|
||||
key_id: str,
|
||||
provider_id: str,
|
||||
provider_name: str,
|
||||
all_models: list[dict],
|
||||
errors: list[str],
|
||||
has_success: bool,
|
||||
) -> str:
|
||||
"""为单个 Key 获取模型并更新 allowed_models,返回结果状态"""
|
||||
"""
|
||||
HTTP 请求完成后更新数据库
|
||||
|
||||
使用新的数据库连接来更新 Key 的 allowed_models
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
provider_id = key.provider_id
|
||||
|
||||
# 获取 Provider 和 Endpoints
|
||||
provider = (
|
||||
db.query(Provider)
|
||||
.options(joinedload(Provider.endpoints))
|
||||
.filter(Provider.id == provider_id)
|
||||
.first()
|
||||
)
|
||||
with create_session() as db:
|
||||
# 重新获取 Key(因为之前的连接已关闭)
|
||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||
if not key:
|
||||
logger.warning(f"Key {key_id} 在更新时不存在")
|
||||
return "error"
|
||||
|
||||
if not provider:
|
||||
logger.warning(f"Provider {provider_id} 不存在,跳过 Key {key.id}")
|
||||
key.last_models_fetch_error = "Provider not found"
|
||||
# 记录获取时间
|
||||
key.last_models_fetch_at = now
|
||||
return "error"
|
||||
|
||||
# Vertex AI 类型不支持自动获取模型(需要使用 Service Account 认证)
|
||||
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
|
||||
if auth_type == "vertex_ai":
|
||||
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key"
|
||||
key.last_models_fetch_at = now
|
||||
logger.info(f"Key {key.id} 为 Vertex AI 类型,跳过自动获取模型")
|
||||
return "skip"
|
||||
# 如果没有任何成功的响应,不更新 allowed_models(保留旧数据)
|
||||
if not has_success:
|
||||
error_msg = "; ".join(errors) if errors else "All endpoints failed"
|
||||
key.last_models_fetch_error = error_msg
|
||||
logger.warning(
|
||||
f"Provider {provider_name} Key {key.id} 所有端点获取失败,保留现有模型列表"
|
||||
)
|
||||
db.commit()
|
||||
return "error"
|
||||
|
||||
# 解密 API Key
|
||||
if not key.api_key:
|
||||
logger.warning(f"Key {key.id} 没有 API Key,跳过")
|
||||
key.last_models_fetch_error = "No API key configured"
|
||||
key.last_models_fetch_at = now
|
||||
return "error"
|
||||
# 有成功的响应,清除错误状态
|
||||
key.last_models_fetch_error = None
|
||||
|
||||
try:
|
||||
api_key_value = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
# 不记录异常详情,避免泄露密钥信息
|
||||
logger.error(f"解密 Key {key.id} 失败")
|
||||
key.last_models_fetch_error = "Decrypt error"
|
||||
key.last_models_fetch_at = now
|
||||
return "error"
|
||||
# 去重获取模型 ID 列表
|
||||
fetched_model_ids: set[str] = set()
|
||||
for model in all_models:
|
||||
model_id = model.get("id")
|
||||
if model_id:
|
||||
fetched_model_ids.add(model_id)
|
||||
|
||||
# 构建 api_format -> endpoint 映射
|
||||
format_to_endpoint: dict[str, object] = {}
|
||||
for endpoint in provider.endpoints: # type: ignore[attr-defined]
|
||||
if endpoint.is_active:
|
||||
format_to_endpoint[endpoint.api_format] = endpoint
|
||||
|
||||
if not format_to_endpoint:
|
||||
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
|
||||
key.last_models_fetch_error = "No active endpoints"
|
||||
key.last_models_fetch_at = now
|
||||
return "error"
|
||||
|
||||
# 使用公共函数构建所有格式的端点配置
|
||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
|
||||
|
||||
# 并发获取模型
|
||||
all_models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
|
||||
|
||||
# 记录获取时间
|
||||
key.last_models_fetch_at = now
|
||||
|
||||
# 如果没有任何成功的响应,不更新 allowed_models(保留旧数据)
|
||||
if not has_success:
|
||||
# 所有端点都失败时,记录错误
|
||||
error_msg = "; ".join(errors) if errors else "All endpoints failed"
|
||||
key.last_models_fetch_error = error_msg
|
||||
logger.warning(
|
||||
f"Provider {provider.name} Key {key.id} 所有端点获取失败,保留现有模型列表"
|
||||
)
|
||||
return "error"
|
||||
|
||||
# 有成功的响应,清除错误状态(部分失败不算失败)
|
||||
key.last_models_fetch_error = None
|
||||
|
||||
# 去重获取模型 ID 列表
|
||||
fetched_model_ids: set[str] = set()
|
||||
for model in all_models:
|
||||
model_id = model.get("id")
|
||||
if model_id:
|
||||
fetched_model_ids.add(model_id)
|
||||
|
||||
logger.info(
|
||||
f"Provider {provider.name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
|
||||
)
|
||||
|
||||
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
|
||||
seen_keys: set[str] = set()
|
||||
unique_models: list[dict] = []
|
||||
for model in all_models:
|
||||
model_id = model.get("id")
|
||||
api_format = model.get("api_format", "")
|
||||
unique_key = f"{model_id}:{api_format}"
|
||||
if model_id and unique_key not in seen_keys:
|
||||
seen_keys.add(unique_key)
|
||||
unique_models.append(model)
|
||||
await set_upstream_models_to_cache(
|
||||
provider_id, # type: ignore[arg-type]
|
||||
key.id, # type: ignore[arg-type]
|
||||
unique_models,
|
||||
)
|
||||
|
||||
# 更新 allowed_models(保留 locked_models)
|
||||
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
|
||||
|
||||
# 如果白名单有变化,触发缓存失效和自动关联检查
|
||||
if has_changed and provider_id:
|
||||
from src.services.model.global_model import on_key_allowed_models_changed
|
||||
|
||||
await on_key_allowed_models_changed(
|
||||
db=db,
|
||||
provider_id=provider_id,
|
||||
allowed_models=list(key.allowed_models or []),
|
||||
logger.info(
|
||||
f"Provider {provider_name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
|
||||
)
|
||||
|
||||
return "success"
|
||||
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
|
||||
seen_keys: set[str] = set()
|
||||
unique_models: list[dict] = []
|
||||
for model in all_models:
|
||||
model_id = model.get("id")
|
||||
api_format = model.get("api_format", "")
|
||||
unique_key = f"{model_id}:{api_format}"
|
||||
if model_id and unique_key not in seen_keys:
|
||||
seen_keys.add(unique_key)
|
||||
unique_models.append(model)
|
||||
await set_upstream_models_to_cache(provider_id, key.id, unique_models)
|
||||
|
||||
# 更新 allowed_models(保留 locked_models)
|
||||
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
|
||||
|
||||
db.commit()
|
||||
|
||||
# 如果白名单有变化,触发缓存失效和自动关联检查
|
||||
if has_changed and provider_id:
|
||||
from src.services.model.global_model import on_key_allowed_models_changed
|
||||
|
||||
# 使用新会话处理后续操作
|
||||
with create_session() as db2:
|
||||
await on_key_allowed_models_changed(
|
||||
db=db2,
|
||||
provider_id=provider_id,
|
||||
allowed_models=list(key.allowed_models or []),
|
||||
)
|
||||
|
||||
return "success"
|
||||
|
||||
def _update_key_allowed_models(self, key: ProviderAPIKey, fetched_model_ids: set[str]) -> bool:
|
||||
"""
|
||||
|
||||
@@ -156,6 +156,8 @@ class CandidateResolver:
|
||||
user_id: str,
|
||||
user_api_key: ApiKey,
|
||||
required_capabilities: dict[str, bool] | None = None,
|
||||
*,
|
||||
expand_retries: bool = True,
|
||||
) -> dict[tuple[int, int], str]:
|
||||
"""
|
||||
为所有候选预先创建 available 状态记录(批量插入优化)
|
||||
@@ -211,9 +213,12 @@ class CandidateResolver:
|
||||
candidate_record_map[(candidate_index, 0)] = record_id
|
||||
else:
|
||||
# max_retries 已从 Endpoint 迁移到 Provider(Endpoint 仍可能保留旧字段用于兼容)
|
||||
max_retries_for_candidate = (
|
||||
int(provider.max_retries or 2) if candidate.is_cached else 1
|
||||
)
|
||||
if not expand_retries:
|
||||
max_retries_for_candidate = 1
|
||||
else:
|
||||
max_retries_for_candidate = (
|
||||
int(provider.max_retries or 2) if candidate.is_cached else 1
|
||||
)
|
||||
|
||||
for retry_index in range(max_retries_for_candidate):
|
||||
record_id = str(uuid.uuid4())
|
||||
|
||||
@@ -38,6 +38,19 @@ BALANCE_CACHE_TTL = 86400
|
||||
# 认证失败缓存 TTL(60 秒,避免频繁重试但允许用户修正后快速重试)
|
||||
AUTH_FAILED_CACHE_TTL = 60
|
||||
|
||||
# 后台余额刷新并发限制(避免启动时耗尽连接池)
|
||||
# 使用较小的值(3)确保不会对连接池造成过大压力
|
||||
_balance_refresh_semaphore: asyncio.Semaphore | None = None
|
||||
|
||||
|
||||
def _get_balance_refresh_semaphore() -> asyncio.Semaphore:
|
||||
"""获取余额刷新信号量(延迟初始化)"""
|
||||
global _balance_refresh_semaphore
|
||||
if _balance_refresh_semaphore is None:
|
||||
# 限制为 3 个并发,确保后台任务不会占用太多连接
|
||||
_balance_refresh_semaphore = asyncio.Semaphore(3)
|
||||
return _balance_refresh_semaphore
|
||||
|
||||
|
||||
def _get_batch_balance_concurrency() -> int:
|
||||
"""
|
||||
@@ -98,6 +111,44 @@ class ProviderOpsService:
|
||||
# 连接器缓存 {provider_id: ProviderConnector}
|
||||
self._connectors: dict[str, ProviderConnector] = {}
|
||||
|
||||
def _release_db_connection_before_await(self) -> None:
|
||||
"""
|
||||
Release pooled DB connection before long awaits (network/Redis).
|
||||
|
||||
SQLAlchemy Session will keep a connection checked out while a transaction is open,
|
||||
even for read-only queries. In async code, this can exhaust the pool if we `await`
|
||||
network I/O while holding that transaction.
|
||||
|
||||
Safety:
|
||||
- Only commits when the session has no pending changes (new/dirty/deleted).
|
||||
- Temporarily disables expire_on_commit to avoid unexpected lazy reloads.
|
||||
"""
|
||||
try:
|
||||
has_pending_changes = bool(self.db.new) or bool(self.db.dirty) or bool(self.db.deleted)
|
||||
except Exception:
|
||||
has_pending_changes = False
|
||||
|
||||
if has_pending_changes:
|
||||
return
|
||||
|
||||
try:
|
||||
if not self.db.in_transaction():
|
||||
return
|
||||
except Exception:
|
||||
return
|
||||
|
||||
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
|
||||
self.db.expire_on_commit = False
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
self.db.expire_on_commit = original_expire_on_commit
|
||||
|
||||
# ==================== 配置管理 ====================
|
||||
|
||||
def get_config(self, provider_id: str) -> ProviderOpsConfig | None:
|
||||
@@ -251,6 +302,9 @@ class ProviderOpsService:
|
||||
if not actual_credentials:
|
||||
return False, "未提供凭据"
|
||||
|
||||
# Avoid holding a DB connection while awaiting network I/O.
|
||||
self._release_db_connection_before_await()
|
||||
|
||||
# 建立连接
|
||||
logger.info(
|
||||
f"尝试连接: provider_id={provider_id}, "
|
||||
@@ -333,6 +387,8 @@ class ProviderOpsService:
|
||||
)
|
||||
connector = self._connectors.get(provider_id)
|
||||
|
||||
# Avoid holding a DB connection while awaiting authentication checks.
|
||||
self._release_db_connection_before_await()
|
||||
if not connector or not await connector.is_authenticated():
|
||||
return ActionResult(
|
||||
status=ActionStatus.AUTH_EXPIRED,
|
||||
@@ -373,6 +429,9 @@ class ProviderOpsService:
|
||||
# 创建操作实例
|
||||
action = architecture.get_action(action_type, merged_config)
|
||||
|
||||
# Avoid holding a DB connection while awaiting the upstream action.
|
||||
self._release_db_connection_before_await()
|
||||
|
||||
# 执行操作
|
||||
async with connector.get_client() as client:
|
||||
result = await action.execute(client)
|
||||
@@ -422,6 +481,9 @@ class ProviderOpsService:
|
||||
Returns:
|
||||
操作结果(可能是缓存的)
|
||||
"""
|
||||
# Avoid holding a DB connection while awaiting Redis/cache I/O.
|
||||
self._release_db_connection_before_await()
|
||||
|
||||
# 尝试从缓存获取
|
||||
cached = await self._get_cached_balance(provider_id)
|
||||
|
||||
@@ -453,7 +515,20 @@ class ProviderOpsService:
|
||||
|
||||
注意:这是一个后台任务,使用独立的短生命周期 session,
|
||||
避免长时间占用连接池资源。
|
||||
|
||||
使用信号量限制并发数,避免启动时多个刷新任务同时运行导致连接池耗尽。
|
||||
"""
|
||||
semaphore = _get_balance_refresh_semaphore()
|
||||
|
||||
# 尝试获取信号量,如果无法立即获取则跳过本次刷新
|
||||
# 这样可以避免在连接池紧张时阻塞
|
||||
try:
|
||||
# 使用 wait_for 设置超时,避免无限等待
|
||||
await asyncio.wait_for(semaphore.acquire(), timeout=5.0)
|
||||
except asyncio.TimeoutError:
|
||||
logger.debug(f"异步刷新余额跳过(并发限制): provider_id={provider_id}")
|
||||
return
|
||||
|
||||
db = None
|
||||
try:
|
||||
# 后台任务需要创建独立的 session,因为原请求的 session 可能已关闭
|
||||
@@ -469,6 +544,8 @@ class ProviderOpsService:
|
||||
db.close()
|
||||
except Exception:
|
||||
pass
|
||||
# 释放信号量
|
||||
semaphore.release()
|
||||
|
||||
async def _clear_balance_cache(self, provider_id: str) -> None:
|
||||
"""清除余额缓存"""
|
||||
@@ -762,6 +839,9 @@ class ProviderOpsService:
|
||||
if not provider_ids:
|
||||
return {}
|
||||
|
||||
# Release the DB connection before awaiting many async cache refreshes.
|
||||
self._release_db_connection_before_await()
|
||||
|
||||
# 使用信号量限制并发数,避免同时发起过多请求耗尽连接池
|
||||
concurrency = _get_batch_balance_concurrency()
|
||||
semaphore = asyncio.Semaphore(concurrency)
|
||||
@@ -822,6 +902,9 @@ class ProviderOpsService:
|
||||
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
|
||||
# Avoid holding a DB connection while awaiting verify pre-processing / network.
|
||||
self._release_db_connection_before_await()
|
||||
|
||||
# 移除 base_url 末尾的斜杠
|
||||
base_url = base_url.rstrip("/")
|
||||
|
||||
|
||||
@@ -136,6 +136,11 @@ class SystemConfigService:
|
||||
"value": [],
|
||||
"description": "邮箱后缀列表,配合 email_suffix_mode 使用",
|
||||
},
|
||||
# 格式转换开关
|
||||
"enable_format_conversion": {
|
||||
"value": False,
|
||||
"description": "全局格式转换开关:开启时强制允许所有提供商的格式转换;关闭时由各提供商自行决定",
|
||||
},
|
||||
"audit_log_retention_days": {
|
||||
"value": 30,
|
||||
"description": "审计日志保留天数,超过此天数的审计日志将被自动清理",
|
||||
@@ -358,6 +363,11 @@ class SystemConfigService:
|
||||
"""获取敏感请求头列表"""
|
||||
return cls.get_config(db, "sensitive_headers", [])
|
||||
|
||||
@classmethod
|
||||
def is_format_conversion_enabled(cls, db: Session) -> bool:
|
||||
"""检查全局格式转换是否启用"""
|
||||
return bool(cls.get_config(db, "enable_format_conversion", True))
|
||||
|
||||
@classmethod
|
||||
def mask_sensitive_headers(cls, db: Session, headers: dict[str, Any]) -> dict[str, Any]:
|
||||
"""脱敏敏感请求头"""
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
- 审计日志清理:定期清理过期的审计日志
|
||||
- 连接池监控:定期检查数据库连接池状态
|
||||
- Pending 状态清理:清理异常的 Pending 状态记录
|
||||
- Gemini 文件映射清理:清理过期的 Gemini 文件→Key 映射
|
||||
|
||||
使用 APScheduler 进行任务调度,支持时区配置。
|
||||
"""
|
||||
@@ -103,6 +104,14 @@ class MaintenanceScheduler:
|
||||
name="审计日志清理",
|
||||
)
|
||||
|
||||
# Gemini 文件映射清理 - 每小时执行
|
||||
scheduler.add_interval_job(
|
||||
self._scheduled_gemini_file_mapping_cleanup,
|
||||
hours=1,
|
||||
job_id="gemini_file_mapping_cleanup",
|
||||
name="Gemini文件映射清理",
|
||||
)
|
||||
|
||||
# Provider 签到任务 - 凌晨 1:05 执行
|
||||
scheduler.add_cron_job(
|
||||
self._scheduled_provider_checkin,
|
||||
@@ -117,8 +126,9 @@ class MaintenanceScheduler:
|
||||
|
||||
async def _run_startup_tasks(self) -> None:
|
||||
"""启动时执行的初始化任务"""
|
||||
# 延迟一点执行,确保系统完全启动
|
||||
await asyncio.sleep(2)
|
||||
# 延迟执行,等待系统完全启动(Redis 连接、其他后台任务稳定)
|
||||
# 增加延迟时间避免与 UsageQueueConsumer 等后台任务竞争数据库连接
|
||||
await asyncio.sleep(10)
|
||||
|
||||
try:
|
||||
logger.info("启动时执行首次清理任务...")
|
||||
@@ -170,6 +180,10 @@ class MaintenanceScheduler:
|
||||
"""审计日志清理任务(定时调用)"""
|
||||
await self._perform_audit_cleanup()
|
||||
|
||||
async def _scheduled_gemini_file_mapping_cleanup(self) -> None:
|
||||
"""Gemini 文件映射清理任务(定时调用)"""
|
||||
await self._perform_gemini_file_mapping_cleanup()
|
||||
|
||||
async def _scheduled_provider_checkin(self) -> None:
|
||||
"""Provider 签到任务(定时调用)"""
|
||||
await self._perform_provider_checkin()
|
||||
@@ -483,6 +497,26 @@ class MaintenanceScheduler:
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def _perform_gemini_file_mapping_cleanup(self) -> None:
|
||||
"""清理过期的 Gemini 文件映射记录"""
|
||||
db = create_session()
|
||||
try:
|
||||
from src.services.gemini_files_mapping import cleanup_expired_mappings
|
||||
|
||||
deleted_count = cleanup_expired_mappings(db)
|
||||
|
||||
if deleted_count > 0:
|
||||
logger.info(f"清理了 {deleted_count} 条过期的 Gemini 文件映射")
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"Gemini 文件映射清理失败: {e}")
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def _perform_provider_checkin(self) -> None:
|
||||
"""执行 Provider 签到任务
|
||||
|
||||
@@ -508,8 +542,21 @@ class MaintenanceScheduler:
|
||||
|
||||
logger.info(f"开始执行 Provider 签到,共 {len(provider_ids)} 个...")
|
||||
|
||||
# 创建 ProviderOpsService 并执行批量余额查询(会触发签到)
|
||||
service = ProviderOpsService(db)
|
||||
# 释放主 session 的连接,避免在整个签到期间占用连接池
|
||||
# (后续每个 provider 将使用独立短生命周期 session)
|
||||
try:
|
||||
if db.in_transaction():
|
||||
db.commit()
|
||||
except Exception:
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
db.close()
|
||||
except Exception:
|
||||
pass
|
||||
db = None
|
||||
|
||||
# 使用信号量限制并发,避免同时发起过多请求
|
||||
concurrency = 3 # 签到任务并发数
|
||||
@@ -518,7 +565,9 @@ class MaintenanceScheduler:
|
||||
async def _checkin_provider(provider_id: str) -> tuple[str, bool, str]:
|
||||
"""执行单个 Provider 的签到"""
|
||||
async with semaphore:
|
||||
task_db = create_session()
|
||||
try:
|
||||
service = ProviderOpsService(task_db)
|
||||
# 触发余额查询(会先执行签到)
|
||||
result = await service.query_balance(provider_id)
|
||||
# 检查签到结果
|
||||
@@ -537,6 +586,11 @@ class MaintenanceScheduler:
|
||||
except Exception as e:
|
||||
logger.warning(f"Provider {provider_id} 签到失败: {e}")
|
||||
return provider_id, False, str(e)
|
||||
finally:
|
||||
try:
|
||||
task_db.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 并行执行签到
|
||||
tasks = [_checkin_provider(pid) for pid in provider_ids]
|
||||
@@ -556,7 +610,8 @@ class MaintenanceScheduler:
|
||||
except Exception as e:
|
||||
logger.exception(f"Provider 签到任务执行失败: {e}")
|
||||
finally:
|
||||
db.close()
|
||||
if db is not None:
|
||||
db.close()
|
||||
|
||||
async def _perform_cleanup(self) -> None:
|
||||
"""执行清理任务"""
|
||||
|
||||
@@ -1,25 +1,13 @@
|
||||
"""
|
||||
异步任务服务层
|
||||
任务服务层(Phase2)
|
||||
|
||||
提供视频/图片/音频等异步任务的:
|
||||
- 提交阶段故障转移(AsyncTaskOrchestrator)
|
||||
- 终态计费与 Usage 写入(VideoTelemetry 等)
|
||||
统一任务框架相关的应用层入口:
|
||||
- 候选提交阶段:`services.candidate.CandidateService`
|
||||
- 终态结算:`services.task.application.TaskApplicationService`
|
||||
"""
|
||||
|
||||
from .orchestrator import (
|
||||
AllCandidatesFailedError,
|
||||
AsyncTaskOrchestrator,
|
||||
CandidateSubmissionError,
|
||||
CandidateUnsupportedError,
|
||||
SubmitOutcome,
|
||||
UpstreamClientRequestError,
|
||||
)
|
||||
from .application import TaskApplicationService
|
||||
|
||||
__all__ = [
|
||||
"AsyncTaskOrchestrator",
|
||||
"SubmitOutcome",
|
||||
"AllCandidatesFailedError",
|
||||
"UpstreamClientRequestError",
|
||||
"CandidateUnsupportedError",
|
||||
"CandidateSubmissionError",
|
||||
"TaskApplicationService",
|
||||
]
|
||||
|
||||
312
src/services/task/application.py
Normal file
312
src/services/task/application.py
Normal file
@@ -0,0 +1,312 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Provider, Usage, User, VideoTask
|
||||
from src.services.billing.dimension_collector_service import DimensionCollectorService
|
||||
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
|
||||
from src.services.billing.rule_service import BillingRuleService
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
class TaskApplicationService:
|
||||
"""
|
||||
TaskApplicationService (Phase2)
|
||||
|
||||
当前仅先收敛"终态结算"入口,用于替代旧版 VideoTelemetry 直写 Usage 的流程。
|
||||
后续将扩展 submit/cancel 并迁移候选编排逻辑。
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
|
||||
async def finalize_video_task(self, task: VideoTask) -> bool:
|
||||
"""
|
||||
更新视频任务的计费信息(轮询完成后调用)。
|
||||
|
||||
异步任务的计费流程:
|
||||
1. 提交成功时:Usage 已结算(billing_status='settled',费用=0)
|
||||
2. 轮询完成时:更新实际费用(成功则计费,失败则保持0)
|
||||
|
||||
返回 True 表示成功更新,False 表示无需更新(如已是最终状态)
|
||||
"""
|
||||
request_id = getattr(task, "request_id", None) or task.id
|
||||
|
||||
existing = self.db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if not existing:
|
||||
# Usage 不存在,尝试创建并结算(兜底逻辑)
|
||||
logger.warning(
|
||||
"Usage not found for video task, creating fallback: task_id=%s request_id=%s",
|
||||
task.id,
|
||||
request_id,
|
||||
)
|
||||
return await self._create_fallback_usage(task, request_id)
|
||||
|
||||
# 检查是否已有计费更新标记(避免重复计费)
|
||||
metadata = existing.request_metadata or {}
|
||||
if metadata.get("billing_updated_at"):
|
||||
logger.debug(
|
||||
"Video task billing already updated: task_id=%s request_id=%s",
|
||||
task.id,
|
||||
request_id,
|
||||
)
|
||||
return False
|
||||
|
||||
# 计算异步任务总耗时(ms)
|
||||
response_time_ms: int | None = None
|
||||
if task.submitted_at and task.completed_at:
|
||||
delta = task.completed_at - task.submitted_at
|
||||
response_time_ms = int(delta.total_seconds() * 1000)
|
||||
|
||||
# === 收集计费维度 ===
|
||||
base_dimensions: dict[str, Any] = {
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size or "",
|
||||
"retry_count": task.retry_count,
|
||||
}
|
||||
|
||||
collector_metadata: dict[str, Any] = {
|
||||
"task": {
|
||||
"id": task.id,
|
||||
"external_task_id": task.external_task_id,
|
||||
"model": task.model,
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size,
|
||||
"retry_count": task.retry_count,
|
||||
"video_size_bytes": task.video_size_bytes,
|
||||
},
|
||||
"result": {
|
||||
"video_url": task.video_url,
|
||||
"video_urls": task.video_urls or [],
|
||||
},
|
||||
}
|
||||
|
||||
dims = DimensionCollectorService(self.db).collect_dimensions(
|
||||
api_format=task.provider_api_format,
|
||||
task_type="video",
|
||||
request=task.original_request_body or {},
|
||||
response=(
|
||||
(task.request_metadata or {}).get("poll_raw_response")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
metadata=collector_metadata,
|
||||
base_dimensions=base_dimensions,
|
||||
)
|
||||
|
||||
# === 计算成本(优先使用冻结的 billing_rule_snapshot)===
|
||||
rule_snapshot = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
|
||||
|
||||
expression = None
|
||||
variables: dict[str, Any] | None = None
|
||||
dimension_mappings: dict[str, dict[str, Any]] | None = None
|
||||
rule_id = None
|
||||
rule_name = None
|
||||
rule_scope = None
|
||||
|
||||
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
|
||||
rule_id = rule_snapshot.get("rule_id")
|
||||
rule_name = rule_snapshot.get("rule_name")
|
||||
rule_scope = rule_snapshot.get("scope")
|
||||
expression = rule_snapshot.get("expression")
|
||||
variables = rule_snapshot.get("variables") or {}
|
||||
dimension_mappings = rule_snapshot.get("dimension_mappings") or {}
|
||||
else:
|
||||
lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=task.provider_id,
|
||||
model_name=task.model,
|
||||
task_type="video",
|
||||
)
|
||||
if lookup:
|
||||
rule = lookup.rule
|
||||
rule_id = rule.id
|
||||
rule_name = rule.name
|
||||
rule_scope = getattr(lookup, "scope", None)
|
||||
expression = rule.expression
|
||||
variables = rule.variables or {}
|
||||
dimension_mappings = rule.dimension_mappings or {}
|
||||
|
||||
billing_snapshot: dict[str, Any] = {
|
||||
"schema_version": "1.0",
|
||||
"rule_id": str(rule_id) if rule_id else None,
|
||||
"rule_name": str(rule_name) if rule_name else None,
|
||||
"scope": str(rule_scope) if rule_scope else None,
|
||||
"expression": str(expression) if expression else None,
|
||||
"dimensions_used": dims,
|
||||
"missing_required": [],
|
||||
"cost": 0.0,
|
||||
"status": "no_rule",
|
||||
"calculated_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
|
||||
cost = 0.0
|
||||
# 只有任务成功时才计费
|
||||
if task.status == "completed" and expression:
|
||||
engine = FormulaEngine()
|
||||
try:
|
||||
result = engine.evaluate(
|
||||
expression=str(expression),
|
||||
variables=variables,
|
||||
dimensions=dims,
|
||||
dimension_mappings=dimension_mappings,
|
||||
strict_mode=config.billing_strict_mode,
|
||||
)
|
||||
billing_snapshot["status"] = result.status
|
||||
billing_snapshot["missing_required"] = result.missing_required
|
||||
if result.status == "complete":
|
||||
cost = float(result.cost)
|
||||
billing_snapshot["cost"] = cost
|
||||
except BillingIncompleteError as exc:
|
||||
# strict_mode=true:标记任务失败并隐藏产物,避免"免费放行"
|
||||
task.status = "failed"
|
||||
task.error_code = "billing_incomplete"
|
||||
task.error_message = f"Missing required dimensions: {exc.missing_required}"
|
||||
task.video_url = None
|
||||
task.video_urls = None
|
||||
billing_snapshot["status"] = "incomplete"
|
||||
billing_snapshot["missing_required"] = exc.missing_required
|
||||
billing_snapshot["cost"] = 0.0
|
||||
except Exception as exc:
|
||||
billing_snapshot["status"] = "incomplete"
|
||||
billing_snapshot["error"] = str(exc)
|
||||
billing_snapshot["cost"] = 0.0
|
||||
|
||||
# 回写到 task.request_metadata 便于审计/重算
|
||||
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
|
||||
metadata = dict(task.request_metadata) if task.request_metadata else {}
|
||||
metadata["billing_snapshot"] = billing_snapshot
|
||||
task.request_metadata = metadata
|
||||
|
||||
# === 更新已结算的 Usage 计费信息 ===
|
||||
updated = UsageService.update_settled_billing(
|
||||
self.db,
|
||||
request_id=request_id,
|
||||
total_cost_usd=cost,
|
||||
request_cost_usd=cost,
|
||||
status="completed" if task.status == "completed" else "failed",
|
||||
status_code=200 if task.status == "completed" else 500,
|
||||
error_message=(
|
||||
None
|
||||
if task.status == "completed"
|
||||
else (task.error_message or task.error_code or "video_task_failed")
|
||||
),
|
||||
response_time_ms=response_time_ms,
|
||||
billing_snapshot=billing_snapshot,
|
||||
extra_metadata={
|
||||
"dimensions": dims,
|
||||
"raw_response_ref": {
|
||||
"video_task_id": task.id,
|
||||
"field": "video_tasks.request_metadata.poll_raw_response",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
if updated:
|
||||
logger.debug(
|
||||
"Updated video task billing: task_id=%s request_id=%s cost=%.6f",
|
||||
task.id,
|
||||
request_id,
|
||||
cost,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Failed to update video task billing (may already be updated): "
|
||||
"task_id=%s request_id=%s",
|
||||
task.id,
|
||||
request_id,
|
||||
)
|
||||
|
||||
return updated
|
||||
|
||||
async def _create_fallback_usage(self, task: VideoTask, request_id: str) -> bool:
|
||||
"""
|
||||
兜底逻辑:当 Usage 不存在时创建完整记录。
|
||||
这种情况理论上不应发生(submit 阶段已创建),但保留以防万一。
|
||||
"""
|
||||
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
|
||||
api_key_obj = (
|
||||
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
|
||||
if task.api_key_id
|
||||
else None
|
||||
)
|
||||
provider_obj = (
|
||||
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
|
||||
if task.provider_id
|
||||
else None
|
||||
)
|
||||
provider_name = provider_obj.name if provider_obj else "unknown"
|
||||
|
||||
# 计算响应时间
|
||||
response_time_ms: int | None = None
|
||||
if task.submitted_at and task.completed_at:
|
||||
delta = task.completed_at - task.submitted_at
|
||||
response_time_ms = int(delta.total_seconds() * 1000)
|
||||
|
||||
try:
|
||||
await UsageService.record_usage_with_custom_cost(
|
||||
db=self.db,
|
||||
user=user_obj,
|
||||
api_key=api_key_obj,
|
||||
provider=provider_name,
|
||||
model=task.model,
|
||||
request_type="video",
|
||||
total_cost_usd=0.0, # 兜底记录不计费
|
||||
request_cost_usd=0.0,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
api_format=task.client_api_format,
|
||||
endpoint_api_format=task.provider_api_format,
|
||||
has_format_conversion=bool(task.format_converted),
|
||||
is_stream=False,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=None,
|
||||
status_code=200 if task.status == "completed" else 500,
|
||||
error_message=(
|
||||
None
|
||||
if task.status == "completed"
|
||||
else (task.error_message or task.error_code or "video_task_failed")
|
||||
),
|
||||
metadata={
|
||||
"fallback_created": True,
|
||||
"video_task_id": task.id,
|
||||
},
|
||||
request_headers=(
|
||||
(task.request_metadata or {}).get("request_headers")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
request_body=task.original_request_body,
|
||||
provider_request_headers=None,
|
||||
response_headers=None,
|
||||
client_response_headers=None,
|
||||
response_body=None,
|
||||
request_id=request_id,
|
||||
provider_id=task.provider_id,
|
||||
provider_endpoint_id=task.endpoint_id,
|
||||
provider_api_key_id=task.key_id,
|
||||
status="completed" if task.status == "completed" else "failed",
|
||||
target_model=None,
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Failed to create fallback usage for video task=%s: %s",
|
||||
task.id,
|
||||
str(exc),
|
||||
)
|
||||
return False
|
||||
36
src/services/task/context.py
Normal file
36
src/services/task/context.py
Normal file
@@ -0,0 +1,36 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class TaskMode(str, Enum):
|
||||
SYNC = "sync"
|
||||
ASYNC = "async"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TaskContext:
|
||||
"""
|
||||
TaskContext (pure DTO)
|
||||
|
||||
- Only primitive types / IDs
|
||||
- Serializable & safe to pass across processes
|
||||
"""
|
||||
|
||||
request_id: str
|
||||
task_type: str # chat/cli/video/image/audio
|
||||
task_mode: TaskMode
|
||||
|
||||
user_id: str
|
||||
api_key_id: str
|
||||
|
||||
client_ip: str = ""
|
||||
user_agent: str = ""
|
||||
start_time: float = 0.0
|
||||
|
||||
api_format: str | None = None
|
||||
model: str | None = None
|
||||
mapped_model: str | None = None
|
||||
|
||||
capability_requirements: dict[str, bool] = field(default_factory=dict)
|
||||
@@ -1,3 +1,7 @@
|
||||
"""Task telemetry implementations for concrete task types (video/image/audio)."""
|
||||
"""Per-task-type implementations (Phase2).
|
||||
|
||||
Currently includes:
|
||||
- video: polling adapter
|
||||
"""
|
||||
|
||||
__all__ = []
|
||||
|
||||
626
src/services/task/impl/video_poller.py
Normal file
626
src/services/task/impl/video_poller.py
Normal file
@@ -0,0 +1,626 @@
|
||||
"""
|
||||
Video task poller adapter.
|
||||
|
||||
Implements the video-specific poll/normalize/update logic used by TaskPollerService.
|
||||
|
||||
优化:HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
|
||||
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
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 (
|
||||
build_upstream_headers_for_endpoint,
|
||||
get_extra_headers_from_endpoint,
|
||||
make_signature_key,
|
||||
)
|
||||
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.database import create_session
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
|
||||
from src.services.task.application import TaskApplicationService
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class VideoPollContext:
|
||||
"""视频轮询上下文,保存 HTTP 请求所需的数据(不依赖数据库会话)"""
|
||||
|
||||
task_id: str
|
||||
external_task_id: str
|
||||
provider_api_format: str
|
||||
base_url: str
|
||||
upstream_key: str
|
||||
headers: dict[str, str]
|
||||
# 用于更新任务的原始数据
|
||||
poll_count: int
|
||||
retry_count: int
|
||||
poll_interval_seconds: int
|
||||
max_poll_count: int
|
||||
current_status: str
|
||||
|
||||
|
||||
# 永久性错误指示词(用于降级判断,不应重试)
|
||||
_PERMANENT_ERROR_INDICATORS = frozenset(
|
||||
{
|
||||
"not found",
|
||||
"404",
|
||||
"unauthorized",
|
||||
"401",
|
||||
"forbidden",
|
||||
"403",
|
||||
"invalid request",
|
||||
"invalid api key",
|
||||
"does not exist",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class PollHTTPError(RuntimeError):
|
||||
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
|
||||
|
||||
def __init__(self, status_code: int, message: str):
|
||||
# 确保错误信息包含状态码
|
||||
full_message = f"HTTP {status_code}: {message}" if message else f"HTTP {status_code}"
|
||||
super().__init__(full_message)
|
||||
self.status_code = status_code
|
||||
self.original_message = message
|
||||
|
||||
|
||||
class VideoTaskPollerAdapter:
|
||||
task_type = "video"
|
||||
|
||||
# scheduler
|
||||
job_id = "task_poller:video"
|
||||
job_name = "视频任务轮询"
|
||||
interval_seconds = config.video_poll_interval_seconds
|
||||
|
||||
# distributed lock
|
||||
lock_key = "task_poller:video:lock"
|
||||
lock_ttl = 60
|
||||
|
||||
# execution
|
||||
batch_size = config.video_poll_batch_size
|
||||
concurrency = config.video_poll_concurrency
|
||||
consecutive_failure_alert_threshold = 5
|
||||
max_backoff_seconds = 300
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._openai_normalizer = OpenAINormalizer()
|
||||
self._gemini_normalizer = GeminiNormalizer()
|
||||
|
||||
def sanitize_error_message(self, message: str) -> str:
|
||||
return sanitize_error_message(message)
|
||||
|
||||
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]:
|
||||
tasks = (
|
||||
db.query(VideoTask)
|
||||
.filter(
|
||||
VideoTask.status.in_(
|
||||
[
|
||||
VideoStatus.SUBMITTED.value,
|
||||
VideoStatus.QUEUED.value,
|
||||
VideoStatus.PROCESSING.value,
|
||||
]
|
||||
),
|
||||
VideoTask.next_poll_at <= now,
|
||||
VideoTask.poll_count < VideoTask.max_poll_count,
|
||||
)
|
||||
.order_by(VideoTask.next_poll_at.asc())
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [t.id for t in tasks]
|
||||
|
||||
def get_task(self, db: Session, task_id: str) -> VideoTask | None:
|
||||
# SQLAlchemy 1.4+ API
|
||||
return db.get(VideoTask, task_id)
|
||||
|
||||
# ==================== 分阶段处理方法(优化数据库连接占用)====================
|
||||
|
||||
async def prepare_poll_context(
|
||||
self, db: Session, task: VideoTask
|
||||
) -> VideoPollContext | InternalVideoPollResult:
|
||||
"""
|
||||
阶段 1:准备轮询上下文(短暂持有数据库连接)
|
||||
|
||||
Returns:
|
||||
VideoPollContext: 成功时返回上下文
|
||||
InternalVideoPollResult: 失败时返回错误结果
|
||||
"""
|
||||
if not task.endpoint_id or not task.key_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_provider_info",
|
||||
error_message="Task missing endpoint_id or key_id",
|
||||
)
|
||||
|
||||
endpoint = self._get_endpoint(db, task.endpoint_id)
|
||||
key = self._get_key(db, task.key_id)
|
||||
|
||||
if not key.api_key:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="provider_config_error",
|
||||
error_message="Provider key not properly configured",
|
||||
)
|
||||
|
||||
try:
|
||||
upstream_key = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
logger.warning("Failed to decrypt provider key for task %s", task.id)
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="decryption_error",
|
||||
error_message="Failed to decrypt provider key",
|
||||
)
|
||||
|
||||
provider_format = (task.provider_api_format or "").strip().lower()
|
||||
if not provider_format:
|
||||
provider_format = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
|
||||
# 构建请求头
|
||||
if provider_format.startswith("gemini:"):
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
else:
|
||||
auth_info = None
|
||||
headers = self._build_headers(provider_format, upstream_key, endpoint, auth_info)
|
||||
|
||||
return VideoPollContext(
|
||||
task_id=task.id,
|
||||
external_task_id=task.external_task_id or "",
|
||||
provider_api_format=provider_format,
|
||||
base_url=endpoint.base_url or "",
|
||||
upstream_key=upstream_key,
|
||||
headers=headers,
|
||||
poll_count=task.poll_count,
|
||||
retry_count=task.retry_count,
|
||||
poll_interval_seconds=task.poll_interval_seconds,
|
||||
max_poll_count=task.max_poll_count,
|
||||
current_status=task.status,
|
||||
)
|
||||
|
||||
async def poll_task_http(self, ctx: VideoPollContext) -> InternalVideoPollResult:
|
||||
"""
|
||||
阶段 2:执行 HTTP 请求(不持有数据库连接)
|
||||
|
||||
Args:
|
||||
ctx: 轮询上下文
|
||||
|
||||
Returns:
|
||||
InternalVideoPollResult: 轮询结果
|
||||
"""
|
||||
if not ctx.external_task_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_external_task_id",
|
||||
error_message="Task missing external_task_id",
|
||||
)
|
||||
|
||||
if ctx.provider_api_format.startswith("gemini:"):
|
||||
return await self._poll_gemini_with_context(ctx)
|
||||
return await self._poll_openai_with_context(ctx)
|
||||
|
||||
async def update_task_after_poll(
|
||||
self,
|
||||
task_id: str,
|
||||
result: InternalVideoPollResult,
|
||||
ctx: VideoPollContext | None,
|
||||
redis_client: Any | None,
|
||||
error_exception: Exception | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
阶段 3:更新数据库(获取新的数据库连接)
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID
|
||||
result: 轮询结果
|
||||
ctx: 轮询上下文(准备阶段就失败时为 None)
|
||||
redis_client: Redis 客户端
|
||||
error_exception: 如果 HTTP 请求失败,传入异常对象
|
||||
"""
|
||||
with create_session() as db:
|
||||
task = db.get(VideoTask, task_id)
|
||||
if not task:
|
||||
logger.warning("Task %s disappeared during poll update", task_id)
|
||||
return
|
||||
|
||||
if error_exception is not None and ctx is not None:
|
||||
# HTTP 请求失败(需要 ctx 来计算 backoff)
|
||||
self._handle_poll_error(task, error_exception, ctx)
|
||||
elif result.status == VideoStatus.COMPLETED:
|
||||
task.status = VideoStatus.COMPLETED.value
|
||||
task.video_url = result.video_url
|
||||
task.video_expires_at = result.expires_at
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
task.progress_percent = 100
|
||||
if result.video_urls:
|
||||
task.video_urls = result.video_urls
|
||||
self._attach_poll_raw_response(task, result)
|
||||
elif result.status == VideoStatus.FAILED:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = result.error_code
|
||||
task.error_message = result.error_message
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
self._attach_poll_raw_response(task, result)
|
||||
else:
|
||||
task.poll_count += 1
|
||||
task.progress_percent = result.progress_percent
|
||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
|
||||
seconds=task.poll_interval_seconds
|
||||
)
|
||||
|
||||
# 超时检查
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
if task.poll_count >= task.max_poll_count and task.status not in [
|
||||
VideoStatus.COMPLETED.value,
|
||||
VideoStatus.FAILED.value,
|
||||
VideoStatus.CANCELLED.value,
|
||||
]:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "poll_timeout"
|
||||
task.error_message = f"Task timed out after {task.poll_count} polls"
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
|
||||
# 终态结算
|
||||
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||
try:
|
||||
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
|
||||
task
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Failed to record video usage for task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
db.commit()
|
||||
|
||||
def _handle_poll_error(self, task: VideoTask, exc: Exception, ctx: VideoPollContext) -> None:
|
||||
"""处理轮询错误"""
|
||||
task.poll_count += 1
|
||||
error_msg = sanitize_error_message(str(exc))
|
||||
logger.warning("Poll error for task %s: %s", task.id, error_msg)
|
||||
task.progress_message = f"Poll error: {error_msg}"
|
||||
|
||||
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
|
||||
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
||||
if is_permanent:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "poll_permanent_error"
|
||||
task.error_message = error_msg
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
else:
|
||||
backoff = min(
|
||||
ctx.poll_interval_seconds * (2 ** min(ctx.retry_count, 5)),
|
||||
self.max_backoff_seconds,
|
||||
)
|
||||
task.retry_count += 1
|
||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
||||
|
||||
async def _poll_openai_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
|
||||
"""使用上下文进行 OpenAI 轮询(不需要数据库)"""
|
||||
url = self._build_openai_url(ctx.base_url, ctx.external_task_id)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.get(url, headers=ctx.headers)
|
||||
if response.status_code >= 400:
|
||||
error_message = self._extract_error_message(response.text, response.status_code)
|
||||
raise PollHTTPError(response.status_code, error_message)
|
||||
|
||||
payload = response.json()
|
||||
return self._openai_normalizer.video_poll_to_internal(payload)
|
||||
|
||||
async def _poll_gemini_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
|
||||
"""使用上下文进行 Gemini 轮询(不需要数据库)"""
|
||||
operation_name = normalize_gemini_operation_id(ctx.external_task_id)
|
||||
url = self._build_gemini_url(ctx.base_url, operation_name)
|
||||
|
||||
logger.debug(
|
||||
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
|
||||
ctx.task_id,
|
||||
ctx.external_task_id,
|
||||
url,
|
||||
)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.get(url, headers=ctx.headers)
|
||||
if response.status_code >= 400:
|
||||
logger.warning(
|
||||
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
|
||||
ctx.task_id,
|
||||
response.status_code,
|
||||
response.text[:500] if response.text else "(empty)",
|
||||
)
|
||||
error_message = self._extract_error_message(response.text, response.status_code)
|
||||
raise PollHTTPError(response.status_code, error_message)
|
||||
|
||||
payload = response.json()
|
||||
return self._gemini_normalizer.video_poll_to_internal(payload)
|
||||
|
||||
# ==================== 旧版方法(保留兼容性)====================
|
||||
|
||||
async def poll_single_task(
|
||||
self, db: Session, task: VideoTask, *, redis_client: Any | None
|
||||
) -> None:
|
||||
"""
|
||||
旧版单任务轮询方法(保留向后兼容)
|
||||
|
||||
注意:此方法在 HTTP 请求期间持有数据库连接,建议使用分阶段方法。
|
||||
"""
|
||||
try:
|
||||
result = await self._poll_task_status(db, task)
|
||||
if result.status == VideoStatus.COMPLETED:
|
||||
task.status = VideoStatus.COMPLETED.value
|
||||
task.video_url = result.video_url
|
||||
task.video_expires_at = result.expires_at
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
task.progress_percent = 100
|
||||
if result.video_urls:
|
||||
task.video_urls = result.video_urls
|
||||
self._attach_poll_raw_response(task, result)
|
||||
elif result.status == VideoStatus.FAILED:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = result.error_code
|
||||
task.error_message = result.error_message
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
self._attach_poll_raw_response(task, result)
|
||||
else:
|
||||
task.poll_count += 1
|
||||
task.progress_percent = result.progress_percent
|
||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
|
||||
seconds=task.poll_interval_seconds
|
||||
)
|
||||
except Exception as exc:
|
||||
task.poll_count += 1
|
||||
error_msg = sanitize_error_message(str(exc))
|
||||
logger.warning("Poll error for task %s: %s", task.id, error_msg)
|
||||
task.progress_message = f"Poll error: {error_msg}"
|
||||
|
||||
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
|
||||
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
||||
if is_permanent:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "poll_permanent_error"
|
||||
task.error_message = error_msg
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
else:
|
||||
backoff = min(
|
||||
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
|
||||
self.max_backoff_seconds,
|
||||
)
|
||||
task.retry_count += 1
|
||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
||||
|
||||
# 超时:超过最大轮询次数且未进入终态
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
if task.poll_count >= task.max_poll_count and task.status not in [
|
||||
VideoStatus.COMPLETED.value,
|
||||
VideoStatus.FAILED.value,
|
||||
VideoStatus.CANCELLED.value,
|
||||
]:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "poll_timeout"
|
||||
task.error_message = f"Task timed out after {task.poll_count} polls"
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
|
||||
# 终态结算
|
||||
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||
try:
|
||||
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
|
||||
task
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Failed to record video usage for task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
|
||||
if not result.raw_response:
|
||||
return
|
||||
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
|
||||
# (直接修改 JSON 字段内部不会自动标记为 dirty)
|
||||
metadata = dict(task.request_metadata) if task.request_metadata else {}
|
||||
metadata["poll_raw_response"] = result.raw_response
|
||||
task.request_metadata = metadata
|
||||
|
||||
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
|
||||
if status_code is not None:
|
||||
return 400 <= status_code < 500 and status_code != 429
|
||||
error_msg = str(exc).lower()
|
||||
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
|
||||
|
||||
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
|
||||
if not task.endpoint_id or not task.key_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_provider_info",
|
||||
error_message="Task missing endpoint_id or key_id",
|
||||
)
|
||||
endpoint = self._get_endpoint(db, task.endpoint_id)
|
||||
key = self._get_key(db, task.key_id)
|
||||
if not key.api_key:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="provider_config_error",
|
||||
error_message="Provider key not properly configured",
|
||||
)
|
||||
try:
|
||||
upstream_key = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
logger.warning("Failed to decrypt provider key for task %s", task.id)
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="decryption_error",
|
||||
error_message="Failed to decrypt provider key",
|
||||
)
|
||||
|
||||
provider_format = (task.provider_api_format or "").strip().lower()
|
||||
if not provider_format:
|
||||
provider_format = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
|
||||
if provider_format.startswith("gemini:"):
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
|
||||
return await self._poll_openai(task, endpoint, upstream_key)
|
||||
|
||||
async def _poll_openai(
|
||||
self,
|
||||
task: VideoTask,
|
||||
endpoint: ProviderEndpoint,
|
||||
upstream_key: str,
|
||||
) -> InternalVideoPollResult:
|
||||
if not task.external_task_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_external_task_id",
|
||||
error_message="Task missing external_task_id",
|
||||
)
|
||||
url = self._build_openai_url(endpoint.base_url, task.external_task_id)
|
||||
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.get(url, headers=headers)
|
||||
if response.status_code >= 400:
|
||||
error_message = self._extract_error_message(response.text, response.status_code)
|
||||
raise PollHTTPError(response.status_code, error_message)
|
||||
|
||||
payload = response.json()
|
||||
return self._openai_normalizer.video_poll_to_internal(payload)
|
||||
|
||||
async def _poll_gemini(
|
||||
self,
|
||||
task: VideoTask,
|
||||
endpoint: ProviderEndpoint,
|
||||
upstream_key: str,
|
||||
auth_info: ProviderAuthInfo | None,
|
||||
) -> InternalVideoPollResult:
|
||||
if not task.external_task_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_external_task_id",
|
||||
error_message="Task missing external_task_id",
|
||||
)
|
||||
operation_name = normalize_gemini_operation_id(task.external_task_id)
|
||||
url = self._build_gemini_url(endpoint.base_url, operation_name)
|
||||
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
|
||||
|
||||
logger.debug(
|
||||
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
|
||||
task.id,
|
||||
task.external_task_id,
|
||||
url,
|
||||
)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.get(url, headers=headers)
|
||||
if response.status_code >= 400:
|
||||
logger.warning(
|
||||
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
|
||||
task.id,
|
||||
response.status_code,
|
||||
response.text[:500] if response.text else "(empty)",
|
||||
)
|
||||
error_message = self._extract_error_message(response.text, response.status_code)
|
||||
raise PollHTTPError(response.status_code, error_message)
|
||||
|
||||
payload = response.json()
|
||||
return self._gemini_normalizer.video_poll_to_internal(payload)
|
||||
|
||||
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
|
||||
base = (base_url or "https://api.openai.com").rstrip("/")
|
||||
if base.endswith("/v1"):
|
||||
return f"{base}/videos/{task_id}"
|
||||
return f"{base}/v1/videos/{task_id}"
|
||||
|
||||
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
|
||||
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
|
||||
if base.endswith("/v1beta"):
|
||||
base = base[: -len("/v1beta")]
|
||||
return f"{base}/v1beta/{operation_name}"
|
||||
|
||||
def _build_headers(
|
||||
self,
|
||||
endpoint_sig: str,
|
||||
upstream_key: str,
|
||||
endpoint: ProviderEndpoint,
|
||||
auth_info: ProviderAuthInfo | None = None,
|
||||
) -> dict[str, str]:
|
||||
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||
headers = build_upstream_headers_for_endpoint(
|
||||
{},
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
)
|
||||
if auth_info:
|
||||
headers.pop("x-goog-api-key", None)
|
||||
headers[auth_info.auth_header] = auth_info.auth_value
|
||||
return headers
|
||||
|
||||
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
|
||||
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
|
||||
if not endpoint:
|
||||
raise RuntimeError("Provider endpoint not found")
|
||||
return endpoint
|
||||
|
||||
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
|
||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||
if not key:
|
||||
raise RuntimeError("Provider key not found")
|
||||
return key
|
||||
|
||||
def _extract_error_message(self, response_text: str | None, status_code: int) -> str:
|
||||
"""从响应中提取有意义的错误信息"""
|
||||
if not response_text:
|
||||
return f"Request failed with status {status_code}"
|
||||
|
||||
# 尝试解析 JSON 格式的错误
|
||||
try:
|
||||
data = json.loads(response_text)
|
||||
# OpenAI 格式: {"error": {"message": "..."}}
|
||||
if isinstance(data.get("error"), dict):
|
||||
error_obj = data["error"]
|
||||
message = error_obj.get("message") or error_obj.get("detail") or str(error_obj)
|
||||
return sanitize_error_message(message)
|
||||
# Gemini 格式: {"error": {"message": "...", "code": 404}}
|
||||
if "message" in data:
|
||||
return sanitize_error_message(data["message"])
|
||||
except (json.JSONDecodeError, TypeError, KeyError):
|
||||
pass
|
||||
|
||||
# 回退到原始文本(截断)
|
||||
return sanitize_error_message(response_text[:500])
|
||||
@@ -1,320 +0,0 @@
|
||||
"""
|
||||
VideoTelemetry(Phase3)
|
||||
|
||||
将 Video 异步任务的“终态计费 + Usage 写入 + required 缺失告警”从 poller 中抽离出来,
|
||||
便于未来 Image/Audio 复用相同框架。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.video_handler_base import sanitize_error_message
|
||||
from src.config.settings import config
|
||||
from src.core.api_format.conversion.internal_video import VideoStatus
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Provider, User, VideoTask
|
||||
from src.services.billing.dimension_collector_service import DimensionCollectorService
|
||||
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
|
||||
from src.services.billing.rule_service import BillingRuleService
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
class VideoTelemetry:
|
||||
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._formula_engine = FormulaEngine()
|
||||
|
||||
async def record_terminal_usage(self, task: VideoTask) -> None:
|
||||
"""
|
||||
为视频任务终态写入 Usage:
|
||||
- COMPLETED: 使用 FormulaEngine 计算 cost(或 no_rule / incomplete -> cost=0)
|
||||
- FAILED: cost=0
|
||||
|
||||
该方法可能会在 strict_mode 缺失 required 维度时将任务降级为 FAILED 并隐藏产物。
|
||||
"""
|
||||
request_id = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
request_id = task.request_metadata.get("request_id")
|
||||
request_id = request_id or task.id
|
||||
|
||||
# 计算异步任务总耗时(ms)
|
||||
response_time_ms = None
|
||||
if task.submitted_at and task.completed_at:
|
||||
delta = task.completed_at - task.submitted_at
|
||||
response_time_ms = int(delta.total_seconds() * 1000)
|
||||
|
||||
# 基础维度(无需 collectors 也可计费)
|
||||
base_dimensions: dict[str, Any] = {
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size or "",
|
||||
"retry_count": task.retry_count,
|
||||
}
|
||||
|
||||
# collectors 可用的 metadata(结构稳定,便于配置 path)
|
||||
collector_metadata: dict[str, Any] = {
|
||||
"task": {
|
||||
"id": task.id,
|
||||
"external_task_id": task.external_task_id,
|
||||
"model": task.model,
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size,
|
||||
"retry_count": task.retry_count,
|
||||
"video_size_bytes": task.video_size_bytes,
|
||||
},
|
||||
"result": {
|
||||
"video_url": task.video_url,
|
||||
"video_urls": task.video_urls or [],
|
||||
},
|
||||
}
|
||||
|
||||
# 维度采集:base + collectors 覆盖/补全
|
||||
dims = DimensionCollectorService(self.db).collect_dimensions(
|
||||
api_format=task.provider_api_format,
|
||||
task_type="video",
|
||||
request=task.original_request_body or {},
|
||||
response=(
|
||||
(task.request_metadata or {}).get("poll_raw_response")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
metadata=collector_metadata,
|
||||
base_dimensions=base_dimensions,
|
||||
)
|
||||
|
||||
# 取冻结的 rule_snapshot;若缺失则回退 DB 查找(兼容旧任务)
|
||||
rule_snapshot = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
|
||||
|
||||
billing_snapshot: dict[str, Any] = {
|
||||
"status": "complete",
|
||||
"missing_required": [],
|
||||
"strict_mode": config.billing_strict_mode,
|
||||
}
|
||||
cost = 0.0
|
||||
|
||||
if task.status == VideoStatus.FAILED.value:
|
||||
billing_snapshot["billed_reason"] = "task_failed"
|
||||
else:
|
||||
# COMPLETED:计算成本
|
||||
expression = None
|
||||
variables = None
|
||||
dimension_mappings = None
|
||||
rule_id = None
|
||||
rule_name = None
|
||||
rule_scope = None
|
||||
|
||||
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
|
||||
rule_id = rule_snapshot.get("rule_id")
|
||||
rule_name = rule_snapshot.get("rule_name")
|
||||
rule_scope = rule_snapshot.get("scope")
|
||||
expression = rule_snapshot.get("expression")
|
||||
variables = rule_snapshot.get("variables")
|
||||
dimension_mappings = rule_snapshot.get("dimension_mappings")
|
||||
else:
|
||||
lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=task.provider_id,
|
||||
model_name=task.model,
|
||||
task_type="video",
|
||||
)
|
||||
if lookup:
|
||||
rule = lookup.rule
|
||||
rule_id = rule.id
|
||||
rule_name = rule.name
|
||||
rule_scope = lookup.scope
|
||||
expression = rule.expression
|
||||
variables = rule.variables
|
||||
dimension_mappings = rule.dimension_mappings
|
||||
|
||||
if not expression:
|
||||
billing_snapshot["status"] = "no_rule"
|
||||
billing_snapshot["cost_breakdown"] = {"total": 0.0}
|
||||
logger.warning(
|
||||
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
task.provider_id,
|
||||
)
|
||||
else:
|
||||
billing_snapshot.update(
|
||||
{
|
||||
"rule_id": rule_id,
|
||||
"rule_name": rule_name,
|
||||
"rule_scope": rule_scope,
|
||||
"expression": expression,
|
||||
"variables": variables or {},
|
||||
}
|
||||
)
|
||||
try:
|
||||
result = self._formula_engine.evaluate(
|
||||
expression=expression,
|
||||
variables=variables or {},
|
||||
dimensions=dims,
|
||||
dimension_mappings=dimension_mappings or {},
|
||||
strict_mode=config.billing_strict_mode,
|
||||
)
|
||||
billing_snapshot["status"] = result.status
|
||||
billing_snapshot["missing_required"] = result.missing_required
|
||||
billing_snapshot["resolved_values"] = result.resolved_values
|
||||
if result.status == "complete":
|
||||
cost = result.cost
|
||||
else:
|
||||
logger.error(
|
||||
"Billing incomplete due to missing required dimensions "
|
||||
"(request_id=%s, model=%s, missing=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
result.missing_required,
|
||||
)
|
||||
cost = 0.0
|
||||
await self._maybe_alert_missing_required(
|
||||
model=task.model,
|
||||
missing_required=result.missing_required,
|
||||
)
|
||||
if result.error:
|
||||
billing_snapshot["error"] = result.error
|
||||
except BillingIncompleteError as exc:
|
||||
logger.error(
|
||||
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
exc.missing_required,
|
||||
)
|
||||
billing_snapshot["status"] = "incomplete"
|
||||
billing_snapshot["missing_required"] = exc.missing_required
|
||||
billing_snapshot["resolved_values"] = {}
|
||||
billing_snapshot["error"] = "strict_mode_missing_required"
|
||||
cost = 0.0
|
||||
|
||||
# strict_mode=true:标记任务失败并隐藏产物,避免"免费放行"
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "billing_incomplete"
|
||||
task.error_message = f"Missing required dimensions: {exc.missing_required}"
|
||||
task.video_url = None
|
||||
task.video_urls = None
|
||||
|
||||
await self._maybe_alert_missing_required(
|
||||
model=task.model,
|
||||
missing_required=exc.missing_required,
|
||||
)
|
||||
|
||||
billing_snapshot["cost_breakdown"] = {"total": cost}
|
||||
|
||||
# 将 billing_snapshot 回写到 task.request_metadata 便于对账(不会影响 usage 的单独存档)
|
||||
if task.request_metadata is None:
|
||||
task.request_metadata = {}
|
||||
if isinstance(task.request_metadata, dict):
|
||||
task.request_metadata["billing_snapshot"] = billing_snapshot
|
||||
|
||||
usage_metadata: dict[str, Any] = {
|
||||
"billing_snapshot": billing_snapshot,
|
||||
"dimensions": dims,
|
||||
"raw_response_ref": {
|
||||
"video_task_id": task.id,
|
||||
"field": "video_tasks.request_metadata.poll_raw_response",
|
||||
},
|
||||
}
|
||||
|
||||
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
|
||||
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
|
||||
api_key_obj = (
|
||||
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
|
||||
if task.api_key_id
|
||||
else None
|
||||
)
|
||||
provider_obj = (
|
||||
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
|
||||
if task.provider_id
|
||||
else None
|
||||
)
|
||||
provider_name = provider_obj.name if provider_obj else "unknown"
|
||||
|
||||
await UsageService.record_usage_with_custom_cost(
|
||||
db=self.db,
|
||||
user=user_obj,
|
||||
api_key=api_key_obj,
|
||||
provider=provider_name,
|
||||
model=task.model,
|
||||
request_type="video",
|
||||
total_cost_usd=cost,
|
||||
request_cost_usd=cost,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
api_format=task.client_api_format,
|
||||
endpoint_api_format=task.provider_api_format,
|
||||
has_format_conversion=bool(task.format_converted),
|
||||
is_stream=False,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=None,
|
||||
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
|
||||
error_message=(
|
||||
None
|
||||
if task.status == VideoStatus.COMPLETED.value
|
||||
else (task.error_message or task.error_code or "video_task_failed")
|
||||
),
|
||||
metadata=usage_metadata,
|
||||
request_headers=(
|
||||
(task.request_metadata or {}).get("request_headers")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
request_body=task.original_request_body,
|
||||
provider_request_headers=None,
|
||||
response_headers=None,
|
||||
client_response_headers=None,
|
||||
response_body=None,
|
||||
request_id=request_id,
|
||||
provider_id=task.provider_id,
|
||||
provider_endpoint_id=task.endpoint_id,
|
||||
provider_api_key_id=task.key_id,
|
||||
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
|
||||
target_model=None,
|
||||
)
|
||||
|
||||
async def _maybe_alert_missing_required(
|
||||
self, *, model: str, missing_required: list[str]
|
||||
) -> None:
|
||||
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
|
||||
if not missing_required:
|
||||
return
|
||||
if not self.redis:
|
||||
logger.error(
|
||||
"Missing required billing dimensions (model=%s): %s", model, missing_required
|
||||
)
|
||||
return
|
||||
|
||||
# 按小时 bucket 聚合
|
||||
now = datetime.now(timezone.utc)
|
||||
hour_bucket = now.strftime("%Y%m%d%H")
|
||||
for dim in missing_required:
|
||||
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
|
||||
try:
|
||||
count = await self.redis.incr(key)
|
||||
if count == 1:
|
||||
await self.redis.expire(key, 3700)
|
||||
if count >= 10:
|
||||
logger.warning(
|
||||
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
|
||||
model,
|
||||
dim,
|
||||
count,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["VideoTelemetry"]
|
||||
27
src/services/task/lifecycle.py
Normal file
27
src/services/task/lifecycle.py
Normal file
@@ -0,0 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class TaskStatus(str, Enum):
|
||||
"""Generic task status (progress)."""
|
||||
|
||||
PENDING = "pending"
|
||||
STREAMING = "streaming"
|
||||
|
||||
SUBMITTED = "submitted"
|
||||
QUEUED = "queued"
|
||||
PROCESSING = "processing"
|
||||
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
EXPIRED = "expired"
|
||||
|
||||
|
||||
class BillingStatus(str, Enum):
|
||||
"""Billing settlement status (Usage.billing_status)."""
|
||||
|
||||
PENDING = "pending"
|
||||
SETTLED = "settled"
|
||||
VOID = "void"
|
||||
@@ -1,635 +0,0 @@
|
||||
"""
|
||||
AsyncTaskOrchestrator
|
||||
|
||||
提交阶段故障转移(多候选尝试):
|
||||
- 目标:拿到 external_task_id 后锁定 provider/endpoint/key,后续轮询不再切换。
|
||||
- 仅覆盖“提交阶段”;轮询阶段由各 task poller 使用已锁定的信息执行。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, RequestCandidate
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
||||
from src.services.orchestration.candidate_resolver import CandidateResolver
|
||||
from src.services.orchestration.error_classifier import ErrorClassifier
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _sanitize(message: str, max_length: int = 200) -> str:
|
||||
if not message:
|
||||
return "request_failed"
|
||||
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SubmitFunc(Protocol):
|
||||
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ExtractExternalTaskIdFunc(Protocol):
|
||||
def __call__(self, payload: dict[str, Any]) -> str | None: ...
|
||||
|
||||
|
||||
class UpstreamClientRequestError(RuntimeError):
|
||||
"""可判定为客户端请求问题(不应 failover)的上游错误。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
response: httpx.Response,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
) -> None:
|
||||
self.response = response
|
||||
self.candidate_keys = candidate_keys
|
||||
super().__init__(f"Upstream client error: HTTP {response.status_code}")
|
||||
|
||||
|
||||
class AllCandidatesFailedError(RuntimeError):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
reason: str,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
last_status_code: int | None = None,
|
||||
) -> None:
|
||||
self.reason = reason
|
||||
self.candidate_keys = candidate_keys
|
||||
self.last_status_code = last_status_code
|
||||
super().__init__(f"All candidates failed: {reason}")
|
||||
|
||||
|
||||
class CandidateUnsupportedError(RuntimeError):
|
||||
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
|
||||
|
||||
|
||||
class CandidateSubmissionError(RuntimeError):
|
||||
"""候选提交异常(网络/解密/解析等)。"""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SubmitOutcome:
|
||||
candidate: ProviderCandidate
|
||||
candidate_keys: list[dict[str, Any]]
|
||||
external_task_id: str
|
||||
rule_lookup: BillingRuleLookupResult | None
|
||||
upstream_payload: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class AsyncTaskOrchestrator:
|
||||
"""
|
||||
异步任务编排器:只负责提交阶段的候选遍历与错误处理策略。
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session, *, redis_client: Redis | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._candidate_resolver: CandidateResolver | None = None
|
||||
self._error_classifier: ErrorClassifier | None = None
|
||||
|
||||
self._cache_scheduler = None
|
||||
# 候选记录映射:{candidate_index: RequestCandidate}
|
||||
self._candidate_records: dict[int, RequestCandidate] = {}
|
||||
|
||||
def _create_candidate_records(
|
||||
self,
|
||||
candidates: list[ProviderCandidate],
|
||||
request_id: str | None,
|
||||
user_api_key: ApiKey,
|
||||
) -> dict[int, RequestCandidate]:
|
||||
"""
|
||||
为所有候选预创建 RequestCandidate 记录。
|
||||
|
||||
Args:
|
||||
candidates: 候选列表
|
||||
request_id: 请求 ID
|
||||
user_api_key: 用户 API Key
|
||||
|
||||
Returns:
|
||||
{candidate_index: RequestCandidate} 映射
|
||||
"""
|
||||
if not request_id:
|
||||
return {}
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
records: dict[int, RequestCandidate] = {}
|
||||
|
||||
for idx, cand in enumerate(candidates):
|
||||
record = RequestCandidate(
|
||||
id=str(uuid.uuid4()),
|
||||
request_id=request_id,
|
||||
candidate_index=idx,
|
||||
retry_index=0,
|
||||
user_id=user_api_key.user_id if user_api_key else None,
|
||||
api_key_id=user_api_key.id if user_api_key else None,
|
||||
provider_id=cand.provider.id,
|
||||
endpoint_id=cand.endpoint.id,
|
||||
key_id=cand.key.id,
|
||||
status="available",
|
||||
is_cached=bool(getattr(cand, "is_cached", False)),
|
||||
created_at=now,
|
||||
)
|
||||
self.db.add(record)
|
||||
records[idx] = record
|
||||
|
||||
try:
|
||||
self.db.flush()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Failed to create candidate records: %s",
|
||||
str(exc),
|
||||
)
|
||||
self.db.rollback()
|
||||
return {}
|
||||
|
||||
return records
|
||||
|
||||
def _update_candidate_record(
|
||||
self,
|
||||
idx: int,
|
||||
*,
|
||||
status: str,
|
||||
skip_reason: str | None = None,
|
||||
status_code: int | None = None,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
started_at: datetime | None = None,
|
||||
finished_at: datetime | None = None,
|
||||
) -> None:
|
||||
"""更新候选记录状态。"""
|
||||
record = self._candidate_records.get(idx)
|
||||
if not record:
|
||||
return
|
||||
|
||||
record.status = status
|
||||
if skip_reason is not None:
|
||||
record.skip_reason = skip_reason
|
||||
if status_code is not None:
|
||||
record.status_code = status_code
|
||||
if error_type is not None:
|
||||
record.error_type = error_type
|
||||
if error_message is not None:
|
||||
record.error_message = error_message
|
||||
if started_at is not None:
|
||||
record.started_at = started_at
|
||||
if finished_at is not None:
|
||||
record.finished_at = finished_at
|
||||
|
||||
try:
|
||||
self.db.flush()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Failed to update candidate record %d: %s",
|
||||
idx,
|
||||
str(exc),
|
||||
)
|
||||
|
||||
def _commit_candidate_records(self) -> None:
|
||||
"""提交候选记录到数据库。"""
|
||||
if not self._candidate_records:
|
||||
return
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Failed to commit candidate records: %s",
|
||||
str(exc),
|
||||
)
|
||||
self.db.rollback()
|
||||
|
||||
async def _ensure_initialized(self) -> None:
|
||||
if self._cache_scheduler is not None:
|
||||
return
|
||||
|
||||
# 使用 SystemConfigService 读取运行时调度策略(与 Chat/CLI 一致)
|
||||
priority_mode = SystemConfigService.get_config(
|
||||
self.db,
|
||||
"provider_priority_mode",
|
||||
"provider",
|
||||
)
|
||||
scheduling_mode = SystemConfigService.get_config(
|
||||
self.db,
|
||||
"scheduling_mode",
|
||||
"cache_affinity",
|
||||
)
|
||||
self._cache_scheduler = await get_cache_aware_scheduler(
|
||||
self.redis,
|
||||
priority_mode=priority_mode,
|
||||
scheduling_mode=scheduling_mode,
|
||||
)
|
||||
|
||||
self._candidate_resolver = CandidateResolver(
|
||||
db=self.db,
|
||||
cache_scheduler=self._cache_scheduler,
|
||||
)
|
||||
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
|
||||
|
||||
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
|
||||
"""
|
||||
判断某个上游 HTTP 错误是否为“客户端错误”(不应 failover)。
|
||||
|
||||
规则:
|
||||
- 401/403/429:一般是 key/权限/限流问题,优先 failover
|
||||
- 其他 4xx:若 ErrorClassifier 判断为客户端请求错误,则停止
|
||||
"""
|
||||
if status_code in (401, 403, 429):
|
||||
return False
|
||||
if 400 <= status_code < 500:
|
||||
assert self._error_classifier is not None
|
||||
return self._error_classifier.is_client_error(error_text)
|
||||
return False
|
||||
|
||||
async def submit_with_failover(
|
||||
self,
|
||||
*,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey,
|
||||
request_id: str | None,
|
||||
task_type: str,
|
||||
submit_func: SubmitFunc,
|
||||
extract_external_task_id: ExtractExternalTaskIdFunc,
|
||||
supported_auth_types: set[str] | None = None,
|
||||
allow_format_conversion: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
max_candidates: int | None = None,
|
||||
) -> SubmitOutcome:
|
||||
"""
|
||||
提交异步任务并在失败时自动尝试下一个候选,直到拿到 external_task_id。
|
||||
|
||||
Returns:
|
||||
SubmitOutcome(包含选中的候选 + external_task_id + candidate_keys + billing rule lookup)
|
||||
|
||||
Raises:
|
||||
UpstreamClientRequestError: 判定为客户端请求错误(不应 failover)
|
||||
ProviderNotAvailableException: 没有可用候选(调度器层面)
|
||||
AllCandidatesFailedError: 有候选但全部提交失败
|
||||
"""
|
||||
await self._ensure_initialized()
|
||||
assert self._candidate_resolver is not None
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] submit_with_failover: "
|
||||
"api_format=%s, model=%s, task_type=%s, request_id=%s",
|
||||
api_format,
|
||||
model_name,
|
||||
task_type,
|
||||
request_id,
|
||||
)
|
||||
|
||||
candidates, _global_model_id = await self._candidate_resolver.fetch_candidates(
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
user_api_key=user_api_key,
|
||||
request_id=request_id,
|
||||
is_stream=False,
|
||||
capability_requirements=capability_requirements,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] fetch_candidates returned %d candidates for model=%s",
|
||||
len(candidates),
|
||||
model_name,
|
||||
)
|
||||
|
||||
# 如果没有候选,直接抛出异常
|
||||
if not candidates:
|
||||
logger.error(
|
||||
"[AsyncTaskOrchestrator] No candidates returned from fetch_candidates for model=%s",
|
||||
model_name,
|
||||
)
|
||||
raise ProviderNotAvailableException("No candidates available")
|
||||
|
||||
if max_candidates is not None and max_candidates > 0:
|
||||
candidates = candidates[:max_candidates]
|
||||
|
||||
# 创建候选记录(用于链路追踪)
|
||||
self._candidate_records = self._create_candidate_records(
|
||||
candidates=candidates,
|
||||
request_id=request_id,
|
||||
user_api_key=user_api_key,
|
||||
)
|
||||
|
||||
candidate_keys: list[dict[str, Any]] = []
|
||||
eligible_count = 0
|
||||
last_status_code: int | None = None
|
||||
|
||||
for idx, cand in enumerate(candidates):
|
||||
submit_started_at = datetime.now(timezone.utc)
|
||||
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
|
||||
candidate_info: dict[str, Any] = {
|
||||
"index": idx,
|
||||
"provider_id": cand.provider.id,
|
||||
"provider_name": cand.provider.name,
|
||||
"endpoint_id": cand.endpoint.id,
|
||||
"key_id": cand.key.id,
|
||||
"key_name": cand.key.name,
|
||||
"auth_type": auth_type,
|
||||
"priority": getattr(cand.key, "priority", 0) or 0,
|
||||
"is_cached": bool(getattr(cand, "is_cached", False)),
|
||||
}
|
||||
candidate_keys.append(candidate_info)
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Checking candidate %d: provider=%s, is_skipped=%s, skip_reason=%s, needs_conversion=%s, auth_type=%s",
|
||||
idx,
|
||||
cand.provider.name,
|
||||
getattr(cand, "is_skipped", False),
|
||||
getattr(cand, "skip_reason", None),
|
||||
getattr(cand, "needs_conversion", False),
|
||||
auth_type,
|
||||
)
|
||||
|
||||
# 调度器层面标记为跳过(健康/熔断/并发等)
|
||||
if getattr(cand, "is_skipped", False):
|
||||
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
|
||||
candidate_info.update(
|
||||
{
|
||||
"skipped": True,
|
||||
"skip_reason": skip_reason,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d skipped: is_skipped=True, reason=%s",
|
||||
idx,
|
||||
cand.skip_reason,
|
||||
)
|
||||
continue
|
||||
|
||||
# 视频/图片等直连 upstream 的 handler 目前不支持跨格式转换
|
||||
if not allow_format_conversion and bool(getattr(cand, "needs_conversion", False)):
|
||||
candidate_info.update(
|
||||
{"skipped": True, "skip_reason": "format_conversion_not_supported"}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx, status="skipped", skip_reason="format_conversion_not_supported"
|
||||
)
|
||||
logger.info("[AsyncTaskOrchestrator] Candidate %d skipped: needs_conversion", idx)
|
||||
continue
|
||||
|
||||
# auth_type 过滤
|
||||
if supported_auth_types is not None and auth_type not in supported_auth_types:
|
||||
skip_reason = f"unsupported_auth_type:{auth_type}"
|
||||
candidate_info.update(
|
||||
{
|
||||
"skipped": True,
|
||||
"skip_reason": skip_reason,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d skipped: unsupported_auth_type=%s",
|
||||
idx,
|
||||
auth_type,
|
||||
)
|
||||
continue
|
||||
|
||||
# billing rule 过滤(可选)
|
||||
rule_lookup: BillingRuleLookupResult | None = None
|
||||
has_billing_rule = True
|
||||
if config.billing_require_rule:
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Checking billing rule for candidate %d (billing_require_rule=True)",
|
||||
idx,
|
||||
)
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=cand.provider.id,
|
||||
model_name=model_name,
|
||||
task_type=task_type,
|
||||
)
|
||||
has_billing_rule = rule_lookup is not None
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Billing rule lookup result: has_rule=%s",
|
||||
has_billing_rule,
|
||||
)
|
||||
if not has_billing_rule:
|
||||
candidate_info.update(
|
||||
{
|
||||
"has_billing_rule": False,
|
||||
"skipped": True,
|
||||
"skip_reason": "billing_rule_missing",
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx, status="skipped", skip_reason="billing_rule_missing"
|
||||
)
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d skipped: billing_rule_missing", idx
|
||||
)
|
||||
continue
|
||||
candidate_info["has_billing_rule"] = has_billing_rule
|
||||
|
||||
logger.info("[AsyncTaskOrchestrator] Candidate %d eligible, attempting submit", idx)
|
||||
eligible_count += 1
|
||||
|
||||
# 更新记录为 pending 状态(开始尝试)
|
||||
self._update_candidate_record(idx, status="pending", started_at=submit_started_at)
|
||||
|
||||
# 尝试提交
|
||||
try:
|
||||
response = await submit_func(cand)
|
||||
except Exception as exc:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
logger.error(
|
||||
"[AsyncTaskOrchestrator] Candidate %d submit exception: %s: %s",
|
||||
idx,
|
||||
type(exc).__name__,
|
||||
str(exc),
|
||||
)
|
||||
error_msg = _sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "exception",
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
error_type=type(exc).__name__,
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
continue
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d submit response: status_code=%d",
|
||||
idx,
|
||||
response.status_code,
|
||||
)
|
||||
|
||||
last_status_code = int(getattr(response, "status_code", 0) or 0)
|
||||
|
||||
# 上游错误:决定是否停止
|
||||
if response.status_code >= 400:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
error_text = ""
|
||||
try:
|
||||
error_text = response.text or ""
|
||||
except Exception:
|
||||
error_text = ""
|
||||
error_msg = _sanitize(error_text)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "http_error",
|
||||
"status_code": response.status_code,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="http_error",
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
if self._should_stop_on_http_error(
|
||||
status_code=response.status_code, error_text=error_text
|
||||
):
|
||||
self._commit_candidate_records()
|
||||
raise UpstreamClientRequestError(
|
||||
response=response,
|
||||
candidate_keys=candidate_keys,
|
||||
)
|
||||
continue
|
||||
|
||||
# 解析任务 ID(200 但缺字段也视为失败并 failover)
|
||||
payload: dict[str, Any] | None = None
|
||||
try:
|
||||
data = response.json()
|
||||
if isinstance(data, dict):
|
||||
payload = data
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d response payload: %s",
|
||||
idx,
|
||||
str(payload)[:500] if payload else "None",
|
||||
)
|
||||
except Exception as exc:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
logger.error(
|
||||
"[AsyncTaskOrchestrator] Candidate %d invalid JSON: %s",
|
||||
idx,
|
||||
str(exc),
|
||||
)
|
||||
error_msg = _sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "invalid_json",
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="invalid_json",
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
continue
|
||||
|
||||
external_task_id = extract_external_task_id(payload or {})
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d extracted task_id: %s",
|
||||
idx,
|
||||
external_task_id,
|
||||
)
|
||||
if not external_task_id:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "empty_task_id",
|
||||
"error_message": "Upstream returned empty task id",
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="empty_task_id",
|
||||
error_message="Upstream returned empty task id",
|
||||
finished_at=finished_at,
|
||||
)
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Candidate %d: empty task_id, payload keys: %s",
|
||||
idx,
|
||||
list(payload.keys()) if payload else [],
|
||||
)
|
||||
continue
|
||||
|
||||
# 成功
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
candidate_info.update({"attempt_status": "success", "selected": True})
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="success",
|
||||
status_code=response.status_code,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
self._commit_candidate_records()
|
||||
return SubmitOutcome(
|
||||
candidate=cand,
|
||||
candidate_keys=candidate_keys,
|
||||
external_task_id=str(external_task_id),
|
||||
rule_lookup=rule_lookup,
|
||||
upstream_payload=payload,
|
||||
)
|
||||
|
||||
# 没有任何候选可尝试
|
||||
if not candidates:
|
||||
raise ProviderNotAvailableException("No candidates available")
|
||||
|
||||
# 提交所有候选记录
|
||||
self._commit_candidate_records()
|
||||
|
||||
if eligible_count == 0:
|
||||
reason = "no_eligible_candidates"
|
||||
if config.billing_require_rule:
|
||||
reason = "no_candidate_with_billing_rule"
|
||||
raise AllCandidatesFailedError(
|
||||
reason=reason,
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
|
||||
raise AllCandidatesFailedError(
|
||||
reason="all_candidates_failed",
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AsyncTaskOrchestrator",
|
||||
"SubmitOutcome",
|
||||
"AllCandidatesFailedError",
|
||||
"UpstreamClientRequestError",
|
||||
"CandidateUnsupportedError",
|
||||
"CandidateSubmissionError",
|
||||
]
|
||||
240
src/services/task/task_poller.py
Normal file
240
src/services/task/task_poller.py
Normal file
@@ -0,0 +1,240 @@
|
||||
"""
|
||||
Task poller (Phase2)
|
||||
|
||||
Provides a generic polling skeleton for async tasks.
|
||||
Currently wired with a video poller adapter.
|
||||
|
||||
优化:HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
|
||||
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format.conversion.internal_video import InternalVideoPollResult
|
||||
from src.core.logger import logger
|
||||
from src.database import create_session
|
||||
from src.services.system.scheduler import get_scheduler
|
||||
from src.services.task.impl.video_poller import VideoPollContext, VideoTaskPollerAdapter
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class TaskPollerAdapter(Protocol):
|
||||
task_type: str
|
||||
|
||||
# scheduler
|
||||
job_id: str
|
||||
job_name: str
|
||||
interval_seconds: int
|
||||
|
||||
# distributed lock (optional, best-effort)
|
||||
lock_key: str
|
||||
lock_ttl: int
|
||||
|
||||
# execution
|
||||
batch_size: int
|
||||
concurrency: int
|
||||
consecutive_failure_alert_threshold: int
|
||||
|
||||
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]: ...
|
||||
|
||||
def get_task(self, db: Session, task_id: str) -> Any | None: ...
|
||||
|
||||
# 分阶段处理方法(推荐使用)
|
||||
async def prepare_poll_context(
|
||||
self, db: Session, task: Any
|
||||
) -> Any: ... # Returns context or error result
|
||||
|
||||
async def poll_task_http(self, ctx: Any) -> Any: ... # Returns poll result
|
||||
|
||||
async def update_task_after_poll(
|
||||
self,
|
||||
task_id: str,
|
||||
result: Any,
|
||||
ctx: Any,
|
||||
redis_client: Any | None,
|
||||
error_exception: Exception | None = None,
|
||||
) -> None: ...
|
||||
|
||||
# 旧版方法(保留兼容性)
|
||||
async def poll_single_task(
|
||||
self, db: Session, task: Any, *, redis_client: Any | None
|
||||
) -> None: ...
|
||||
|
||||
def sanitize_error_message(self, message: str) -> str: ...
|
||||
|
||||
|
||||
class TaskPollerService:
|
||||
"""Generic background poller for async tasks."""
|
||||
|
||||
def __init__(self, adapter: TaskPollerAdapter) -> None:
|
||||
self.adapter = adapter
|
||||
self._lock = asyncio.Lock()
|
||||
self.redis: Any | None = None
|
||||
self._semaphore: asyncio.Semaphore | None = None
|
||||
self._consecutive_failures = 0
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._semaphore is None:
|
||||
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
|
||||
|
||||
# lazy import to avoid redis hard dependency in local runs
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
if self.redis is None:
|
||||
self.redis = await get_redis_client(require_redis=False)
|
||||
|
||||
scheduler = get_scheduler()
|
||||
scheduler.add_interval_job(
|
||||
self.poll_pending_tasks,
|
||||
seconds=self.adapter.interval_seconds,
|
||||
job_id=self.adapter.job_id,
|
||||
name=self.adapter.job_name,
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
scheduler = get_scheduler()
|
||||
scheduler.remove_job(self.adapter.job_id)
|
||||
|
||||
async def poll_pending_tasks(self) -> None:
|
||||
async with self._lock:
|
||||
token = await self._acquire_redis_lock()
|
||||
if token is None:
|
||||
return
|
||||
|
||||
try:
|
||||
with create_session() as db:
|
||||
now = datetime.now(timezone.utc)
|
||||
task_ids = self.adapter.list_due_task_ids(
|
||||
db, now=now, limit=self.adapter.batch_size
|
||||
)
|
||||
|
||||
if not task_ids:
|
||||
self._consecutive_failures = 0
|
||||
return
|
||||
|
||||
poll_results: list[bool] = []
|
||||
|
||||
if self._semaphore is None:
|
||||
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
|
||||
semaphore = self._semaphore
|
||||
|
||||
async def poll_with_semaphore(task_id: str) -> None:
|
||||
async with semaphore:
|
||||
try:
|
||||
# ========== 阶段 1:准备数据(短暂持有连接)==========
|
||||
with create_session() as task_db:
|
||||
task_obj = self.adapter.get_task(task_db, task_id)
|
||||
if not task_obj:
|
||||
logger.warning(
|
||||
"[%s] Task %s disappeared during poll",
|
||||
self.adapter.task_type,
|
||||
task_id,
|
||||
)
|
||||
poll_results.append(True)
|
||||
return
|
||||
|
||||
ctx_or_result = await self.adapter.prepare_poll_context(
|
||||
task_db, task_obj
|
||||
)
|
||||
|
||||
# 检查是否是错误结果(而非上下文)
|
||||
if isinstance(ctx_or_result, InternalVideoPollResult):
|
||||
# 准备阶段就失败了,直接更新任务状态
|
||||
await self.adapter.update_task_after_poll(
|
||||
task_id=task_id,
|
||||
result=ctx_or_result,
|
||||
ctx=None, # type: ignore[arg-type]
|
||||
redis_client=self.redis,
|
||||
)
|
||||
poll_results.append(True)
|
||||
return
|
||||
|
||||
ctx: VideoPollContext = ctx_or_result
|
||||
|
||||
# ========== 阶段 2:HTTP 请求(不持有数据库连接)==========
|
||||
error_exception: Exception | None = None
|
||||
try:
|
||||
result = await self.adapter.poll_task_http(ctx)
|
||||
except Exception as http_exc:
|
||||
# HTTP 请求失败,记录异常以便后续处理
|
||||
error_exception = http_exc
|
||||
result = InternalVideoPollResult(
|
||||
status=None, # type: ignore[arg-type]
|
||||
error_message=str(http_exc),
|
||||
)
|
||||
|
||||
# ========== 阶段 3:更新数据库(获取新连接)==========
|
||||
await self.adapter.update_task_after_poll(
|
||||
task_id=task_id,
|
||||
result=result,
|
||||
ctx=ctx,
|
||||
redis_client=self.redis,
|
||||
error_exception=error_exception,
|
||||
)
|
||||
poll_results.append(True)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"[%s] Unexpected error polling task %s: %s",
|
||||
self.adapter.task_type,
|
||||
task_id,
|
||||
self.adapter.sanitize_error_message(str(exc)),
|
||||
)
|
||||
poll_results.append(False)
|
||||
|
||||
async with asyncio.TaskGroup() as tg:
|
||||
for tid in task_ids:
|
||||
tg.create_task(poll_with_semaphore(tid))
|
||||
|
||||
batch_failures = sum(1 for r in poll_results if r is False)
|
||||
if batch_failures == len(task_ids):
|
||||
self._consecutive_failures += 1
|
||||
if (
|
||||
self._consecutive_failures
|
||||
>= self.adapter.consecutive_failure_alert_threshold
|
||||
):
|
||||
logger.error(
|
||||
"[ALERT] %s poller: %d consecutive batches failed.",
|
||||
self.adapter.task_type,
|
||||
self._consecutive_failures,
|
||||
)
|
||||
else:
|
||||
self._consecutive_failures = 0
|
||||
finally:
|
||||
await self._release_redis_lock(token)
|
||||
|
||||
async def _acquire_redis_lock(self) -> str | None:
|
||||
if not self.redis:
|
||||
return "no_redis"
|
||||
token = str(uuid4())
|
||||
acquired = await self.redis.set(
|
||||
self.adapter.lock_key, token, nx=True, ex=self.adapter.lock_ttl
|
||||
)
|
||||
return token if acquired else None
|
||||
|
||||
async def _release_redis_lock(self, token: str) -> None:
|
||||
if not self.redis or token == "no_redis":
|
||||
return
|
||||
script = """
|
||||
if redis.call('GET', KEYS[1]) == ARGV[1] then
|
||||
return redis.call('DEL', KEYS[1])
|
||||
end
|
||||
return 0
|
||||
"""
|
||||
await self.redis.eval(script, 1, self.adapter.lock_key, token)
|
||||
|
||||
|
||||
_task_poller: TaskPollerService | None = None
|
||||
|
||||
|
||||
def get_task_poller() -> TaskPollerService:
|
||||
global _task_poller
|
||||
if _task_poller is None:
|
||||
_task_poller = TaskPollerService(VideoTaskPollerAdapter())
|
||||
return _task_poller
|
||||
@@ -576,6 +576,12 @@ class UsageService:
|
||||
"""更新已存在的 Usage 记录(内部方法)"""
|
||||
# 更新关键字段
|
||||
existing_usage.provider_name = usage_params["provider_name"]
|
||||
existing_usage.model = usage_params["model"]
|
||||
existing_usage.request_type = usage_params["request_type"]
|
||||
existing_usage.api_format = usage_params["api_format"]
|
||||
existing_usage.endpoint_api_format = usage_params["endpoint_api_format"]
|
||||
existing_usage.has_format_conversion = usage_params["has_format_conversion"]
|
||||
existing_usage.is_stream = usage_params["is_stream"]
|
||||
existing_usage.status = usage_params["status"]
|
||||
existing_usage.status_code = usage_params["status_code"]
|
||||
existing_usage.error_message = usage_params["error_message"]
|
||||
@@ -621,6 +627,10 @@ class UsageService:
|
||||
existing_usage.provider_endpoint_id = usage_params["provider_endpoint_id"]
|
||||
existing_usage.provider_api_key_id = usage_params["provider_api_key_id"]
|
||||
|
||||
# 更新元数据(如 billing_snapshot/dimensions 等)
|
||||
if usage_params.get("request_metadata") is not None:
|
||||
existing_usage.request_metadata = usage_params["request_metadata"]
|
||||
|
||||
# 更新模型映射信息
|
||||
if target_model is not None:
|
||||
existing_usage.target_model = target_model
|
||||
@@ -1000,6 +1010,11 @@ class UsageService:
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||
)
|
||||
|
||||
# 结算标记:record_usage_async 写入的 Usage 通常为终态记录
|
||||
if status not in ("pending", "streaming"):
|
||||
usage.billing_status = "settled"
|
||||
usage.finalized_at = datetime.now(timezone.utc)
|
||||
|
||||
db.commit() # 立即提交事务,释放数据库锁
|
||||
return usage
|
||||
|
||||
@@ -1172,6 +1187,11 @@ class UsageService:
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||
)
|
||||
|
||||
# 结算标记:终态请求写入 settled + finalized_at
|
||||
if status not in ("pending", "streaming"):
|
||||
usage.billing_status = "settled"
|
||||
usage.finalized_at = datetime.now(timezone.utc)
|
||||
|
||||
# 提交事务
|
||||
try:
|
||||
db.commit()
|
||||
@@ -1297,17 +1317,39 @@ class UsageService:
|
||||
is_free_tier=is_free_tier,
|
||||
)
|
||||
|
||||
# Upsert(与 record_usage 保持一致)
|
||||
# Upsert(并发幂等:优先用 billing_status 作为结算闸门)
|
||||
from sqlalchemy import update
|
||||
|
||||
existing_usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if existing_usage:
|
||||
# 避免重复记账:若已是终态记录,直接返回(批量接口也采用该策略)
|
||||
if existing_usage.status not in ("pending", "streaming"):
|
||||
# 避免重复记账:若已结算/作废,直接返回(防止并发重复加计数)
|
||||
if getattr(existing_usage, "billing_status", None) in ("settled", "void"):
|
||||
logger.debug(
|
||||
"record_usage_with_custom_cost: request_id=%s already finalized (status=%s), skip",
|
||||
"record_usage_with_custom_cost: request_id=%s already finalized (billing_status=%s), skip",
|
||||
request_id,
|
||||
existing_usage.status,
|
||||
getattr(existing_usage, "billing_status", None),
|
||||
)
|
||||
return existing_usage
|
||||
|
||||
# 并发闸门:只有 billing_status='pending' 的那一次调用可以继续
|
||||
now = datetime.now(timezone.utc)
|
||||
claim = db.execute(
|
||||
update(Usage)
|
||||
.where(
|
||||
Usage.request_id == request_id,
|
||||
Usage.billing_status == "pending",
|
||||
)
|
||||
.values(billing_status="settled", finalized_at=now)
|
||||
)
|
||||
if claim.rowcount != 1:
|
||||
# 已被其他 worker 抢先处理(或被 VOID)
|
||||
latest = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
return latest or existing_usage
|
||||
|
||||
# 同步 ORM 对象(避免后续代码读到旧值)
|
||||
existing_usage.billing_status = "settled"
|
||||
existing_usage.finalized_at = now
|
||||
|
||||
cls._update_existing_usage(existing_usage, usage_params, target_model)
|
||||
usage = existing_usage
|
||||
else:
|
||||
@@ -1382,6 +1424,11 @@ class UsageService:
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||
)
|
||||
|
||||
# 结算标记:record_usage_with_custom_cost 写入/更新的 Usage 通常为终态记录
|
||||
if status not in ("pending", "streaming"):
|
||||
usage.billing_status = "settled"
|
||||
usage.finalized_at = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
@@ -2147,35 +2194,31 @@ class UsageService:
|
||||
# ========== 请求状态追踪方法 ==========
|
||||
|
||||
@classmethod
|
||||
def create_pending_usage(
|
||||
def begin_pending_usage(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
user: User | None,
|
||||
api_key: ApiKey | None,
|
||||
model: str,
|
||||
*,
|
||||
is_stream: bool = False,
|
||||
request_type: str = "chat",
|
||||
api_format: str | None = None,
|
||||
request_headers: dict[str, Any] | None = None,
|
||||
request_body: Any | None = None,
|
||||
) -> Usage:
|
||||
"""
|
||||
创建 pending 状态的使用记录(在请求开始时调用)
|
||||
创建(或返回已有)pending Usage 记录,但**不提交事务**。
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
request_id: 请求ID
|
||||
user: 用户对象
|
||||
api_key: API Key 对象
|
||||
model: 模型名称
|
||||
is_stream: 是否流式请求
|
||||
api_format: API 格式
|
||||
request_headers: 请求头
|
||||
request_body: 请求体
|
||||
|
||||
Returns:
|
||||
创建的 Usage 记录
|
||||
适用场景:
|
||||
- ApplicationService 在同一事务内创建 pending usage + task + candidates
|
||||
- submit 幂等:重复调用同一 request_id 时返回已有记录
|
||||
"""
|
||||
existing = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if existing:
|
||||
return existing
|
||||
|
||||
# 根据配置决定是否记录请求详情
|
||||
should_log_headers = SystemConfigService.should_log_headers(db)
|
||||
should_log_body = SystemConfigService.should_log_body(db)
|
||||
@@ -2204,21 +2247,359 @@ class UsageService:
|
||||
output_tokens=0,
|
||||
total_tokens=0,
|
||||
total_cost_usd=0.0,
|
||||
request_type="chat",
|
||||
request_type=request_type,
|
||||
api_format=api_format,
|
||||
is_stream=is_stream,
|
||||
status="pending",
|
||||
billing_status="pending",
|
||||
request_headers=processed_request_headers,
|
||||
request_body=processed_request_body,
|
||||
)
|
||||
|
||||
db.add(usage)
|
||||
db.flush()
|
||||
return usage
|
||||
|
||||
@classmethod
|
||||
def create_pending_usage(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
user: User | None,
|
||||
api_key: ApiKey | None,
|
||||
model: str,
|
||||
is_stream: bool = False,
|
||||
request_type: str = "chat",
|
||||
api_format: str | None = None,
|
||||
request_headers: dict[str, Any] | None = None,
|
||||
request_body: Any | None = None,
|
||||
) -> Usage:
|
||||
"""
|
||||
创建 pending 状态的使用记录(在请求开始时调用)
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
request_id: 请求ID
|
||||
user: 用户对象
|
||||
api_key: API Key 对象
|
||||
model: 模型名称
|
||||
is_stream: 是否流式请求
|
||||
api_format: API 格式
|
||||
request_headers: 请求头
|
||||
request_body: 请求体
|
||||
|
||||
Returns:
|
||||
创建的 Usage 记录
|
||||
"""
|
||||
usage = cls.begin_pending_usage(
|
||||
db,
|
||||
request_id=request_id,
|
||||
user=user,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
is_stream=is_stream,
|
||||
request_type=request_type,
|
||||
api_format=api_format,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
)
|
||||
db.commit()
|
||||
|
||||
logger.debug(f"创建 pending 使用记录: request_id={request_id}, model={model}")
|
||||
|
||||
return usage
|
||||
|
||||
# ========== billing_status 并发幂等 finalize ==========
|
||||
|
||||
@classmethod
|
||||
def finalize_settled(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
*,
|
||||
total_cost_usd: float,
|
||||
request_cost_usd: float | None = None,
|
||||
status: str = "completed",
|
||||
status_code: int = 200,
|
||||
error_message: str | None = None,
|
||||
response_time_ms: int | None = None,
|
||||
billing_snapshot: dict[str, Any] | None = None,
|
||||
extra_metadata: dict[str, Any] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
并发安全的幂等 finalize(settled)。
|
||||
|
||||
约定:
|
||||
- 仅当 billing_status='pending' 时才会生效(rowcount==1)
|
||||
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||
"""
|
||||
from sqlalchemy import update
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
cost = float(total_cost_usd)
|
||||
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
|
||||
|
||||
result = db.execute(
|
||||
update(Usage)
|
||||
.where(
|
||||
Usage.request_id == request_id,
|
||||
Usage.billing_status == "pending",
|
||||
)
|
||||
.values(
|
||||
billing_status="settled",
|
||||
finalized_at=now,
|
||||
total_cost_usd=cost,
|
||||
request_cost_usd=request_cost,
|
||||
status=status,
|
||||
status_code=status_code,
|
||||
error_message=error_message,
|
||||
response_time_ms=response_time_ms,
|
||||
)
|
||||
)
|
||||
if result.rowcount != 1:
|
||||
return False
|
||||
|
||||
# 写入审计快照(只在本次 finalize 生效时执行)
|
||||
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if usage:
|
||||
metadata = usage.request_metadata or {}
|
||||
if billing_snapshot is not None:
|
||||
metadata["billing_snapshot"] = billing_snapshot
|
||||
if extra_metadata:
|
||||
metadata.update(extra_metadata)
|
||||
usage.request_metadata = metadata
|
||||
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def finalize_void(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
*,
|
||||
reason: str | None = None,
|
||||
status_code: int = 499,
|
||||
) -> bool:
|
||||
"""
|
||||
并发安全的幂等 finalize(void,不收费)。
|
||||
|
||||
约定:
|
||||
- 仅当 billing_status='pending' 时才会生效(rowcount==1)
|
||||
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||
"""
|
||||
from sqlalchemy import update
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
result = db.execute(
|
||||
update(Usage)
|
||||
.where(
|
||||
Usage.request_id == request_id,
|
||||
Usage.billing_status == "pending",
|
||||
)
|
||||
.values(
|
||||
billing_status="void",
|
||||
finalized_at=now,
|
||||
total_cost_usd=0.0,
|
||||
request_cost_usd=0.0,
|
||||
status="cancelled",
|
||||
status_code=status_code,
|
||||
error_message=reason,
|
||||
response_time_ms=None,
|
||||
)
|
||||
)
|
||||
return result.rowcount == 1
|
||||
|
||||
@classmethod
|
||||
def finalize_submitted(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
*,
|
||||
provider_name: str,
|
||||
provider_id: str | None = None,
|
||||
provider_endpoint_id: str | None = None,
|
||||
provider_api_key_id: str | None = None,
|
||||
response_time_ms: int | None = None,
|
||||
status_code: int = 200,
|
||||
endpoint_api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
response_body: Any | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
异步任务提交成功时的幂等结算。
|
||||
|
||||
将 pending 使用记录标记为 settled,费用暂时为 0。
|
||||
后续轮询完成后通过 update_settled_billing 更新实际费用。
|
||||
|
||||
约定:
|
||||
- 仅当 billing_status='pending' 时才会生效(rowcount==1)
|
||||
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||
"""
|
||||
from sqlalchemy import update
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# 处理响应头和响应体
|
||||
should_log_headers = SystemConfigService.should_log_headers(db)
|
||||
should_log_body = SystemConfigService.should_log_body(db)
|
||||
|
||||
processed_provider_headers = None
|
||||
if should_log_headers and provider_request_headers:
|
||||
processed_provider_headers = SystemConfigService.mask_sensitive_headers(
|
||||
db, provider_request_headers
|
||||
)
|
||||
|
||||
processed_response_headers = None
|
||||
if should_log_headers and response_headers:
|
||||
processed_response_headers = dict(response_headers)
|
||||
|
||||
processed_response_body = None
|
||||
if should_log_body and response_body:
|
||||
processed_response_body = SystemConfigService.truncate_body(
|
||||
db, response_body, is_request=False
|
||||
)
|
||||
|
||||
values: dict[str, Any] = {
|
||||
"billing_status": "settled",
|
||||
"finalized_at": now,
|
||||
"total_cost_usd": 0.0,
|
||||
"request_cost_usd": 0.0,
|
||||
"status": "completed",
|
||||
"status_code": status_code,
|
||||
"response_time_ms": response_time_ms,
|
||||
"provider_name": provider_name,
|
||||
"provider_id": provider_id,
|
||||
"provider_endpoint_id": provider_endpoint_id,
|
||||
"provider_api_key_id": provider_api_key_id,
|
||||
"endpoint_api_format": endpoint_api_format,
|
||||
}
|
||||
|
||||
if processed_provider_headers is not None:
|
||||
values["provider_request_headers"] = processed_provider_headers
|
||||
if processed_response_headers is not None:
|
||||
values["response_headers"] = processed_response_headers
|
||||
if processed_response_body is not None:
|
||||
values["response_body"] = processed_response_body
|
||||
|
||||
result = db.execute(
|
||||
update(Usage)
|
||||
.where(
|
||||
Usage.request_id == request_id,
|
||||
Usage.billing_status == "pending",
|
||||
)
|
||||
.values(**values)
|
||||
)
|
||||
return result.rowcount == 1
|
||||
|
||||
@classmethod
|
||||
def update_settled_billing(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
*,
|
||||
total_cost_usd: float,
|
||||
request_cost_usd: float | None = None,
|
||||
status: str = "completed",
|
||||
status_code: int = 200,
|
||||
error_message: str | None = None,
|
||||
response_time_ms: int | None = None,
|
||||
billing_snapshot: dict[str, Any] | None = None,
|
||||
extra_metadata: dict[str, Any] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
更新已结算记录的计费信息(用于异步任务轮询完成后)。
|
||||
|
||||
与 finalize_settled 不同:
|
||||
- finalize_settled: pending -> settled(首次结算)
|
||||
- update_settled_billing: settled -> settled(更新费用)
|
||||
|
||||
约定:
|
||||
- 仅当 billing_status='settled' 时才会生效
|
||||
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||
"""
|
||||
from sqlalchemy import update
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
cost = float(total_cost_usd)
|
||||
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
|
||||
|
||||
values: dict[str, Any] = {
|
||||
"total_cost_usd": cost,
|
||||
"request_cost_usd": request_cost,
|
||||
"status": status,
|
||||
"status_code": status_code,
|
||||
}
|
||||
if error_message is not None:
|
||||
values["error_message"] = error_message
|
||||
if response_time_ms is not None:
|
||||
values["response_time_ms"] = response_time_ms
|
||||
|
||||
result = db.execute(
|
||||
update(Usage)
|
||||
.where(
|
||||
Usage.request_id == request_id,
|
||||
Usage.billing_status == "settled",
|
||||
)
|
||||
.values(**values)
|
||||
)
|
||||
if result.rowcount != 1:
|
||||
return False
|
||||
|
||||
# 写入审计快照
|
||||
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
|
||||
if usage:
|
||||
metadata = usage.request_metadata or {}
|
||||
if billing_snapshot is not None:
|
||||
metadata["billing_snapshot"] = billing_snapshot
|
||||
if extra_metadata:
|
||||
metadata.update(extra_metadata)
|
||||
metadata["billing_updated_at"] = now.isoformat()
|
||||
usage.request_metadata = metadata
|
||||
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def void_settled(
|
||||
cls,
|
||||
db: Session,
|
||||
request_id: str,
|
||||
*,
|
||||
reason: str | None = None,
|
||||
status_code: int = 499,
|
||||
) -> bool:
|
||||
"""
|
||||
将已结算的记录作废(用于异步任务取消)。
|
||||
|
||||
与 finalize_void 不同:
|
||||
- finalize_void: pending -> void(未结算时作废)
|
||||
- void_settled: settled -> void(已结算后取消,费用归零)
|
||||
|
||||
约定:
|
||||
- 仅当 billing_status='settled' 时才会生效
|
||||
- 不在本方法内 commit,由调用方决定事务提交时机
|
||||
"""
|
||||
from sqlalchemy import update
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
result = db.execute(
|
||||
update(Usage)
|
||||
.where(
|
||||
Usage.request_id == request_id,
|
||||
Usage.billing_status == "settled",
|
||||
)
|
||||
.values(
|
||||
billing_status="void",
|
||||
finalized_at=now,
|
||||
total_cost_usd=0.0,
|
||||
request_cost_usd=0.0,
|
||||
status="cancelled",
|
||||
status_code=status_code,
|
||||
error_message=reason,
|
||||
)
|
||||
)
|
||||
return result.rowcount == 1
|
||||
|
||||
@classmethod
|
||||
def update_usage_status(
|
||||
cls,
|
||||
@@ -2235,6 +2616,7 @@ class UsageService:
|
||||
api_format: str | None = None,
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool | None = None,
|
||||
status_code: int | None = None,
|
||||
) -> Usage | None:
|
||||
"""
|
||||
快速更新使用记录状态
|
||||
@@ -2253,6 +2635,7 @@ class UsageService:
|
||||
api_format: API 格式(可选,用于获取按格式配置的倍率)
|
||||
endpoint_api_format: 端点原生 API 格式(可选)
|
||||
has_format_conversion: 是否发生了格式转换(可选)
|
||||
status_code: HTTP 状态码(可选)
|
||||
|
||||
Returns:
|
||||
更新后的 Usage 记录,如果未找到则返回 None
|
||||
@@ -2295,6 +2678,16 @@ class UsageService:
|
||||
usage.endpoint_api_format = endpoint_api_format
|
||||
if has_format_conversion is not None:
|
||||
usage.has_format_conversion = has_format_conversion
|
||||
if status_code is not None:
|
||||
usage.status_code = status_code
|
||||
|
||||
# 结算状态:当请求进入终态时,将 billing_status 标记为 settled
|
||||
# 注意:取消是否应 VOID/部分结算由更高层策略决定;这里默认终态均视为已结算。
|
||||
if status in ("completed", "failed", "cancelled"):
|
||||
if getattr(usage, "billing_status", None) == "pending":
|
||||
usage.billing_status = "settled"
|
||||
if getattr(usage, "finalized_at", None) is None:
|
||||
usage.finalized_at = datetime.now(timezone.utc)
|
||||
|
||||
db.commit()
|
||||
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
"""
|
||||
视频相关服务
|
||||
"""
|
||||
|
||||
from src.services.video.task_poller import VideoTaskPollerService, get_video_task_poller
|
||||
|
||||
__all__ = [
|
||||
"VideoTaskPollerService",
|
||||
"get_video_task_poller",
|
||||
]
|
||||
@@ -1,446 +0,0 @@
|
||||
"""
|
||||
视频任务后台轮询服务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import (
|
||||
build_upstream_headers_for_endpoint,
|
||||
get_extra_headers_from_endpoint,
|
||||
make_signature_key,
|
||||
)
|
||||
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.database import create_session
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
|
||||
from src.services.system.scheduler import get_scheduler
|
||||
from src.services.task.impl.video_telemetry import VideoTelemetry
|
||||
|
||||
# 永久性错误指示词(用于降级判断,不应重试)
|
||||
_PERMANENT_ERROR_INDICATORS = frozenset(
|
||||
{
|
||||
"not found",
|
||||
"404",
|
||||
"unauthorized",
|
||||
"401",
|
||||
"forbidden",
|
||||
"403",
|
||||
"invalid request",
|
||||
"invalid api key",
|
||||
"does not exist",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class PollHTTPError(RuntimeError):
|
||||
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
|
||||
|
||||
def __init__(self, status_code: int, message: str):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
class VideoTaskPollerService:
|
||||
"""后台轮询视频生成任务状态"""
|
||||
|
||||
LOCK_KEY = "video_task_poller:lock"
|
||||
LOCK_TTL = 60
|
||||
MAX_BACKOFF_SECONDS = 300
|
||||
# 连续失败告警阈值
|
||||
CONSECUTIVE_FAILURE_ALERT_THRESHOLD = 5
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._lock = asyncio.Lock()
|
||||
self.redis = None
|
||||
self._openai_normalizer = OpenAINormalizer()
|
||||
self._gemini_normalizer = GeminiNormalizer()
|
||||
# 追踪连续失败次数(用于告警)
|
||||
self._consecutive_failures = 0
|
||||
# 从配置读取参数
|
||||
self._batch_size = config.video_poll_batch_size
|
||||
self._concurrency = config.video_poll_concurrency
|
||||
# Semaphore 延迟初始化,避免在事件循环外创建
|
||||
self._semaphore: asyncio.Semaphore | None = None
|
||||
|
||||
async def start(self) -> None:
|
||||
# 在事件循环内初始化 Semaphore
|
||||
if self._semaphore is None:
|
||||
self._semaphore = asyncio.Semaphore(self._concurrency)
|
||||
if self.redis is None:
|
||||
self.redis = await get_redis_client(require_redis=False)
|
||||
|
||||
scheduler = get_scheduler()
|
||||
scheduler.add_interval_job(
|
||||
self.poll_pending_tasks,
|
||||
seconds=config.video_poll_interval_seconds,
|
||||
job_id="video_task_poller",
|
||||
name="视频任务轮询",
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""停止轮询服务"""
|
||||
scheduler = get_scheduler()
|
||||
scheduler.remove_job("video_task_poller")
|
||||
|
||||
async def poll_pending_tasks(self) -> None:
|
||||
async with self._lock:
|
||||
token = await self._acquire_redis_lock()
|
||||
if token is None:
|
||||
return
|
||||
|
||||
try:
|
||||
with create_session() as db:
|
||||
now = datetime.now(timezone.utc)
|
||||
tasks = (
|
||||
db.query(VideoTask)
|
||||
.filter(
|
||||
VideoTask.status.in_(
|
||||
[
|
||||
VideoStatus.SUBMITTED.value,
|
||||
VideoStatus.QUEUED.value,
|
||||
VideoStatus.PROCESSING.value,
|
||||
]
|
||||
),
|
||||
VideoTask.next_poll_at <= now,
|
||||
VideoTask.poll_count < VideoTask.max_poll_count,
|
||||
)
|
||||
.order_by(VideoTask.next_poll_at.asc())
|
||||
.limit(self._batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
if not tasks:
|
||||
# 无任务时重置连续失败计数
|
||||
self._consecutive_failures = 0
|
||||
return
|
||||
|
||||
# 提取任务 ID 列表,释放查询 session 后逐个轮询
|
||||
task_ids = [t.id for t in tasks]
|
||||
|
||||
# 并发轮询:每个任务使用独立 session,避免共享 session 的并发风险
|
||||
poll_results: list[bool] = []
|
||||
|
||||
# 确保 semaphore 已初始化(在 start 中初始化,此处防御性检查)
|
||||
if self._semaphore is None:
|
||||
self._semaphore = asyncio.Semaphore(self._concurrency)
|
||||
semaphore = self._semaphore
|
||||
|
||||
async def poll_with_semaphore(task_id: str) -> None:
|
||||
"""带信号量的轮询,结果写入 poll_results"""
|
||||
async with semaphore:
|
||||
try:
|
||||
with create_session() as task_db:
|
||||
task_obj = task_db.query(VideoTask).get(task_id)
|
||||
if not task_obj:
|
||||
logger.warning("Task %s disappeared during poll", task_id)
|
||||
poll_results.append(True)
|
||||
return
|
||||
await self._poll_single_task(task_db, task_obj)
|
||||
task_db.commit()
|
||||
poll_results.append(True)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Unexpected error polling task %s: %s",
|
||||
task_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
poll_results.append(False)
|
||||
|
||||
async with asyncio.TaskGroup() as tg:
|
||||
for tid in task_ids:
|
||||
tg.create_task(poll_with_semaphore(tid))
|
||||
|
||||
batch_failures = sum(1 for r in poll_results if r is False)
|
||||
|
||||
# 更新连续失败计数并检查告警阈值
|
||||
if batch_failures == len(task_ids):
|
||||
self._consecutive_failures += 1
|
||||
if self._consecutive_failures >= self.CONSECUTIVE_FAILURE_ALERT_THRESHOLD:
|
||||
logger.error(
|
||||
"[ALERT] Video task poller: %d consecutive batches failed. "
|
||||
"Provider connectivity or configuration issue suspected.",
|
||||
self._consecutive_failures,
|
||||
)
|
||||
else:
|
||||
self._consecutive_failures = 0
|
||||
finally:
|
||||
await self._release_redis_lock(token)
|
||||
|
||||
async def _poll_single_task(self, db: Session, task: VideoTask) -> None:
|
||||
try:
|
||||
result = await self._poll_task_status(db, task)
|
||||
if result.status == VideoStatus.COMPLETED:
|
||||
task.status = VideoStatus.COMPLETED.value
|
||||
task.video_url = result.video_url
|
||||
task.video_expires_at = result.expires_at
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
task.progress_percent = 100
|
||||
# 存储多视频 URL(Gemini sampleCount > 1 时)
|
||||
if result.video_urls:
|
||||
task.video_urls = result.video_urls
|
||||
# 保存上游原始响应(用于审计/重算)
|
||||
self._attach_poll_raw_response(task, result)
|
||||
elif result.status == VideoStatus.FAILED:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = result.error_code
|
||||
task.error_message = result.error_message
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
self._attach_poll_raw_response(task, result)
|
||||
else:
|
||||
task.poll_count += 1
|
||||
task.progress_percent = result.progress_percent
|
||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
|
||||
seconds=task.poll_interval_seconds
|
||||
)
|
||||
except Exception as exc:
|
||||
task.poll_count += 1
|
||||
error_msg = sanitize_error_message(str(exc))
|
||||
logger.warning("Poll error for task %s: %s", task.id, error_msg)
|
||||
task.progress_message = f"Poll error: {error_msg}"
|
||||
|
||||
# 区分临时性错误和永久性错误
|
||||
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
|
||||
is_permanent = self._is_permanent_error(exc, status_code=status_code)
|
||||
if is_permanent:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "poll_permanent_error"
|
||||
task.error_message = error_msg
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
else:
|
||||
# 临时性错误:指数退避重试
|
||||
backoff = min(
|
||||
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
|
||||
self.MAX_BACKOFF_SECONDS,
|
||||
)
|
||||
task.retry_count += 1
|
||||
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
|
||||
|
||||
# 检查是否超过最大轮询次数(超时)
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
if task.poll_count >= task.max_poll_count and task.status not in [
|
||||
VideoStatus.COMPLETED.value,
|
||||
VideoStatus.FAILED.value,
|
||||
VideoStatus.CANCELLED.value,
|
||||
]:
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "poll_timeout"
|
||||
task.error_message = f"Task timed out after {task.poll_count} polls"
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
|
||||
# 终态写入 Usage(复用外层 per-task session)
|
||||
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||
try:
|
||||
await VideoTelemetry(db, redis_client=self.redis).record_terminal_usage(task)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Failed to record video usage for task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
|
||||
if not result.raw_response:
|
||||
return
|
||||
if task.request_metadata is None:
|
||||
task.request_metadata = {}
|
||||
# 仅在终态写一次,避免污染 request_metadata
|
||||
task.request_metadata["poll_raw_response"] = result.raw_response
|
||||
|
||||
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
|
||||
"""判断是否为永久性错误(不应重试)"""
|
||||
# 优先使用 HTTP 状态码判断
|
||||
if status_code is not None:
|
||||
# 4xx 客户端错误(除 429 限流)通常是永久性错误
|
||||
return 400 <= status_code < 500 and status_code != 429
|
||||
|
||||
# 降级到字符串匹配
|
||||
error_msg = str(exc).lower()
|
||||
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
|
||||
|
||||
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
|
||||
if not task.endpoint_id or not task.key_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_provider_info",
|
||||
error_message="Task missing endpoint_id or key_id",
|
||||
)
|
||||
endpoint = self._get_endpoint(db, task.endpoint_id)
|
||||
key = self._get_key(db, task.key_id)
|
||||
if not key.api_key:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="provider_config_error",
|
||||
error_message="Provider key not properly configured",
|
||||
)
|
||||
try:
|
||||
upstream_key = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
logger.warning("Failed to decrypt provider key for task %s", task.id)
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="decryption_error",
|
||||
error_message="Failed to decrypt provider key",
|
||||
)
|
||||
|
||||
provider_format = (task.provider_api_format or "").strip().lower()
|
||||
if not provider_format:
|
||||
provider_format = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
|
||||
if provider_format.startswith("gemini:"):
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
|
||||
return await self._poll_openai(task, endpoint, upstream_key)
|
||||
|
||||
async def _poll_openai(
|
||||
self,
|
||||
task: VideoTask,
|
||||
endpoint: ProviderEndpoint,
|
||||
upstream_key: str,
|
||||
) -> InternalVideoPollResult:
|
||||
if not task.external_task_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_external_task_id",
|
||||
error_message="Task missing external_task_id",
|
||||
)
|
||||
url = self._build_openai_url(endpoint.base_url, task.external_task_id)
|
||||
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.get(url, headers=headers)
|
||||
if response.status_code >= 400:
|
||||
raise PollHTTPError(
|
||||
response.status_code,
|
||||
sanitize_error_message(response.text or "Poll error"),
|
||||
)
|
||||
|
||||
payload = response.json()
|
||||
return self._openai_normalizer.video_poll_to_internal(payload)
|
||||
|
||||
async def _poll_gemini(
|
||||
self,
|
||||
task: VideoTask,
|
||||
endpoint: ProviderEndpoint,
|
||||
upstream_key: str,
|
||||
auth_info: ProviderAuthInfo | None,
|
||||
) -> InternalVideoPollResult:
|
||||
if not task.external_task_id:
|
||||
return InternalVideoPollResult(
|
||||
status=VideoStatus.FAILED,
|
||||
error_code="missing_external_task_id",
|
||||
error_message="Task missing external_task_id",
|
||||
)
|
||||
operation_name = normalize_gemini_operation_id(task.external_task_id)
|
||||
url = self._build_gemini_url(endpoint.base_url, operation_name)
|
||||
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.get(url, headers=headers)
|
||||
if response.status_code >= 400:
|
||||
raise PollHTTPError(
|
||||
response.status_code,
|
||||
sanitize_error_message(response.text or "Poll error"),
|
||||
)
|
||||
|
||||
payload = response.json()
|
||||
return self._gemini_normalizer.video_poll_to_internal(payload)
|
||||
|
||||
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
|
||||
base = (base_url or "https://api.openai.com").rstrip("/")
|
||||
if base.endswith("/v1"):
|
||||
return f"{base}/videos/{task_id}"
|
||||
return f"{base}/v1/videos/{task_id}"
|
||||
|
||||
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
|
||||
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
|
||||
if base.endswith("/v1beta"):
|
||||
base = base[: -len("/v1beta")]
|
||||
return f"{base}/v1beta/{operation_name}"
|
||||
|
||||
def _build_headers(
|
||||
self,
|
||||
endpoint_sig: str,
|
||||
upstream_key: str,
|
||||
endpoint: ProviderEndpoint,
|
||||
auth_info: ProviderAuthInfo | None = None,
|
||||
) -> dict[str, str]:
|
||||
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||
headers = build_upstream_headers_for_endpoint(
|
||||
{},
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
)
|
||||
if auth_info:
|
||||
headers.pop("x-goog-api-key", None)
|
||||
headers[auth_info.auth_header] = auth_info.auth_value
|
||||
return headers
|
||||
|
||||
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
|
||||
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
|
||||
if not endpoint:
|
||||
raise RuntimeError("Provider endpoint not found")
|
||||
return endpoint
|
||||
|
||||
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
|
||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||
if not key:
|
||||
raise RuntimeError("Provider key not found")
|
||||
return key
|
||||
|
||||
async def _acquire_redis_lock(self) -> str | None:
|
||||
if not self.redis:
|
||||
return "no_redis"
|
||||
token = str(uuid4())
|
||||
acquired = await self.redis.set(self.LOCK_KEY, token, nx=True, ex=self.LOCK_TTL)
|
||||
return token if acquired else None
|
||||
|
||||
async def _release_redis_lock(self, token: str) -> None:
|
||||
if not self.redis or token == "no_redis":
|
||||
return
|
||||
script = """
|
||||
if redis.call('GET', KEYS[1]) == ARGV[1] then
|
||||
return redis.call('DEL', KEYS[1])
|
||||
end
|
||||
return 0
|
||||
"""
|
||||
await self.redis.eval(script, 1, self.LOCK_KEY, token)
|
||||
|
||||
|
||||
_video_task_poller: VideoTaskPollerService | None = None
|
||||
|
||||
|
||||
def get_video_task_poller() -> VideoTaskPollerService:
|
||||
global _video_task_poller
|
||||
if _video_task_poller is None:
|
||||
_video_task_poller = VideoTaskPollerService()
|
||||
return _video_task_poller
|
||||
Reference in New Issue
Block a user