"""Extend API orchestration handler. Mirrors ``RetakeHandler``'s dual API/local dispatch. Extend appends (``mode="end"``) or prepends (``mode="start"``) freshly generated frames to a source video. The cloud ``/v1/extend`` endpoint and the local PyTorch wrapper share this entry point. """ from __future__ import annotations import uuid from datetime import datetime from pathlib import Path from tempfile import TemporaryDirectory from threading import RLock from api_types import ( ExtendMode, ExtendRequest, ExtendResponse, RetakeCancelledResponse, RetakeExtendModel, RetakePayloadResponse, RetakeVideoResponse, TargetResolution, ) from _routes._errors import HTTPError from api_model_specs import FORCED_API_MODEL_MAP from handlers.base import StateHandlerBase from handlers.generation_handler import GenerationHandler from handlers.rolling_extend import assemble_rolling_result, extract_context, plan_rolling_extend from handlers.pipelines_handler import PipelinesHandler from handlers.text_handler import TextHandler from handlers.video_resolution import ( TIME_FACTOR, correct_frame_count, read_source_metadata, resolve_target_resolution, validate_source_video_path, ) from runtime_config.ltx_capabilities import local_caps, supports from runtime_config.model_download_specs import resolve_active_ltx_model_id from runtime_config.runtime_config import RuntimeConfig from services.generation_interrupt import GenerationCancelledError, is_cancel_exception from services.ltx_api_client.ltx_api_client import LTXAPIClientError from services.interfaces import LTXAPIClient from state.app_state_types import AppState from state.app_settings import should_video_generate_with_ltx_api # Cloud caps a single extend at 20s; mirror it locally. Minimum is 2s (cloud min). _MIN_DURATION = 2.0 _MAX_DURATION = 20.0 class ExtendHandler(StateHandlerBase): def __init__( self, state: AppState, lock: RLock, ltx_api_client: LTXAPIClient, config: RuntimeConfig, generation_handler: GenerationHandler, pipelines_handler: PipelinesHandler, text_handler: TextHandler, ) -> None: super().__init__(state, lock, config) self._ltx_api_client = ltx_api_client self._generation = generation_handler self._pipelines = pipelines_handler self._text = text_handler def run(self, req: ExtendRequest) -> ExtendResponse: video_path = req.video_path duration = req.duration prompt = req.prompt mode = req.mode if duration < _MIN_DURATION: raise HTTPError(400, f"duration must be at least {int(_MIN_DURATION)} seconds") if duration > _MAX_DURATION: raise HTTPError(400, f"duration must be at most {int(_MAX_DURATION)} seconds") video_file = validate_source_video_path(video_path) if should_video_generate_with_ltx_api( force_api_generations=self.config.force_api_generations, settings=self.state.app_settings, ): # The cloud preserves source resolution (no resolution param); resolution # selection is local-only. return self._run_api_extend( video_file=video_file, duration=duration, prompt=prompt, mode=mode, model=req.model, ) model_id = resolve_active_ltx_model_id( self.models_dir, self.state.app_settings.active_ltx_model_id ) if model_id is None: raise HTTPError(409, "NO_DOWNLOADED_LTX_MODEL") if not supports(local_caps(model_id), "extend"): raise HTTPError( 409, "Extend is not supported for the active LTX model.", code="UNSUPPORTED_EXTEND", ) return self._run_local_extend( video_file=video_file, duration=duration, prompt=prompt, mode=mode, resolution=req.resolution ) def _run_api_extend( self, *, video_file: Path, duration: float, prompt: str, mode: ExtendMode, model: RetakeExtendModel, ) -> ExtendResponse: api_key = self.state.app_settings.ltx_api_key if not api_key: raise HTTPError(400, "LTX API key not configured. Set it in Settings.") with self._generation.reserved_generation_start(): # Drive the generation state machine so the result is recoverable via # /api/generation/progress if the page unmounts mid-generation (mirrors retake). try: self._generation.start_api_generation(uuid.uuid4().hex[:8]) except RuntimeError as exc: # Lost a race with a concurrent generation between the check above and here; # surface the same 409 as the guard, not a bare 500 (the running generation # owns the state now, so don't fail_generation). raise HTTPError(409, str(exc)) from exc try: self._generation.update_progress("inference", 55, None, None) result = self._ltx_api_client.extend( api_key=api_key, video_path=str(video_file), duration=duration, prompt=prompt, mode=mode, model=FORCED_API_MODEL_MAP[model], ) if result.video_bytes is not None: self._generation.update_progress("downloading_output", 85, None, None) output = self.config.outputs_dir / f"extend_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:8]}.mp4" try: with open(output, "wb") as out: out.write(result.video_bytes) self._generation.update_progress("complete", 100, None, None) self._generation.complete_generation(str(output)) except Exception: output.unlink(missing_ok=True) # don't strand a half-written .mp4 raise return RetakeVideoResponse(status="complete", video_path=str(output)) if result.result_payload is not None: self._generation.update_progress("complete", 100, None, None) self._generation.complete_generation(None) return RetakePayloadResponse(status="complete", result=result.result_payload) raise HTTPError(500, "Extend API returned no result") except LTXAPIClientError as exc: self._generation.fail_generation(exc.detail) raise HTTPError(exc.status_code, exc.detail) from exc except HTTPError as exc: self._generation.fail_generation(exc.detail) raise except Exception as exc: self._generation.fail_generation(str(exc)) raise def _run_local_extend( self, *, video_file: Path, duration: float, prompt: str, mode: ExtendMode, resolution: TargetResolution | None, ) -> ExtendResponse: with self._generation.reserved_generation_start(): fps, source_width, source_height, source_frames = read_source_metadata(str(video_file)) target_frames = correct_frame_count(source_frames) extend_frames = self._duration_to_extend_frames(duration, fps) try: rolling_plan = plan_rolling_extend(target_frames, extend_frames, fps) except ValueError as exc: raise HTTPError(400, str(exc)) from exc target_width, target_height = resolve_target_resolution(resolution, source_width, source_height) try: self._text.prepare_text_encoding(prompt, enhance_prompt=False) except RuntimeError as exc: raise HTTPError(400, str(exc)) from exc generation_id = uuid.uuid4().hex[:8] seed = self._resolve_seed() output_path = self.config.outputs_dir / f"extend_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{generation_id}.mp4" try: pipeline_state = self._pipelines.load_retake_pipeline(distilled=True) self._generation.start_generation(generation_id) self._generation.update_progress("loading_model", 5, 0, 1) self._generation.update_progress("inference", 15, 0, 1) if rolling_plan.rolling: with TemporaryDirectory(prefix="rolling-extend-", dir=self.config.outputs_dir) as temp: context_path = Path(temp) / "context.mp4" window_output = Path(temp) / "generated-window.mp4" extract_context( video_file, context_path, mode=mode, fps=fps, source_frames=target_frames, context_frames=rolling_plan.context_frames, ) pipeline_state.pipeline.extend( video_path=str(context_path), prompt=prompt, extend_frames=extend_frames, mode=mode, seed=seed, output_path=str(window_output), negative_prompt=self.config.default_negative_prompt, regenerate_audio=True, enhance_prompt=False, distilled=True, target_width=target_width, target_height=target_height, target_frames=rolling_plan.context_frames, ) self._generation.update_progress("assembling_video", 95, 0, 1) assemble_rolling_result( video_file, window_output, output_path, mode=mode, fps=fps, source_frames=target_frames, context_frames=rolling_plan.context_frames, extend_frames=extend_frames, ) else: pipeline_state.pipeline.extend( video_path=str(video_file), prompt=prompt, extend_frames=extend_frames, mode=mode, seed=seed, output_path=str(output_path), negative_prompt=self.config.default_negative_prompt, regenerate_audio=True, enhance_prompt=False, distilled=True, target_width=target_width, target_height=target_height, target_frames=target_frames, ) # Denoiser interrupt cannot abort VAE decode / ffmpeg; a Stop after the last # denoise step still finishes encode, then this check drops the file. if self._generation.is_generation_cancelled(): output_path.unlink(missing_ok=True) raise GenerationCancelledError() self._generation.update_progress("complete", 100, 1, 1) self._generation.complete_generation(str(output_path)) return RetakeVideoResponse(status="complete", video_path=str(output_path)) except HTTPError: self._generation.fail_generation("Extend generation failed") raise except Exception as exc: self._generation.fail_generation(str(exc)) if is_cancel_exception(exc): return RetakeCancelledResponse(status="cancelled") raise HTTPError(500, f"Generation error: {exc}") from exc finally: self._text.clear_api_embeddings() @staticmethod def _duration_to_extend_frames(duration: float, fps: float) -> int: # Snap up to the nearest multiple of TIME_FACTOR so the output stays 8k+1: the # source is 8k+1, so adding a multiple of 8 keeps the padded latent integer-sized. frames = round(duration * fps) return ((frames + TIME_FACTOR - 1) // TIME_FACTOR) * TIME_FACTOR