mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
0
_deprecated_py_src/services/task/video/__init__.py
Normal file
0
_deprecated_py_src/services/task/video/__init__.py
Normal file
324
_deprecated_py_src/services/task/video/billing.py
Normal file
324
_deprecated_py_src/services/task/video/billing.py
Normal 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)
|
||||
266
_deprecated_py_src/services/task/video/cancel.py
Normal file
266
_deprecated_py_src/services/task/video/cancel.py
Normal 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
|
||||
38
_deprecated_py_src/services/task/video/facade.py
Normal file
38
_deprecated_py_src/services/task/video/facade.py
Normal 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)
|
||||
134
_deprecated_py_src/services/task/video/operations.py
Normal file
134
_deprecated_py_src/services/task/video/operations.py
Normal 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)
|
||||
656
_deprecated_py_src/services/task/video/poller_adapter.py
Normal file
656
_deprecated_py_src/services/task/video/poller_adapter.py
Normal 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])
|
||||
Reference in New Issue
Block a user