refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构

- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1735 additions and 2433 deletions

View File

@@ -0,0 +1,324 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
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, Usage, User
from src.services.usage.service import UsageService
class VideoTaskBillingService:
"""视频任务计费/结算服务。"""
def __init__(self, db: Session) -> None:
self.db = db
async def _create_fallback_usage_for_video_task(self, task: Any, request_id: str) -> bool:
"""
Fallback: create a Usage row if it's missing (should be rare).
This keeps behavior compatible with the old Phase2 finalize logic.
"""
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 getattr(task, "api_key_id", None)
else None
)
provider_obj = (
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
if getattr(task, "provider_id", None)
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
response_time_ms: int | None = None
if getattr(task, "submitted_at", None) and getattr(task, "completed_at", None):
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
request_headers: dict[str, Any] | None = None
if isinstance(getattr(task, "request_metadata", None), dict):
task_meta = task.request_metadata
for header_key in ("request_headers", "headers", "original_headers"):
raw_headers = task_meta.get(header_key)
if isinstance(raw_headers, dict):
request_headers = dict(raw_headers)
break
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(getattr(task, "format_converted", False)),
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=request_headers,
request_body=getattr(task, "original_request_body", None),
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=getattr(task, "provider_id", None),
provider_endpoint_id=getattr(task, "endpoint_id", None),
provider_api_key_id=getattr(task, "key_id", None),
status="completed" if task.status == "completed" else "failed",
target_model=None,
finalized_at=getattr(task, "completed_at", None),
)
return True
except Exception as exc:
logger.exception(
"Failed to create fallback usage for video task={}: {}",
task.id,
str(exc),
)
return False
async def finalize_video_task(self, task: Any) -> bool:
"""
Update billing/usage for a completed/failed video task.
Async video billing flow:
- Submit success: Usage is already settled with cost=0
- Poll completion: update actual cost (success -> bill, failure -> keep 0)
Returns True when updated, False when skipped (already finalized).
"""
from datetime import datetime, timezone
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
request_id = getattr(task, "request_id", None) or getattr(task, "id", None)
if not request_id:
return False
# Advisory check无锁仅用于快速跳过实际状态转换由 update_settled_billing 的
# with_for_update() 保证原子性。
existing = self.db.query(Usage).filter(Usage.request_id == request_id).first()
if not existing:
logger.warning(
"Usage not found for video task, creating fallback: task_id={} request_id={}",
getattr(task, "id", None),
request_id,
)
return await self._create_fallback_usage_for_video_task(task, request_id)
if getattr(existing, "billing_status", None) != "pending":
logger.debug(
"Skip video task billing finalize because Usage is already terminal: task_id={} request_id={} billing_status={}",
getattr(task, "id", None),
request_id,
getattr(existing, "billing_status", None),
)
return False
response_time_ms: int | None = None
if getattr(task, "submitted_at", None) and getattr(task, "completed_at", None):
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
base_dimensions: dict[str, Any] = {
"duration_seconds": getattr(task, "duration_seconds", None),
"resolution": getattr(task, "resolution", None),
"aspect_ratio": getattr(task, "aspect_ratio", None),
"size": getattr(task, "size", None) or "",
"retry_count": getattr(task, "retry_count", 0),
}
collector_metadata: dict[str, Any] = {
"task": {
"id": getattr(task, "id", None),
"external_task_id": getattr(task, "external_task_id", None),
"model": getattr(task, "model", None),
"duration_seconds": getattr(task, "duration_seconds", None),
"resolution": getattr(task, "resolution", None),
"aspect_ratio": getattr(task, "aspect_ratio", None),
"size": getattr(task, "size", None),
"retry_count": getattr(task, "retry_count", 0),
"video_size_bytes": getattr(task, "video_size_bytes", None),
},
"result": {
"video_url": getattr(task, "video_url", None),
"video_urls": getattr(task, "video_urls", None) or [],
},
}
poll_raw = None
if isinstance(getattr(task, "request_metadata", None), dict):
poll_raw = task.request_metadata.get("poll_raw_response")
dims = DimensionCollectorService(self.db).collect_dimensions(
api_format=getattr(task, "provider_api_format", None),
task_type="video",
request=getattr(task, "original_request_body", None) or {},
response=poll_raw if isinstance(poll_raw, dict) else None,
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# Prefer frozen rule snapshot from submit stage.
rule_snapshot = None
if isinstance(getattr(task, "request_metadata", None), 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=getattr(task, "provider_id", None),
model_name=getattr(task, "model", None),
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = getattr(rule, "id", None)
rule_name = getattr(rule, "name", None)
rule_scope = getattr(lookup, "scope", None)
expression = getattr(rule, "expression", None)
variables = getattr(rule, "variables", None) or {}
dimension_mappings = getattr(rule, "dimension_mappings", None) 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
is_success = str(getattr(task, "status", "")) in {
VideoStatus.COMPLETED.value,
"completed",
}
if is_success 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: mark task failed and hide artifacts (avoid free pass)
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
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["cost"] = 0.0
cost = 0.0
except Exception as exc:
billing_snapshot["status"] = "incomplete"
billing_snapshot["error"] = str(exc)
billing_snapshot["cost"] = 0.0
cost = 0.0
# Write back to task.request_metadata for audit/recalc.
task_meta = dict(task.request_metadata) if getattr(task, "request_metadata", None) else {}
task_meta["billing_snapshot"] = billing_snapshot
task.request_metadata = task_meta
updated = UsageService.update_settled_billing(
self.db,
request_id=request_id,
total_cost_usd=cost,
request_cost_usd=cost,
status="completed" if str(getattr(task, "status", "")) == "completed" else "failed",
status_code=200 if str(getattr(task, "status", "")) == "completed" else 500,
error_message=(
None
if str(getattr(task, "status", "")) == "completed"
else (
getattr(task, "error_message", None)
or getattr(task, "error_code", None)
or "video_task_failed"
)
),
response_time_ms=response_time_ms,
billing_snapshot=billing_snapshot,
extra_metadata={
"dimensions": dims,
"raw_response_ref": {
"video_task_id": getattr(task, "id", None),
"field": "video_tasks.request_metadata.poll_raw_response",
},
},
finalized_at=getattr(task, "completed_at", None),
)
if updated:
logger.debug(
"Updated video task billing: task_id={} request_id={} cost={:.6f}",
getattr(task, "id", None),
request_id,
cost,
)
else:
logger.warning(
"Failed to update video task billing (may already be updated): "
"task_id={} request_id={}",
getattr(task, "id", None),
request_id,
)
return bool(updated)

View File

@@ -0,0 +1,266 @@
from __future__ import annotations
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 ProviderAPIKey, ProviderEndpoint
from src.services.usage.service import UsageService
class VideoTaskCancelService:
"""视频任务取消服务(上游取消 + 本地状态与计费回写)。"""
def __init__(self, db: Session) -> None:
self.db = db
async def cancel_task(
self,
*,
task: Any,
task_id: str,
original_headers: dict[str, str] | None = None,
) -> Any:
"""
Cancel a video task (best-effort) and void its Usage (no charge).
Returns:
- None on success
- upstream httpx.Response when upstream returns an error (status >= 400)
"""
import json
from datetime import datetime, timezone
import httpx
from fastapi import HTTPException
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 VideoStatus
from src.core.crypto import crypto_service
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
current_status = str(getattr(task, "status", "") or "")
non_cancellable_statuses = {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}
if current_status in non_cancellable_statuses:
raise HTTPException(
status_code=409,
detail=f"Task cannot be cancelled in status: {current_status}",
)
external_task_id = getattr(task, "external_task_id", None)
if not external_task_id:
raise HTTPException(status_code=500, detail="Task missing external_task_id")
endpoint = (
self.db.query(ProviderEndpoint).filter(ProviderEndpoint.id == task.endpoint_id).first()
)
key = self.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
if not endpoint or not key:
raise HTTPException(status_code=500, detail="Provider endpoint or key not found")
if not getattr(key, "api_key", None):
raise HTTPException(status_code=500, detail="Provider key not configured")
upstream_key = crypto_service.decrypt(key.api_key)
extra_headers = get_extra_headers_from_endpoint(endpoint)
raw_family = str(getattr(endpoint, "api_family", "") or "").strip().lower()
raw_kind = str(getattr(endpoint, "endpoint_kind", "") or "").strip().lower()
provider_format = (
make_signature_key(raw_family, raw_kind)
if raw_family and raw_kind
else str(
getattr(endpoint, "api_format", "")
or getattr(task, "provider_api_format", "")
or ""
)
)
provider_format_norm = provider_format.strip().lower()
headers = build_upstream_headers_for_endpoint(
original_headers or {},
provider_format,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
async def _try_rust_cancel_response(
*,
method: str,
url: str,
request_headers: dict[str, str],
body: Any,
content_type: str | None = None,
) -> httpx.Response:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
build_execution_plan_body,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return httpx.Response(
status_code=503,
request=httpx.Request(method, url, headers=request_headers),
json={"error": {"message": "Video 取消仅支持 Rust executor"}},
)
final_headers = dict(request_headers)
if (
body is not None
and content_type
and not any(str(key).lower() == "content-type" for key in final_headers)
):
final_headers["content-type"] = content_type
try:
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=str(getattr(task, "request_id", "") or task_id),
candidate_id=None,
provider_name=provider_format_norm.split(":", 1)[0],
provider_id=str(getattr(endpoint, "provider_id", "") or ""),
endpoint_id=str(getattr(endpoint, "id", "") or ""),
key_id=str(getattr(key, "id", "") or ""),
method=method,
url=url,
headers=final_headers,
body=build_execution_plan_body(body, content_type=content_type),
stream=False,
provider_api_format=provider_format,
client_api_format=provider_format,
model_name=str(getattr(task, "model", "") or ""),
content_type=content_type,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=300_000,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[VideoCancel] Rust executor unavailable task={} method={} url={}: {}",
getattr(task, "id", task_id),
method,
url,
str(exc),
)
return httpx.Response(
status_code=503,
request=httpx.Request(method, url, headers=final_headers),
json={"error": {"message": "执行器暂时不可用,请稍后重试"}},
)
response_headers = dict(result.headers)
if result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(result.response_json, ensure_ascii=False).encode("utf-8")
elif result.response_body_bytes is not None:
response_body = result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=result.status_code,
request=httpx.Request(method, url, headers=final_headers),
headers=response_headers,
content=response_body,
)
if provider_format_norm.startswith("openai:"):
upstream_url = build_provider_url(endpoint, is_stream=False, key=key)
upstream_url = f"{upstream_url.rstrip('/')}/{str(external_task_id).lstrip('/')}"
response = await _try_rust_cancel_response(
method="DELETE",
url=upstream_url,
request_headers=headers,
body=None,
)
if response.status_code >= 400:
return response
elif provider_format_norm.startswith("gemini:"):
# Gemini cancel endpoint supports both:
# - operations/{id}:cancel
# - models/{model}/operations/{id}:cancel
operation_name = str(external_task_id)
if not (
operation_name.startswith("operations/") or operation_name.startswith("models/")
):
operation_name = f"operations/{operation_name}"
base = (
getattr(endpoint, "base_url", None) or "https://generativelanguage.googleapis.com"
).rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
upstream_url = f"{base}/v1beta/{operation_name}:cancel"
auth_info = await get_provider_auth(endpoint, key)
if auth_info:
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
response = await _try_rust_cancel_response(
method="POST",
url=upstream_url,
request_headers=headers,
body={},
content_type="application/json",
)
if response.status_code >= 400:
return response
else:
raise HTTPException(
status_code=400,
detail=f"Cancel not supported for provider format: {provider_format}",
)
now = datetime.now(timezone.utc)
task.status = VideoStatus.CANCELLED.value
task.completed_at = getattr(task, "completed_at", None) or now
task.updated_at = now
# Void Usage (no charge)
try:
voided = UsageService.finalize_void(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
finalized_at=task.completed_at,
)
if not voided:
logger.warning(
"Skip voiding video usage because billing is already terminal: task_id={} request_id={}",
getattr(task, "id", task_id),
getattr(task, "request_id", None),
)
except Exception as exc:
logger.warning(
"Failed to void usage for cancelled task={}: {}",
getattr(task, "id", task_id),
str(exc),
)
self.db.commit()
return None

View File

@@ -0,0 +1,38 @@
from __future__ import annotations
from typing import Any
from src.services.task.core.schema import TaskStatusResult
from src.services.task.video.operations import VideoTaskOperationsService
class TaskVideoFacadeService:
"""视频任务门面服务(向后兼容 TaskService 的视频公开方法)。"""
def __init__(self, video_ops: VideoTaskOperationsService) -> None:
self._video_ops = video_ops
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
return await self._video_ops.poll(task_id, user_id=user_id)
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
return await self._video_ops.poll_now(task_id, user_id=user_id)
async def cancel(
self,
task_id: str,
*,
user_id: str,
original_headers: dict[str, str] | None = None,
) -> Any:
return await self._video_ops.cancel(
task_id,
user_id=user_id,
original_headers=original_headers,
)
async def finalize_video_task(self, task: Any) -> bool:
return await self._video_ops.finalize_video_task(task)
async def finalize(self, task_id: str) -> bool:
return await self._video_ops.finalize(task_id)

View File

@@ -0,0 +1,134 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import VideoTask
from src.services.task.core.exceptions import TaskNotFoundError
from src.services.task.core.schema import TaskStatusResult
from src.services.task.video.billing import VideoTaskBillingService
from src.services.task.video.cancel import VideoTaskCancelService
class VideoTaskOperationsService:
"""视频任务相关应用服务(轮询/取消/终态结算)。"""
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
self._billing_ops = VideoTaskBillingService(db)
self._cancel_ops = VideoTaskCancelService(db)
def _extract_short_id(self, task_id: str) -> str:
# Keep the parsing rule consistent with handlers:
# - models/{model}/operations/{short_id}
# - operations/{short_id}
# - {short_id}
return task_id.rsplit("/", 1)[-1] if "/" in task_id else task_id
def _get_video_task_for_user(self, task_id: str, *, user_id: str) -> Any:
"""
Resolve a video task by:
- internal UUID (VideoTask.id)
- external operation id (VideoTask.short_id)
"""
task = (
self.db.query(VideoTask)
.filter(VideoTask.id == task_id, VideoTask.user_id == user_id)
.first()
)
if task:
return task
short_id = self._extract_short_id(task_id)
task = (
self.db.query(VideoTask)
.filter(VideoTask.short_id == short_id, VideoTask.user_id == user_id)
.first()
)
if not task:
raise TaskNotFoundError(task_id)
return task
async def poll(self, task_id: str, *, user_id: str) -> TaskStatusResult:
"""Read task status from DB (does not trigger polling)."""
task = self._get_video_task_for_user(task_id, user_id=user_id)
result_url = None
if getattr(task, "status", None) == "completed":
result_url = getattr(task, "video_url", None)
error_message = None
if getattr(task, "status", None) == "failed":
error_message = getattr(task, "error_message", None) or getattr(
task, "error_code", None
)
return TaskStatusResult(
task_id=str(getattr(task, "id", task_id)),
status=str(getattr(task, "status", "unknown")),
progress_percent=int(getattr(task, "progress_percent", 0) or 0),
result_url=result_url,
error_message=str(error_message) if error_message else None,
provider_id=(
str(getattr(task, "provider_id", None))
if getattr(task, "provider_id", None)
else None
),
provider_name=(
str(getattr(task, "provider_name", None))
if getattr(task, "provider_name", None)
else None
),
endpoint_id=(
str(getattr(task, "endpoint_id", None))
if getattr(task, "endpoint_id", None)
else None
),
key_id=str(getattr(task, "key_id", None)) if getattr(task, "key_id", None) else None,
)
async def poll_now(self, task_id: str, *, user_id: str) -> TaskStatusResult:
"""
Trigger a single polling attempt (best-effort), then return latest DB status.
Note: this uses the poller adapter's single-task method and may hold a DB
connection during the upstream HTTP request; keep usage low.
"""
from src.services.task.video.poller_adapter import VideoTaskPollerAdapter
task = self._get_video_task_for_user(task_id, user_id=user_id)
adapter = VideoTaskPollerAdapter()
await adapter.poll_single_task(self.db, task, redis_client=self.redis)
self.db.commit()
return await self.poll(task_id, user_id=user_id)
async def cancel(
self,
task_id: str,
*,
user_id: str,
original_headers: dict[str, str] | None = None,
) -> Any:
from fastapi import HTTPException
try:
task = self._get_video_task_for_user(task_id, user_id=user_id)
except TaskNotFoundError:
raise HTTPException(status_code=404, detail="Video task not found")
return await self._cancel_ops.cancel_task(
task=task,
task_id=task_id,
original_headers=original_headers,
)
async def finalize_video_task(self, task: Any) -> bool:
return await self._billing_ops.finalize_video_task(task)
async def finalize(self, task_id: str) -> bool:
"""Finalize a task by internal id (best-effort)."""
task = self.db.query(VideoTask).filter(VideoTask.id == task_id).first()
if not task:
return False
return await self.finalize_video_task(task)

View File

@@ -0,0 +1,656 @@
"""
Video task poller adapter.
Implements the video-specific poll/normalize/update logic used by TaskPollerService.
优化HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy.orm import Session
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.core.provider_auth_types import ProviderAuthInfo
from src.core.video_utils import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.provider.auth import get_provider_auth
from src.services.provider.provider_context import resolve_provider_proxy
@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
proxy_config: dict[str, Any] | None = None
delegate_config: dict[str, Any] | None = None
proxy_snapshot: Any = None
# 永久性错误指示词(用于降级判断,不应重试)
_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
VideoTaskFinalizeFn = Callable[[Session, VideoTask, Any | None], Awaitable[None]]
async def _default_finalize_video_task(
db: Session,
task: VideoTask,
redis_client: Any | None,
) -> None:
"""默认终态结算逻辑(延迟导入,避免 task 模块循环依赖)。"""
from src.services.task.video.operations import VideoTaskOperationsService
await VideoTaskOperationsService(db, redis_client=redis_client).finalize_video_task(task)
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, finalize_video_task_fn: VideoTaskFinalizeFn | None = None) -> None:
self._openai_normalizer = OpenAINormalizer()
self._gemini_normalizer = GeminiNormalizer()
self._finalize_video_task = finalize_video_task_fn or _default_finalize_video_task
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 {}", 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)
proxy_config, delegate_config, proxy_snapshot = await self._build_transport_context(
endpoint=endpoint,
key=key,
)
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,
proxy_config=proxy_config,
delegate_config=delegate_config,
proxy_snapshot=proxy_snapshot,
)
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 {} disappeared during poll update", task_id)
return
if task.status in {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}:
logger.debug(
"Skip poll update for terminal task {} with status {}",
task_id,
task.status,
)
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
if result.video_duration_seconds is not None:
task.video_duration_seconds = result.video_duration_seconds
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 self._finalize_video_task(db, task, redis_client)
except Exception as exc:
logger.exception(
"Failed to record video usage for task={}: {}",
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 {}: {}", 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)
payload = await self._try_rust_poll_payload(ctx=ctx, url=url)
if payload is None:
raise PollHTTPError(503, "Video 轮询仅支持 Rust executor")
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={} external_id={} url={}",
ctx.task_id,
ctx.external_task_id,
url,
)
payload = await self._try_rust_poll_payload(ctx=ctx, url=url)
if payload is None:
raise PollHTTPError(503, "Video 轮询仅支持 Rust executor")
return self._gemini_normalizer.video_poll_to_internal(payload)
async def _try_rust_poll_payload(
self,
*,
ctx: VideoPollContext,
url: str,
) -> dict[str, Any] | None:
import httpx
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return None
try:
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=f"video-poll-{ctx.task_id}",
candidate_id=None,
provider_name=ctx.provider_api_format.split(":", 1)[0],
provider_id="",
endpoint_id="",
key_id="",
method="GET",
url=url,
headers=dict(ctx.headers),
body=ExecutionPlanBody(),
stream=False,
provider_api_format=ctx.provider_api_format,
client_api_format=ctx.provider_api_format,
model_name="video-poll",
proxy=ctx.proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=300_000,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError, ValueError) as exc:
logger.warning(
"[VideoPoller] Rust poll fallback task={} url={} error={}",
ctx.task_id,
url,
sanitize_error_message(str(exc)),
)
return None
if result.status_code >= 400:
response_text = ""
if result.response_json is not None:
response_text = json.dumps(result.response_json, ensure_ascii=False)
elif result.response_body_bytes is not None:
response_text = result.response_body_bytes.decode("utf-8", errors="replace")
raise PollHTTPError(
result.status_code,
self._extract_error_message(response_text, result.status_code),
)
if isinstance(result.response_json, dict):
return result.response_json
if result.response_body_bytes is not None:
try:
payload = json.loads(result.response_body_bytes.decode("utf-8"))
if isinstance(payload, dict):
return payload
except Exception:
logger.warning(
"[VideoPoller] Rust poll returned non-json body task={} url={}",
ctx.task_id,
url,
)
return None
# ==================== 旧版方法(保留兼容性)====================
async def poll_single_task(
self, db: Session, task: VideoTask, *, redis_client: Any | None
) -> None:
"""
兼容入口:复用三阶段轮询流程,避免维护重复逻辑。
"""
if task.status in {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}:
logger.debug(
"Skip legacy poll for terminal task {} with status {}", task.id, task.status
)
return
ctx_or_result = await self.prepare_poll_context(db, task)
if isinstance(ctx_or_result, InternalVideoPollResult):
await self.update_task_after_poll(
task_id=task.id,
result=ctx_or_result,
ctx=None,
redis_client=redis_client,
)
return
ctx = ctx_or_result
error_exception: Exception | None = None
try:
result = await self.poll_task_http(ctx)
except Exception as http_exc:
error_exception = http_exc
result = InternalVideoPollResult(
status=None, # type: ignore[arg-type]
error_message=str(http_exc),
)
await self.update_task_after_poll(
task_id=task.id,
result=result,
ctx=ctx,
redis_client=redis_client,
error_exception=error_exception,
)
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:
ctx_or_result = await self.prepare_poll_context(db, task)
if isinstance(ctx_or_result, InternalVideoPollResult):
return ctx_or_result
return await self.poll_task_http(ctx_or_result)
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,
header_rules=getattr(endpoint, "header_rules", None),
)
if auth_info:
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
return headers
async def _build_transport_context(
self,
*,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
) -> tuple[dict[str, Any] | None, dict[str, Any] | None, Any]:
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import ExecutionProxySnapshot
try:
effective_proxy = resolve_effective_proxy(
resolve_provider_proxy(endpoint=endpoint, key=key),
getattr(key, "proxy", None),
)
if not effective_proxy or not effective_proxy.get("enabled", True):
effective_proxy = await get_system_proxy_config_async()
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
proxy_url: str | None = None
if effective_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(effective_proxy)
proxy_info = await resolve_proxy_info_async(effective_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
return effective_proxy, delegate_cfg, proxy_snapshot
except Exception as exc:
logger.warning(
"[VideoPoller] Failed to build transport context endpoint={} key={}: {}",
getattr(endpoint, "id", None),
getattr(key, "id", None),
sanitize_error_message(str(exc)),
)
return None, None, None
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])