diff --git a/apps/dreamverse/arch.md b/apps/dreamverse/arch.md index eb532f17d8..c94cc96255 100644 --- a/apps/dreamverse/arch.md +++ b/apps/dreamverse/arch.md @@ -299,18 +299,47 @@ There are three related prompt paths in the current system: ## Initial Image And Segment Handling -The frontend currently sends `initial_image` as part of session init or +The frontend sends `initial_image` and, for first/last frame mode, +`last_frame_image` as part of `session_init_v2`, `project_init_v1`, or `simple_generate`. The server: -- validates and persists the image -- uses it only for segment 1 when present +- validates and persists the images +- uses `initial_image` only for segment 1 when present - keeps continuation state for later segments in the GPU worker This means the runtime, not the frontend, decides how segment 1 image conditioning and later continuation conditioning are applied. +## Creation Studio Config + +The lobby creation studio sends model, mode, aspect ratio, resolution, and +duration with session init. The server parses these fields into a per-session +creation config and echoes the resolved values back on `gpu_assigned` and +`ltx2_stream_start` as `creation_config`. + +Incoming fields on `session_init_v2` and `project_init_v1`: + +- `generation_mode`: `t2va`, `fl2va`, or `ref2va` (canonical upstream IDs from #1834) +- `model_id`: `fast-ltx2`, `fast-ltx23`, or `fast-h3` +- `aspect_ratio`: one of `21:9`, `16:9`, `4:3`, `1:1`, `3:4`, `9:16` +- `resolution`: one of `480p`, `720p`, `1080p`, `4k` +- `duration_sec`: `5`, `10`, or `15` +- `initial_image`: optional image payload for reference / first-frame modes +- `last_frame_image`: optional image payload for first/last frame mode + +Echoed `creation_config` includes the resolved frame size, +`num_frames`, and `generation_segment_cap` derived from `duration_sec`. + +Mode validation: + +- `ref2va` requires `initial_image` +- `fl2va` requires both `initial_image` and `last_frame_image` + +Per-step generation uses the resolved `frame_width`, `frame_height`, and +`num_frames` from the session creation config. + ## Websocket Contract The websocket is the main integration surface between UI and runtime. diff --git a/apps/dreamverse/dreamverse/creation_capabilities.py b/apps/dreamverse/dreamverse/creation_capabilities.py new file mode 100644 index 0000000000..afa8facb75 --- /dev/null +++ b/apps/dreamverse/dreamverse/creation_capabilities.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from dreamverse.config import MODEL_REGISTRY + +# Canonical upstream wire IDs. FL2VA is tracked in #1834 but not wired on Dreamverse +# streaming backends yet. +LTX_LOBBY_GENERATION_MODES = frozenset({"t2va", "ref2va"}) +H3_LOBBY_GENERATION_MODES = frozenset({"t2va", "ref2va"}) + +LTX_LOBBY_ASPECT_RATIOS = frozenset({"21:9", "16:9", "4:3", "1:1", "3:4", "9:16"}) + +# Realtime FastLTX serving is validated through 1080p-class outputs; 4K is rejected +# until the runtime path is tested on Dreamverse GPUs. +LTX_LOBBY_RESOLUTIONS = frozenset({"480p", "720p", "1080p"}) + +# FastH3 serves a fixed 768x1344 (16:9-class) output; lobby resolution is nominal. +H3_LOBBY_ASPECT_RATIOS = frozenset({"16:9"}) +H3_LOBBY_RESOLUTIONS = frozenset({"720p"}) + +LOBBY_DURATION_SEC = frozenset({5, 10, 15}) + +UNSUPPORTED_GENERATION_MODE_MESSAGES = { + "fl2va": "First/last frame mode (FL2VA) is not supported yet.", +} + + +@dataclass(frozen=True) +class ModelCreationCapabilities: + generation_modes: frozenset[str] + aspect_ratios: frozenset[str] + resolutions: frozenset[str] + duration_sec: frozenset[int] + unsupported_generation_modes: frozenset[str] = frozenset({"fl2va"}) + + def as_dict(self) -> dict[str, object]: + unsupported = { + mode: UNSUPPORTED_GENERATION_MODE_MESSAGES[mode] + for mode in sorted(self.unsupported_generation_modes) + if mode in UNSUPPORTED_GENERATION_MODE_MESSAGES + } + return { + "generation_modes": sorted(self.generation_modes), + "aspect_ratios": sorted(self.aspect_ratios), + "resolutions": sorted(self.resolutions), + "duration_sec": sorted(self.duration_sec), + "unsupported_generation_modes": unsupported, + "reference_assets": { + "mime_types": ["image/png", "image/jpeg", "image/webp"], + "max_bytes": 15 * 1024 * 1024, + }, + } + + +LTX_MODEL_CREATION_CAPABILITIES = ModelCreationCapabilities( + generation_modes=LTX_LOBBY_GENERATION_MODES, + aspect_ratios=LTX_LOBBY_ASPECT_RATIOS, + resolutions=LTX_LOBBY_RESOLUTIONS, + duration_sec=LOBBY_DURATION_SEC, +) + +H3_MODEL_CREATION_CAPABILITIES = ModelCreationCapabilities( + generation_modes=H3_LOBBY_GENERATION_MODES, + aspect_ratios=H3_LOBBY_ASPECT_RATIOS, + resolutions=H3_LOBBY_RESOLUTIONS, + duration_sec=LOBBY_DURATION_SEC, +) + +MODEL_CREATION_CAPABILITIES: dict[str, ModelCreationCapabilities] = { + "fast-ltx2": LTX_MODEL_CREATION_CAPABILITIES, + "fast-ltx23": LTX_MODEL_CREATION_CAPABILITIES, + "fast-h3": H3_MODEL_CREATION_CAPABILITIES, +} + + +def capabilities_for_model(model_id: str) -> ModelCreationCapabilities: + if model_id not in MODEL_REGISTRY: + raise ValueError(f"Unknown model_id: {model_id}") + return MODEL_CREATION_CAPABILITIES.get(model_id, LTX_MODEL_CREATION_CAPABILITIES) + + +def lobby_capabilities_as_dict() -> dict[str, object]: + model_ids = sorted(MODEL_REGISTRY.keys()) + models = {model_id: capabilities_for_model(model_id).as_dict() for model_id in model_ids} + union_modes: set[str] = set() + union_aspects: set[str] = set() + union_resolutions: set[str] = set() + union_durations: set[int] = set() + for caps in MODEL_CREATION_CAPABILITIES.values(): + union_modes.update(caps.generation_modes) + union_aspects.update(caps.aspect_ratios) + union_resolutions.update(caps.resolutions) + union_durations.update(caps.duration_sec) + return { + "model_ids": model_ids, + "models": models, + "generation_modes": sorted(union_modes), + "aspect_ratios": sorted(union_aspects), + "resolutions": sorted(union_resolutions), + "duration_sec": sorted(union_durations), + "unsupported_generation_modes": dict(UNSUPPORTED_GENERATION_MODE_MESSAGES), + "reference_assets": { + "mime_types": ["image/png", "image/jpeg", "image/webp"], + "max_bytes": 15 * 1024 * 1024, + }, + } + + +# Backward-compatible alias used in tests. +LOBBY_CREATION_CAPABILITIES = lobby_capabilities_as_dict() + + +def validate_lobby_creation_config( + *, + model_id: str, + generation_mode: str, + aspect_ratio: str, + resolution: str, + duration_sec: int, +) -> None: + if model_id not in MODEL_REGISTRY: + raise ValueError(f"Unknown model_id: {model_id}") + + caps = capabilities_for_model(model_id) + + if generation_mode in caps.unsupported_generation_modes: + raise ValueError(UNSUPPORTED_GENERATION_MODE_MESSAGES[generation_mode]) + if generation_mode not in caps.generation_modes: + raise ValueError(f"Unsupported generation_mode: {generation_mode}") + + if aspect_ratio not in caps.aspect_ratios: + raise ValueError(f"Unsupported aspect_ratio: {aspect_ratio}") + if resolution not in caps.resolutions: + raise ValueError(f"Unsupported resolution: {resolution}") + if duration_sec not in caps.duration_sec: + raise ValueError("duration_sec must be 5, 10, or 15.") diff --git a/apps/dreamverse/dreamverse/generation_contracts.py b/apps/dreamverse/dreamverse/generation_contracts.py index 782b23304d..cde10e990a 100644 --- a/apps/dreamverse/dreamverse/generation_contracts.py +++ b/apps/dreamverse/dreamverse/generation_contracts.py @@ -36,6 +36,10 @@ def generate_step( segment_idx: int, image_path: str | None, reset_conditioning: bool, + *, + frame_width: int | None = None, + frame_height: int | None = None, + num_frames: int | None = None, ) -> StepResult: ... diff --git a/apps/dreamverse/dreamverse/generation_worker.py b/apps/dreamverse/dreamverse/generation_worker.py index 1585940ae1..b52508f8b7 100644 --- a/apps/dreamverse/dreamverse/generation_worker.py +++ b/apps/dreamverse/dreamverse/generation_worker.py @@ -80,6 +80,10 @@ def generate_step( segment_idx: int, image_path: str | None, reset_conditioning: bool, + *, + frame_width: int | None = None, + frame_height: int | None = None, + num_frames: int | None = None, ) -> StepResult: """Generate one segment through the selected model backend.""" return self._require_backend().generate_step( @@ -87,6 +91,9 @@ def generate_step( segment_idx, image_path, reset_conditioning, + frame_width=frame_width, + frame_height=frame_height, + num_frames=num_frames, ) def warmup(self, prompt: str) -> dict[str, float]: diff --git a/apps/dreamverse/dreamverse/gpu_pool.py b/apps/dreamverse/dreamverse/gpu_pool.py index a524fe16c7..ee30be0cb5 100644 --- a/apps/dreamverse/dreamverse/gpu_pool.py +++ b/apps/dreamverse/dreamverse/gpu_pool.py @@ -189,6 +189,9 @@ def handle_command(cmd: Command): segment_idx, image_path=payload.image_path, reset_conditioning=payload.reset_conditioning, + frame_width=payload.frame_width, + frame_height=payload.frame_height, + num_frames=payload.num_frames, ) head_trim_frames = step_result.head_trim_frames head_trim_audio_frames = step_result.head_trim_audio_frames @@ -753,6 +756,10 @@ async def user_step( segment_idx: int = 1, image_path: str | None = None, reset_conditioning: bool = False, + *, + frame_width: int | None = None, + frame_height: int | None = None, + num_frames: int | None = None, ) -> dict[str, float]: """Execute a generation step for a specific user. @@ -766,6 +773,9 @@ async def user_step( segment_idx=segment_idx, image_path=image_path, reset_conditioning=bool(reset_conditioning), + frame_width=frame_width, + frame_height=frame_height, + num_frames=num_frames, ) response = await self._send_command_tagged(Command(CommandType.USER_STEP, payload=payload, user_id=user_id), timeout=1800.0) diff --git a/apps/dreamverse/dreamverse/ltx2_generation.py b/apps/dreamverse/dreamverse/ltx2_generation.py index 59f74680c1..ecd9caf5ed 100644 --- a/apps/dreamverse/dreamverse/ltx2_generation.py +++ b/apps/dreamverse/dreamverse/ltx2_generation.py @@ -454,6 +454,10 @@ def generate_step( segment_idx: int, image_path: str | None, reset_conditioning: bool, + *, + frame_width: int | None = None, + frame_height: int | None = None, + num_frames: int | None = None, ) -> StepResult: """Execute one generation step; snapshot state for the next segment.""" timings: dict = {} @@ -464,9 +468,9 @@ def generate_step( prompt=prompt, negative_prompt="", save_video=False, - height=FRAME_HEIGHT, - width=FRAME_WIDTH, - num_frames=NUM_FRAMES, + height=frame_height or FRAME_HEIGHT, + width=frame_width or FRAME_WIDTH, + num_frames=num_frames or NUM_FRAMES, fps=24, num_inference_steps=NUM_INFERENCE_STEPS, guidance_scale=1.0, diff --git a/apps/dreamverse/dreamverse/main.py b/apps/dreamverse/dreamverse/main.py index 65d4f99640..f3b3c8313e 100644 --- a/apps/dreamverse/dreamverse/main.py +++ b/apps/dreamverse/dreamverse/main.py @@ -33,6 +33,7 @@ prompt_config_router, curated_presets_router, ) +from dreamverse.routes.creation import creation_router from dreamverse.session.controller import SessionController @@ -92,6 +93,7 @@ async def lifespan(app: FastAPI): app.include_router(build_health_router(lambda: runtime.gpu_pool)) app.include_router(internal_monitor_router) app.include_router(prompt_config_router) +app.include_router(creation_router) if DEVTOOLS_ENABLED: app.include_router(curated_presets_router) diff --git a/apps/dreamverse/dreamverse/minimax_h3_generation.py b/apps/dreamverse/dreamverse/minimax_h3_generation.py index 6c99b40b72..e1af5da99b 100644 --- a/apps/dreamverse/dreamverse/minimax_h3_generation.py +++ b/apps/dreamverse/dreamverse/minimax_h3_generation.py @@ -200,6 +200,10 @@ def generate_step( segment_idx: int, image_path: str | None, reset_conditioning: bool, + *, + frame_width: int | None = None, + frame_height: int | None = None, + num_frames: int | None = None, ) -> StepResult: """Generate one synchronized FastH3 segment and retain its last frame. @@ -207,6 +211,7 @@ def generate_step( conditioned frame and its matching audio duration are trimmed before streaming so adjacent segments do not duplicate media. """ + del frame_width, frame_height, num_frames if self.generator is None: raise RuntimeError("FastH3 generator is not initialized.") conditioning_image, uses_continuation = self._select_conditioning_image( diff --git a/apps/dreamverse/dreamverse/mock_server.py b/apps/dreamverse/dreamverse/mock_server.py index 1d788c88e2..204e387602 100644 --- a/apps/dreamverse/dreamverse/mock_server.py +++ b/apps/dreamverse/dreamverse/mock_server.py @@ -31,6 +31,8 @@ from dreamverse._deps import require_dreamverse_runtime_deps from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP +from dreamverse.creation_capabilities import lobby_capabilities_as_dict +from dreamverse.session_creation_config import parse_session_creation_config, validate_generation_mode_assets from dreamverse.session_init_image import cleanup_session_init_image, persist_session_init_image LATENCY_MS = 200 @@ -225,6 +227,11 @@ async def prompt_system_config(): } +@app.get("/creation-capabilities") +async def creation_capabilities(): + return lobby_capabilities_as_dict() + + @app.get("/curated-presets") async def curated_presets(): presets = [ @@ -290,6 +297,8 @@ async def websocket_endpoint(websocket: WebSocket): send_lock = asyncio.Lock() stop_event = asyncio.Event() session_init_image = None + session_last_frame_image = None + session_creation_config = None async def ws_send_json(payload: dict) -> None: async with send_lock: @@ -348,6 +357,7 @@ async def session_timeout() -> None: try: session_init_image = persist_session_init_image(init_data.get("initial_image")) + session_last_frame_image = persist_session_init_image(init_data.get("last_frame_image")) except ValueError as exc: await ws_send_json({ "type": "error", @@ -356,13 +366,31 @@ async def session_timeout() -> None: await websocket.close(code=1003, reason="Invalid initial image") return + try: + session_creation_config = parse_session_creation_config(init_data) + validate_generation_mode_assets( + session_creation_config.generation_mode, + has_initial_image=session_init_image is not None, + has_last_frame_image=session_last_frame_image is not None, + ) + except ValueError as exc: + await ws_send_json({ + "type": "error", + "message": str(exc), + }) + await websocket.close(code=1003, reason="Invalid creation config") + return + timeout_task = asyncio.create_task(session_timeout()) - await ws_send_json({ + gpu_assigned_payload: dict[str, object] = { "type": "gpu_assigned", "gpu_id": 0, "session_timeout": SESSION_TIMEOUT_SECONDS, - }) + } + if session_creation_config is not None: + gpu_assigned_payload["creation_config"] = session_creation_config.as_dict() + await ws_send_json(gpu_assigned_payload) raw_prompt_queue: asyncio.Queue[PromptSubmission] = asyncio.Queue() ready_prompt_queue: asyncio.Queue[ReadyPrompt] = asyncio.Queue() @@ -391,8 +419,16 @@ def replace_session_image(initial_image_payload: object) -> None: if previous_session_image is not None: cleanup_session_init_image(previous_session_image) + def replace_last_frame_image(last_frame_payload: object) -> None: + nonlocal session_last_frame_image + next_last_frame_image = persist_session_init_image(last_frame_payload) + previous_last_frame_image = session_last_frame_image + session_last_frame_image = next_last_frame_image + if previous_last_frame_image is not None: + cleanup_session_init_image(previous_last_frame_image) + async def send_stream_start(seed_reason: str) -> None: - await ws_send_json({ + stream_start_payload: dict[str, object] = { "type": "ltx2_stream_start", "total_segments": len(curated_prompts), "preset_id": preset_id, @@ -400,8 +436,15 @@ async def send_stream_start(seed_reason: str) -> None: "live_mode": True, "loop_generation_enabled": loop_generation_enabled, "loop_iteration": loop_iteration, - "generation_segment_cap": 0, - }) + "generation_segment_cap": ( + session_creation_config.generation_segment_cap + if session_creation_config is not None + else GENERATION_SEGMENT_CAP + ), + } + if session_creation_config is not None: + stream_start_payload["creation_config"] = session_creation_config.as_dict() + await ws_send_json(stream_start_payload) if seed_reason == "init": await ws_send_json({ "type": "seed_prompts_updated", @@ -509,6 +552,7 @@ async def apply_project_init_payload(payload: dict[str, object], ) -> bool: nonlocal project_active nonlocal project_stream_started nonlocal pending_project_end + nonlocal session_creation_config next_initial_rollout_prompt = str(payload.get("initial_rollout_prompt") or "").strip() next_preset_id = str(payload.get("preset_id") or "").strip() @@ -520,6 +564,21 @@ async def apply_project_init_payload(payload: dict[str, object], ) -> bool: try: replace_session_image(payload.get("initial_image")) + replace_last_frame_image(payload.get("last_frame_image")) + except ValueError as exc: + await ws_send_json({ + "type": "error", + "message": str(exc), + }) + return False + + try: + session_creation_config = parse_session_creation_config(payload) + validate_generation_mode_assets( + session_creation_config.generation_mode, + has_initial_image=session_init_image is not None, + has_last_frame_image=session_last_frame_image is not None, + ) except ValueError as exc: await ws_send_json({ "type": "error", @@ -1182,6 +1241,7 @@ async def generation_loop() -> None: finally: stop_event.set() cleanup_session_init_image(session_init_image) + cleanup_session_init_image(session_last_frame_image) for static_dir in FRONTEND_STATIC_DIR_CANDIDATES: diff --git a/apps/dreamverse/dreamverse/routes/creation.py b/apps/dreamverse/dreamverse/routes/creation.py new file mode 100644 index 0000000000..4c73d7046b --- /dev/null +++ b/apps/dreamverse/dreamverse/routes/creation.py @@ -0,0 +1,14 @@ +"""Creation studio capability routes.""" + +from __future__ import annotations + +from fastapi import APIRouter + +from dreamverse.creation_capabilities import lobby_capabilities_as_dict + +creation_router = APIRouter(tags=["creation"]) + + +@creation_router.get("/creation-capabilities") +async def creation_capabilities() -> dict[str, object]: + return lobby_capabilities_as_dict() diff --git a/apps/dreamverse/dreamverse/session/controller.py b/apps/dreamverse/dreamverse/session/controller.py index e3985938d8..461fa19b22 100644 --- a/apps/dreamverse/dreamverse/session/controller.py +++ b/apps/dreamverse/dreamverse/session/controller.py @@ -27,6 +27,7 @@ from fastapi import WebSocket, WebSocketDisconnect from dreamverse.gpu_pool import GPUSlot from dreamverse.session_init_image import cleanup_session_init_image, persist_session_init_image +from dreamverse.session_creation_config import parse_session_creation_config, validate_generation_mode_assets from dreamverse.worker_ipc import MediaChunk, MediaComplete, MediaInit from dreamverse.config import ( @@ -156,6 +157,9 @@ def get_first_blocked_prompt(prompts: list[str]): prompt_worker_task: asyncio.Task | None = None rewrite_seed_prompts_task: asyncio.Task | None = None session_init_image = None + session_last_frame_image = None + session_creation_config = None + session_generation_segment_cap = GENERATION_SEGMENT_CAP async def session_timeout(): """Close the session after timeout.""" @@ -237,6 +241,7 @@ async def cancel_task(task: asyncio.Task | None): try: session_init_image = persist_session_init_image(init_data.get("initial_image")) + session_last_frame_image = persist_session_init_image(init_data.get("last_frame_image")) except ValueError as exc: await ws_send_json({ "type": "error", @@ -245,6 +250,22 @@ async def cancel_task(task: asyncio.Task | None): await websocket.close(code=1003, reason="Invalid initial image") return + try: + session_creation_config = parse_session_creation_config(init_data) + session_generation_segment_cap = session_creation_config.generation_segment_cap + validate_generation_mode_assets( + session_creation_config.generation_mode, + has_initial_image=session_init_image is not None, + has_last_frame_image=session_last_frame_image is not None, + ) + except ValueError as exc: + await ws_send_json({ + "type": "error", + "message": str(exc), + }) + await websocket.close(code=1003, reason="Invalid creation config") + return + if preset_id: print(f"Client {client_id[:8]} selected preset: {preset_id} " f"label={preset_label or '(unset)'} " @@ -256,6 +277,16 @@ async def cancel_task(task: asyncio.Task | None): if session_init_image is not None: print(f"Client {client_id[:8]} uploaded initial image: " f"{session_init_image.display_name}") + if session_last_frame_image is not None: + print(f"Client {client_id[:8]} uploaded last frame image: " + f"{session_last_frame_image.display_name}") + if session_creation_config is not None: + print(f"Client {client_id[:8]} creation config: " + f"model={session_creation_config.model_id}, " + f"mode={session_creation_config.generation_mode}, " + f"size={session_creation_config.frame_width}x{session_creation_config.frame_height}, " + f"duration={session_creation_config.duration_sec}s, " + f"segment_cap={session_creation_config.generation_segment_cap}") # Acquire a GPU slot. gpu_id, slot = await self.gpu_pool.acquire(client_id, websocket) @@ -264,14 +295,20 @@ async def cancel_task(task: asyncio.Task | None): timeout_task = asyncio.create_task(session_timeout()) # Join the engine on this GPU. - await slot.join_user(client_id, model_id=ACTIVE_MODEL_ID) + await slot.join_user( + client_id, + model_id=session_creation_config.model_id if session_creation_config is not None else ACTIVE_MODEL_ID, + ) # Notify client they're connected to a GPU. - await ws_send_json({ + gpu_assigned_payload: dict[str, object] = { "type": "gpu_assigned", "gpu_id": gpu_id, "session_timeout": SESSION_TIMEOUT_SECONDS, - }) + } + if session_creation_config is not None: + gpu_assigned_payload["creation_config"] = session_creation_config.as_dict() + await ws_send_json(gpu_assigned_payload) await log_event( "gpu_assigned", { @@ -315,6 +352,14 @@ def replace_session_init_image(initial_image_payload: object) -> None: if previous_session_init_image is not None: cleanup_session_init_image(previous_session_init_image) + def replace_last_frame_image(last_frame_payload: object) -> None: + nonlocal session_last_frame_image + next_last_frame_image = persist_session_init_image(last_frame_payload) + previous_last_frame_image = session_last_frame_image + session_last_frame_image = next_last_frame_image + if previous_last_frame_image is not None: + cleanup_session_init_image(previous_last_frame_image) + async def schedule_simple_generate_request(payload: dict[str, object]) -> None: nonlocal preset_id nonlocal preset_label @@ -452,6 +497,8 @@ async def apply_project_init_payload(payload: dict[str, object]) -> bool: nonlocal project_active nonlocal project_stream_started nonlocal pending_project_end + nonlocal session_creation_config + nonlocal session_generation_segment_cap next_initial_rollout_prompt = str(payload.get("initial_rollout_prompt") or "").strip() next_enhancement_enabled = bool(payload.get("enhancement_enabled", True)) @@ -498,6 +545,22 @@ async def apply_project_init_payload(payload: dict[str, object]) -> bool: try: replace_session_init_image(payload.get("initial_image")) + replace_last_frame_image(payload.get("last_frame_image")) + except ValueError as exc: + await ws_send_json({ + "type": "error", + "message": str(exc), + }) + return False + + try: + session_creation_config = parse_session_creation_config(payload) + session_generation_segment_cap = session_creation_config.generation_segment_cap + validate_generation_mode_assets( + session_creation_config.generation_mode, + has_initial_image=session_init_image is not None, + has_last_frame_image=session_last_frame_image is not None, + ) except ValueError as exc: await ws_send_json({ "type": "error", @@ -941,7 +1004,7 @@ async def websocket_reader_loop(): "segment_cap": _resolve_generation_segment_cap( single_clip_mode=single_clip_mode, - cap=GENERATION_SEGMENT_CAP, + cap=session_generation_segment_cap, ), }) continue @@ -1281,27 +1344,22 @@ def pick_next_prompt_nowait() -> ReadyPrompt | None: )) else: project_stream_started = True - await ws_send_json({ - "type": - "ltx2_stream_start", - "total_segments": - len(curated_prompts), - "preset_id": - preset_id, - "stream_mode": - "av_fmp4", - "live_mode": - True, - "loop_generation_enabled": - loop_generation_enabled, - "loop_iteration": - loop_iteration, - "generation_segment_cap": - _resolve_generation_segment_cap( + stream_start_payload: dict[str, object] = { + "type": "ltx2_stream_start", + "total_segments": len(curated_prompts), + "preset_id": preset_id, + "stream_mode": "av_fmp4", + "live_mode": True, + "loop_generation_enabled": loop_generation_enabled, + "loop_iteration": loop_iteration, + "generation_segment_cap": _resolve_generation_segment_cap( single_clip_mode=single_clip_mode, - cap=GENERATION_SEGMENT_CAP, + cap=session_generation_segment_cap, ), - }) + } + if session_creation_config is not None: + stream_start_payload["creation_config"] = session_creation_config.as_dict() + await ws_send_json(stream_start_payload) await ws_send_json({ "type": "seed_prompts_updated", "prompts": seed_prompt_memory, @@ -1340,27 +1398,22 @@ def pick_next_prompt_nowait() -> ReadyPrompt | None: loop_iteration += 1 project_stream_started = True - await ws_send_json({ - "type": - "ltx2_stream_start", - "total_segments": - len(curated_prompts), - "preset_id": - preset_id, - "stream_mode": - "av_fmp4", - "live_mode": - True, - "loop_generation_enabled": - loop_generation_enabled, - "loop_iteration": - loop_iteration, - "generation_segment_cap": - _resolve_generation_segment_cap( + restart_stream_payload: dict[str, object] = { + "type": "ltx2_stream_start", + "total_segments": len(curated_prompts), + "preset_id": preset_id, + "stream_mode": "av_fmp4", + "live_mode": True, + "loop_generation_enabled": loop_generation_enabled, + "loop_iteration": loop_iteration, + "generation_segment_cap": _resolve_generation_segment_cap( single_clip_mode=single_clip_mode, - cap=GENERATION_SEGMENT_CAP, + cap=session_generation_segment_cap, ), - }) + } + if session_creation_config is not None: + restart_stream_payload["creation_config"] = session_creation_config.as_dict() + await ws_send_json(restart_stream_payload) if nonlocal_reason == "loop_restart": await ws_send_json({ "type": "loop_restarted", @@ -1389,13 +1442,14 @@ def pick_next_prompt_nowait() -> ReadyPrompt | None: pending_simple_prompt_submission = None if (not single_clip_mode and not generation_cap_blocked and not rollout_waiting_for_rewrite - and GENERATION_SEGMENT_CAP > 0 and generated_segment_count >= GENERATION_SEGMENT_CAP): + and session_generation_segment_cap > 0 + and generated_segment_count >= session_generation_segment_cap): loop_generation_enabled = False rollout_waiting_for_rewrite = True _main_print( "INFO", f"Segment cap reached for client {client_id[:8]} " - f"(cap_segments={GENERATION_SEGMENT_CAP}, " + f"(cap_segments={session_generation_segment_cap}, " f"generated_segments={generated_segment_count}); " "waiting for rollout rewrite", ) @@ -1620,6 +1674,9 @@ def pick_next_prompt_nowait() -> ReadyPrompt | None: pending_reset_conditioning = False step_image_path = (str(session_init_image.file_path) if segment_idx == 1 and session_init_image is not None else None) + step_frame_width = session_creation_config.frame_width if session_creation_config is not None else None + step_frame_height = session_creation_config.frame_height if session_creation_config is not None else None + step_num_frames = session_creation_config.num_frames if session_creation_config is not None else None step_task = asyncio.create_task( slot.user_step( client_id, @@ -1627,6 +1684,9 @@ def pick_next_prompt_nowait() -> ReadyPrompt | None: segment_idx=segment_idx, image_path=step_image_path, reset_conditioning=step_reset_conditioning, + frame_width=step_frame_width, + frame_height=step_frame_height, + num_frames=step_num_frames, )) segment_generation_active = True try: @@ -1808,3 +1868,4 @@ def pick_next_prompt_nowait() -> ReadyPrompt | None: await self.gpu_pool.release(client_id) finally: cleanup_session_init_image(session_init_image) + cleanup_session_init_image(session_last_frame_image) diff --git a/apps/dreamverse/dreamverse/session_creation_config.py b/apps/dreamverse/dreamverse/session_creation_config.py new file mode 100644 index 0000000000..0bf9cc0a1b --- /dev/null +++ b/apps/dreamverse/dreamverse/session_creation_config.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from dreamverse.config import FRAME_HEIGHT, FRAME_WIDTH, GENERATION_SEGMENT_CAP, MODEL_REGISTRY, NUM_FRAMES +from dreamverse.creation_capabilities import validate_lobby_creation_config + +LTX_LOBBY_MODEL_IDS = frozenset(MODEL_REGISTRY.keys()) +SUPPORTED_GENERATION_MODES = frozenset({"t2va", "fl2va", "ref2va"}) +SUPPORTED_ASPECT_RATIOS = frozenset({"21:9", "16:9", "4:3", "1:1", "3:4", "9:16"}) +SUPPORTED_RESOLUTIONS = frozenset({"480p", "720p", "1080p", "4k"}) +SEGMENT_DURATION_SEC = 5 + + +@dataclass(frozen=True) +class SessionCreationConfig: + model_id: str + generation_mode: str + aspect_ratio: str + resolution: str + duration_sec: int + frame_width: int + frame_height: int + num_frames: int + generation_segment_cap: int + + def as_dict(self) -> dict[str, object]: + return { + "model_id": self.model_id, + "generation_mode": self.generation_mode, + "aspect_ratio": self.aspect_ratio, + "resolution": self.resolution, + "duration_sec": self.duration_sec, + "frame_width": self.frame_width, + "frame_height": self.frame_height, + "num_frames": self.num_frames, + "generation_segment_cap": self.generation_segment_cap, + } + + +def _round_to_multiple(value: float, multiple: int = 32) -> int: + rounded = int(round(value / multiple)) * multiple + return max(multiple, rounded) + + +def _resolution_base(resolution: str) -> int: + return { + "480p": 480, + "720p": 720, + "1080p": 1080, + "4k": 2160, + }.get(resolution, 720) + + +def resolve_frame_size(aspect_ratio: str, resolution: str) -> tuple[int, int]: + if aspect_ratio == "16:9" and resolution == "1080p": + return FRAME_WIDTH, FRAME_HEIGHT + + base = _resolution_base(resolution) + width_ratio, height_ratio = { + "21:9": (21, 9), + "16:9": (16, 9), + "4:3": (4, 3), + "1:1": (1, 1), + "3:4": (3, 4), + "9:16": (9, 16), + }.get(aspect_ratio, (16, 9)) + + if width_ratio >= height_ratio: + height = _round_to_multiple(base) + width = _round_to_multiple(height * width_ratio / height_ratio) + else: + width = _round_to_multiple(base) + height = _round_to_multiple(width * height_ratio / width_ratio) + return width, height + + +def duration_sec_to_segment_cap(duration_sec: int, *, global_cap: int = GENERATION_SEGMENT_CAP) -> int: + requested = max(1, int(round(duration_sec / SEGMENT_DURATION_SEC + 0.0001))) + if global_cap <= 0: + return requested + return max(1, min(requested, global_cap)) + + +def parse_session_creation_config(payload: dict[str, object]) -> SessionCreationConfig: + raw_model_id = str(payload.get("model_id") or "").strip() + model_id = raw_model_id if raw_model_id in LTX_LOBBY_MODEL_IDS else "fast-ltx23" + + generation_mode = str(payload.get("generation_mode") or "t2va").strip() + if generation_mode not in SUPPORTED_GENERATION_MODES: + raise ValueError(f"Unsupported generation_mode: {generation_mode}") + + aspect_ratio = str(payload.get("aspect_ratio") or "16:9").strip() + if aspect_ratio not in SUPPORTED_ASPECT_RATIOS: + raise ValueError(f"Unsupported aspect_ratio: {aspect_ratio}") + + resolution = str(payload.get("resolution") or "720p").strip() + if resolution not in SUPPORTED_RESOLUTIONS: + raise ValueError(f"Unsupported resolution: {resolution}") + + try: + duration_sec = int(payload.get("duration_sec") or SEGMENT_DURATION_SEC) + except (TypeError, ValueError) as exc: + raise ValueError("duration_sec must be an integer.") from exc + if duration_sec not in {5, 10, 15}: + raise ValueError("duration_sec must be 5, 10, or 15.") + + validate_lobby_creation_config( + model_id=model_id, + generation_mode=generation_mode, + aspect_ratio=aspect_ratio, + resolution=resolution, + duration_sec=duration_sec, + ) + + if model_id not in MODEL_REGISTRY: + raise ValueError(f"Unsupported model_id: {model_id}") + + frame_width, frame_height = resolve_frame_size(aspect_ratio, resolution) + return SessionCreationConfig( + model_id=model_id, + generation_mode=generation_mode, + aspect_ratio=aspect_ratio, + resolution=resolution, + duration_sec=duration_sec, + frame_width=frame_width, + frame_height=frame_height, + num_frames=NUM_FRAMES, + generation_segment_cap=duration_sec_to_segment_cap(duration_sec), + ) + + +def validate_generation_mode_assets( + generation_mode: str, + *, + has_initial_image: bool, + has_last_frame_image: bool, +) -> None: + if generation_mode == "ref2va" and not has_initial_image: + raise ValueError("Ref2VA mode requires a reference image.") diff --git a/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py b/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py new file mode 100644 index 0000000000..b2253ea60e --- /dev/null +++ b/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py @@ -0,0 +1,74 @@ +import pytest + +from dreamverse.creation_capabilities import ( + capabilities_for_model, + lobby_capabilities_as_dict, + validate_lobby_creation_config, +) + + +def test_lobby_capabilities_include_all_registry_models(): + caps = lobby_capabilities_as_dict() + assert set(caps["model_ids"]) == {"fast-ltx2", "fast-ltx23", "fast-h3"} + assert "fl2va" not in caps["generation_modes"] + assert "4k" not in caps["resolutions"] + + +def test_fast_h3_capabilities_use_fixed_geometry(): + h3_caps = capabilities_for_model("fast-h3") + assert h3_caps.generation_modes == frozenset({"t2va", "ref2va"}) + assert h3_caps.aspect_ratios == frozenset({"16:9"}) + assert h3_caps.resolutions == frozenset({"720p"}) + + +def test_validate_lobby_creation_config_accepts_supported_t2va(): + validate_lobby_creation_config( + model_id="fast-ltx23", + generation_mode="t2va", + aspect_ratio="16:9", + resolution="1080p", + duration_sec=5, + ) + + +def test_validate_lobby_creation_config_accepts_fast_h3(): + validate_lobby_creation_config( + model_id="fast-h3", + generation_mode="ref2va", + aspect_ratio="16:9", + resolution="720p", + duration_sec=10, + ) + + +def test_validate_lobby_creation_config_rejects_fl2va(): + with pytest.raises(ValueError, match="FL2VA"): + validate_lobby_creation_config( + model_id="fast-ltx23", + generation_mode="fl2va", + aspect_ratio="16:9", + resolution="720p", + duration_sec=5, + ) + + +def test_validate_lobby_creation_config_rejects_4k(): + with pytest.raises(ValueError, match="Unsupported resolution"): + validate_lobby_creation_config( + model_id="fast-ltx2", + generation_mode="t2va", + aspect_ratio="16:9", + resolution="4k", + duration_sec=10, + ) + + +def test_validate_lobby_creation_config_rejects_invalid_h3_aspect(): + with pytest.raises(ValueError, match="Unsupported aspect_ratio"): + validate_lobby_creation_config( + model_id="fast-h3", + generation_mode="t2va", + aspect_ratio="9:16", + resolution="720p", + duration_sec=5, + ) diff --git a/apps/dreamverse/dreamverse/tests/test_session_creation_config.py b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py new file mode 100644 index 0000000000..4295374c3e --- /dev/null +++ b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py @@ -0,0 +1,89 @@ +import pytest + +from dreamverse.session_creation_config import ( + duration_sec_to_segment_cap, + parse_session_creation_config, + resolve_frame_size, + validate_generation_mode_assets, +) + + +def test_parse_session_creation_config_defaults(): + config = parse_session_creation_config({}) + assert config.model_id == "fast-ltx23" + assert config.generation_mode == "t2va" + assert config.aspect_ratio == "16:9" + assert config.resolution == "720p" + assert config.duration_sec == 5 + assert config.generation_segment_cap == 1 + + +def test_parse_session_creation_config_maps_duration_to_segment_cap(): + config = parse_session_creation_config( + { + "model_id": "fast-ltx2", + "generation_mode": "ref2va", + "aspect_ratio": "9:16", + "resolution": "480p", + "duration_sec": 15, + }, + ) + assert config.model_id == "fast-ltx2" + assert config.generation_mode == "ref2va" + assert config.generation_segment_cap == 3 + assert config.frame_width >= 480 + assert config.frame_height >= 480 + + +def test_resolve_frame_size_uses_model_default_for_1080p_landscape(): + width, height = resolve_frame_size("16:9", "1080p") + assert (width, height) == (1920, 1088) + + +def test_duration_sec_to_segment_cap_respects_global_cap(): + assert duration_sec_to_segment_cap(15, global_cap=2) == 2 + + +def test_parse_session_creation_config_accepts_fast_h3(): + config = parse_session_creation_config( + { + "model_id": "fast-h3", + "generation_mode": "t2va", + "aspect_ratio": "16:9", + "resolution": "720p", + "duration_sec": 10, + }, + ) + assert config.model_id == "fast-h3" + assert config.generation_mode == "t2va" + assert config.generation_segment_cap == 2 + + +def test_parse_session_creation_config_rejects_fl2va(): + with pytest.raises(ValueError, match="FL2VA"): + parse_session_creation_config( + { + "generation_mode": "fl2va", + "aspect_ratio": "16:9", + "resolution": "720p", + "duration_sec": 5, + }, + ) + + +def test_parse_session_creation_config_rejects_4k(): + with pytest.raises(ValueError, match="Unsupported resolution"): + parse_session_creation_config( + { + "generation_mode": "t2va", + "aspect_ratio": "16:9", + "resolution": "4k", + "duration_sec": 5, + }, + ) + + +def test_validate_generation_mode_assets(): + validate_generation_mode_assets("t2va", has_initial_image=False, has_last_frame_image=False) + with pytest.raises(ValueError, match="Ref2VA"): + validate_generation_mode_assets("ref2va", has_initial_image=False, has_last_frame_image=False) diff --git a/apps/dreamverse/dreamverse/worker_ipc.py b/apps/dreamverse/dreamverse/worker_ipc.py index cd01867c2c..6494602ba4 100644 --- a/apps/dreamverse/dreamverse/worker_ipc.py +++ b/apps/dreamverse/dreamverse/worker_ipc.py @@ -147,6 +147,9 @@ class UserStepPayload: segment_idx: int image_path: str | None reset_conditioning: bool + frame_width: int | None = None + frame_height: int | None = None + num_frames: int | None = None @dataclass(frozen=True) diff --git a/apps/dreamverse/web/e2e/frontend-shell.spec.ts b/apps/dreamverse/web/e2e/frontend-shell.spec.ts index 3f5886c26b..74e6d502c0 100644 --- a/apps/dreamverse/web/e2e/frontend-shell.spec.ts +++ b/apps/dreamverse/web/e2e/frontend-shell.spec.ts @@ -28,6 +28,6 @@ test.describe('frontend shell', () => { await expect(page.getByText('Direct scenes in seconds')).toBeVisible({ timeout: 30_000 }); await expect(page.getByRole('button', { name: /FastLTX/i }).first()).toBeVisible({ timeout: 30_000 }); - await expect(page.getByText('Describe your video or mention elements')).toBeVisible({ timeout: 30_000 }); + await expect(continuation).toHaveAttribute('placeholder', /Describe your video or mention elements/i); }); }); diff --git a/apps/dreamverse/web/next.config.ts b/apps/dreamverse/web/next.config.ts index 6a7e42e3fa..67b70ff222 100644 --- a/apps/dreamverse/web/next.config.ts +++ b/apps/dreamverse/web/next.config.ts @@ -42,6 +42,10 @@ const nextConfig: NextConfig = { source: '/prompt-system-config', destination: `${backendUrl}/prompt-system-config`, }, + { + source: '/creation-capabilities', + destination: `${backendUrl}/creation-capabilities`, + }, { source: '/curated-presets', destination: `${backendUrl}/curated-presets`, diff --git a/apps/dreamverse/web/src/app/page.tsx b/apps/dreamverse/web/src/app/page.tsx index 27f666c40a..7172866be9 100644 --- a/apps/dreamverse/web/src/app/page.tsx +++ b/apps/dreamverse/web/src/app/page.tsx @@ -32,6 +32,18 @@ import { buildRewritePromptWindowSnapshotFromPrompts, normalizePromptWindowSnapshot, } from "@/lib/prompts/promptWindowSnapshot"; +import { + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, + clampLobbySelectionToCapabilities, + parseLobbyCapabilitiesBundle, + resolveModelCapabilities, + validateLobbyCreationSelection, + type LobbyCapabilitiesBundle, +} from "@/lib/creationCapabilities"; +import { + buildCreationInitPayload, + parseEchoedCreationConfig, +} from "@/lib/creationPayload"; import rawPresets from "@/lib/storyPresetsData"; import { cn } from "@/lib/utils"; import { createWebSocketConnection, detachAndCloseWebSocket } from "@/lib/ws/client"; @@ -241,6 +253,9 @@ export default function Page() { const [ttffValueMs, setTtffValueMs] = useState(null); const ttffIntervalRef = useRef | null>(null); const pendingInitialPromptRef = useRef(""); + const referenceFileRef = useRef(null); + const firstFrameFileRef = useRef(null); + const lastFrameFileRef = useRef(null); const lastArchivedReplayKeyRef = useRef(""); const [sidebarOpen, setSidebarOpen] = useState(false); const [creationModelId, setCreationModelId] = useState("fast-ltx23"); @@ -248,6 +263,13 @@ export default function Page() { const [creationAspectRatio, setCreationAspectRatio] = useState("16:9"); const [creationResolution, setCreationResolution] = useState("720p"); const [creationDurationSec, setCreationDurationSec] = useState(5); + const [lobbyCapabilitiesBundle, setLobbyCapabilitiesBundle] = useState( + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, + ); + const activeModelCapabilities = useMemo( + () => resolveModelCapabilities(lobbyCapabilitiesBundle, creationModelId), + [lobbyCapabilitiesBundle, creationModelId], + ); const [sessionCreationConfig, setSessionCreationConfig] = useState({ modelId: "fast-ltx23", modeId: "t2v", @@ -300,14 +322,17 @@ export default function Page() { } function handleReferenceSelect(file: File | null) { + referenceFileRef.current = file; setPreviewUrl(setReferencePreviewUrl, file); } function handleFirstFrameSelect(file: File | null) { + firstFrameFileRef.current = file; setPreviewUrl(setFirstFramePreviewUrl, file); } function handleLastFrameSelect(file: File | null) { + lastFrameFileRef.current = file; setPreviewUrl(setLastFramePreviewUrl, file); } @@ -505,6 +530,65 @@ export default function Page() { setRuntimeReady(true); }, []); + function applyLobbyCapabilitiesBundle(bundle: LobbyCapabilitiesBundle) { + setLobbyCapabilitiesBundle(bundle); + const clamped = clampLobbySelectionToCapabilities({ + capabilities: resolveModelCapabilities(bundle, creationModelId), + modelId: creationModelId, + modeId: creationModeId, + aspectRatio: creationAspectRatio, + resolution: creationResolution, + durationSec: creationDurationSec, + }); + setCreationModelId(clamped.modelId); + setCreationModeId(clamped.modeId); + setCreationAspectRatio(clamped.aspectRatio); + setCreationResolution(clamped.resolution); + setCreationDurationSec(clamped.durationSec); + } + + function handleCreationModelChange(modelId: CreationModelId) { + const clamped = clampLobbySelectionToCapabilities({ + capabilities: resolveModelCapabilities(lobbyCapabilitiesBundle, modelId), + modelId, + modeId: creationModeId, + aspectRatio: creationAspectRatio, + resolution: creationResolution, + durationSec: creationDurationSec, + }); + setCreationModelId(clamped.modelId); + setCreationModeId(clamped.modeId); + setCreationAspectRatio(clamped.aspectRatio); + setCreationResolution(clamped.resolution); + setCreationDurationSec(clamped.durationSec); + } + + useEffect(() => { + if (!runtimeReady) return; + let cancelled = false; + void fetch("/creation-capabilities", { + headers: { Accept: "application/json" }, + cache: "no-store", + }) + .then(async (response) => { + if (!response.ok) return DEFAULT_LOBBY_CAPABILITIES_BUNDLE; + return parseLobbyCapabilitiesBundle(await response.json()); + }) + .then((bundle) => { + if (!cancelled) { + applyLobbyCapabilitiesBundle(bundle); + } + }) + .catch(() => { + if (!cancelled) { + applyLobbyCapabilitiesBundle(DEFAULT_LOBBY_CAPABILITIES_BUNDLE); + } + }); + return () => { + cancelled = true; + }; + }, [runtimeReady]); + useEffect(() => { if (!runtimeReady || initializedRef.current) return; initializedRef.current = true; @@ -1740,9 +1824,19 @@ export default function Page() { resetPlaybackState(); } - function buildProjectInitPayload(type: "session_init_v2" | "project_init_v1") { + async function buildProjectInitPayload(type: "session_init_v2" | "project_init_v1") { const segmentPrompts = getSessionInitPrompts(); setSeedPrompts(segmentPrompts); + const creationPayload = await buildCreationInitPayload({ + modelId: creationModelId, + modeId: creationModeId, + aspectRatio: creationAspectRatio, + resolution: creationResolution, + durationSec: creationDurationSec, + referenceFile: referenceFileRef.current, + firstFrameFile: firstFrameFileRef.current, + lastFrameFile: lastFrameFileRef.current, + }); return { type, generation_mode: toGenerationMode(creationModeId), @@ -1750,24 +1844,24 @@ export default function Page() { preset_label: getInitialPresetLabel(), curated_prompts: segmentPrompts, initial_rollout_prompt: normalizeInitialPrompt(pendingInitialPromptRef.current), - initial_image: null, single_clip_mode: false, enhancement_enabled: sessionStore.get().enhancementEnabled, auto_extension_enabled: sessionStore.get().autoExtensionEnabled, loop_generation_enabled: sessionStore.get().loopGenerationEnabled, + ...creationPayload, }; } - function sendSessionInitMessage() { + async function sendSessionInitMessage() { const ws = wsRef.current; if (!ws) return; - ws.send(JSON.stringify(buildProjectInitPayload("session_init_v2"))); + ws.send(JSON.stringify(await buildProjectInitPayload("session_init_v2"))); } - function sendProjectInitMessage() { + async function sendProjectInitMessage() { const ws = wsRef.current; if (!ws || ws.readyState !== WebSocket.OPEN) return; - ws.send(JSON.stringify(buildProjectInitPayload("project_init_v1"))); + ws.send(JSON.stringify(await buildProjectInitPayload("project_init_v1"))); } function sendEndProjectKeepSession() { @@ -1804,6 +1898,9 @@ export default function Page() { return; } const normalizedEvent = normalizeSocketMessage(decoded.data); + if (decoded.data?.type === "gpu_assigned" || decoded.data?.type === "ltx2_stream_start") { + applyEchoedCreationConfig(decoded.data); + } await applyNormalizedSocketEvent(normalizedEvent, { sessionStore, promptWindowStore, @@ -1846,7 +1943,12 @@ export default function Page() { onOpen: () => { opened = true; sessionStore.patch({ connected: true, connecting: false }); - sendSessionInitMessage(); + void sendSessionInitMessage().catch((error) => { + console.error("Failed to send session init payload:", error); + recoverFailedSessionStart( + error instanceof Error ? error.message : "Failed to prepare session settings.", + ); + }); }, onMessage: (event: MessageEvent) => { wsMessageQueueRef.current = wsMessageQueueRef.current @@ -1929,6 +2031,14 @@ export default function Page() { }); } + function applyEchoedCreationConfig(data: unknown) { + const echoed = parseEchoedCreationConfig(data); + if (!echoed) { + return; + } + setSessionCreationConfig(echoed); + } + function beginProjectLocally({ force = false } = {}) { if (!force && !canStartSession) return; if (sessionStore.get().sessionStarted || sessionStore.get().projectResetPending) return false; @@ -1992,13 +2102,34 @@ export default function Page() { } async function joinSession({ force = false } = {}) { + const validationError = validateLobbyCreationSelection({ + capabilities: activeModelCapabilities, + modelId: creationModelId, + modeId: creationModeId, + aspectRatio: creationAspectRatio, + resolution: creationResolution, + durationSec: creationDurationSec, + referenceFile: referenceFileRef.current, + firstFrameFile: firstFrameFileRef.current, + lastFrameFile: lastFrameFileRef.current, + }); + if (validationError) { + showPreSessionNotice(validationError); + return; + } + if ( wsRef.current && wsRef.current.readyState === WebSocket.OPEN && sessionStore.get().connected ) { if (!beginProjectLocally({ force })) return; - sendProjectInitMessage(); + try { + await sendProjectInitMessage(); + } catch (error) { + console.error("Failed to send project init payload:", error); + showPreSessionNotice(error instanceof Error ? error.message : "Failed to prepare session settings."); + } return; } showPreSessionNotice(""); @@ -2020,7 +2151,12 @@ export default function Page() { && wsRef.current.readyState === WebSocket.OPEN && sessionStore.get().connected ) { - sendProjectInitMessage(); + try { + await sendProjectInitMessage(); + } catch (error) { + console.error("Failed to send project init payload:", error); + showPreSessionNotice(error instanceof Error ? error.message : "Failed to prepare session settings."); + } return; } connectWebSocket(); @@ -2767,10 +2903,11 @@ export default function Page() { lastFramePreviewUrl={lastFramePreviewUrl} mentionOptions={mentionOptions} storyPresets={lobbyStoryPresets} + capabilities={activeModelCapabilities} onValueChange={(value) => sessionStore.patch({ livePromptDraft: value })} onSubmit={() => void joinSession()} onKeyDown={handleLivePromptKeydown} - onModelChange={setCreationModelId} + onModelChange={handleCreationModelChange} onModeChange={setCreationModeId} onAspectRatioChange={setCreationAspectRatio} onResolutionChange={setCreationResolution} @@ -2797,6 +2934,7 @@ export default function Page() { sessionNotice={sessionNotice as string} projectResetPending={projectResetPending as boolean} sessionCreationConfig={sessionCreationConfig} + configPillsReadOnly onSessionModelChange={(modelId) => setSessionCreationConfig((current) => ({ ...current, modelId }))} onSessionModeChange={(modeId) => setSessionCreationConfig((current) => ({ ...current, modeId }))} onSessionAspectRatioChange={(aspectRatio) => setSessionCreationConfig((current) => ({ ...current, aspectRatio }))} diff --git a/apps/dreamverse/web/src/components/creation/CreationComposer.tsx b/apps/dreamverse/web/src/components/creation/CreationComposer.tsx index df3183603d..e2140cbf83 100644 --- a/apps/dreamverse/web/src/components/creation/CreationComposer.tsx +++ b/apps/dreamverse/web/src/components/creation/CreationComposer.tsx @@ -23,6 +23,8 @@ import { CREATION_MODELS, CREATION_MODES, RESOLUTIONS, + UNSUPPORTED_CREATION_MODES, + UNSUPPORTED_RESOLUTIONS, modeRequiresReference, modeUsesDualFrames, type AspectRatioId, @@ -33,6 +35,14 @@ import { formatDurationLabel, formatResolutionLabel, } from "@/lib/creationConfig"; +import { + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, + isSupportedCreationMode, + isSupportedResolution, + resolveModelCapabilities, + unsupportedModeNotice, + type LobbyCreationCapabilities, +} from "@/lib/creationCapabilities"; import { cn } from "@/lib/utils"; const PROMPT_MAX_LENGTH = 500; @@ -64,6 +74,7 @@ interface CreationComposerProps { onLastFrameSelect?: (file: File | null) => void; onSpeechTranscript?: (text: string) => void; onSpeechInterimChange?: (text: string) => void; + capabilities?: LobbyCreationCapabilities; } export default function CreationComposer({ @@ -93,6 +104,7 @@ export default function CreationComposer({ onLastFrameSelect, onSpeechTranscript, onSpeechInterimChange, + capabilities = resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, modelId), }: CreationComposerProps) { const inputRef = useRef(null); const [sttBusy, setSttBusy] = useState(false); @@ -100,8 +112,38 @@ export default function CreationComposer({ const [mentionOpen, setMentionOpen] = useState(false); const [mentionStart, setMentionStart] = useState(null); - const selectedModel = CREATION_MODELS.find((model) => model.id === modelId) ?? CREATION_MODELS[0]; - const selectedMode = CREATION_MODES.find((mode) => mode.id === modeId) ?? CREATION_MODES[0]; + const availableModels = useMemo( + () => CREATION_MODELS.filter((model) => capabilities.model_ids.includes(model.id)), + [capabilities.model_ids], + ); + const availableModes = useMemo( + () => CREATION_MODES.filter((mode) => isSupportedCreationMode(mode.id, capabilities)), + [capabilities], + ); + const unavailableModes = useMemo( + () => + UNSUPPORTED_CREATION_MODES.filter( + (mode) => unsupportedModeNotice(mode.id, capabilities) !== null, + ), + [capabilities], + ); + const availableAspectRatios = useMemo( + () => ASPECT_RATIOS.filter((ratio) => capabilities.aspect_ratios.includes(ratio)), + [capabilities.aspect_ratios], + ); + const availableResolutions = useMemo( + () => RESOLUTIONS.filter((item) => isSupportedResolution(item, capabilities)), + [capabilities], + ); + const unavailableResolutions = useMemo( + () => UNSUPPORTED_RESOLUTIONS.filter((item) => !isSupportedResolution(item, capabilities)), + [capabilities], + ); + const durationMin = capabilities.duration_sec[0] ?? 5; + const durationMax = capabilities.duration_sec[capabilities.duration_sec.length - 1] ?? 15; + + const selectedModel = availableModels.find((model) => model.id === modelId) ?? availableModels[0]; + const selectedMode = availableModes.find((mode) => mode.id === modeId) ?? availableModes[0]; const usesDualFrames = modeUsesDualFrames(modeId); const requiresReference = modeRequiresReference(modeId); const referenceMissing = requiresReference && !referencePreviewUrl; @@ -278,7 +320,7 @@ export default function CreationComposer({ Model - {CREATION_MODELS.map((model) => ( + {availableModels.map((model) => ( onModelChange(model.id)} className="flex-col items-start gap-1 py-2.5"> {model.label} @@ -301,12 +343,21 @@ export default function CreationComposer({ Mode - {CREATION_MODES.map((mode) => ( + {availableModes.map((mode) => ( onModeChange(mode.id)} className="flex-col items-start gap-1 py-2.5"> {mode.label} {mode.description} ))} + {unavailableModes.length > 0 && } + {unavailableModes.map((mode) => ( + + {mode.label} + + {unsupportedModeNotice(mode.id, capabilities) ?? mode.description} + + + ))} @@ -320,7 +371,7 @@ export default function CreationComposer({

Aspect ratio

- {ASPECT_RATIOS.map((ratio) => ( + {availableAspectRatios.map((ratio) => ( + ))}
@@ -363,11 +425,11 @@ export default function CreationComposer({

Total duration

- onDurationChange(values[0] ?? 5)} /> + onDurationChange(values[0] ?? durationMin)} />
- 5s + {formatDurationLabel(durationMin)} {formatDurationLabel(durationSec)} - 15s + {formatDurationLabel(durationMax)}
diff --git a/apps/dreamverse/web/src/components/creation/CreationStudio.tsx b/apps/dreamverse/web/src/components/creation/CreationStudio.tsx index cec2859321..9083d3bcd6 100644 --- a/apps/dreamverse/web/src/components/creation/CreationStudio.tsx +++ b/apps/dreamverse/web/src/components/creation/CreationStudio.tsx @@ -12,6 +12,7 @@ import { type MentionOption, type ResolutionId, } from "@/lib/creationConfig"; +import type { LobbyCreationCapabilities } from "@/lib/creationCapabilities"; interface CreationStudioProps { value: string; @@ -44,6 +45,7 @@ interface CreationStudioProps { onSpeechTranscript?: (text: string) => void; onSpeechInterimChange?: (text: string) => void; onOpenProjects?: () => void; + capabilities?: LobbyCreationCapabilities; } export default function CreationStudio({ @@ -52,6 +54,7 @@ export default function CreationStudio({ storyPresets = [], onPresetGenerate, isGenerating = false, + capabilities, ...composerProps }: CreationStudioProps) { return ( @@ -59,7 +62,7 @@ export default function CreationStudio({
- + {storyPresets.length > 0 && onPresetGenerate && ( )} diff --git a/apps/dreamverse/web/src/lib/creationCapabilities.test.ts b/apps/dreamverse/web/src/lib/creationCapabilities.test.ts new file mode 100644 index 0000000000..3ba1576d8a --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationCapabilities.test.ts @@ -0,0 +1,80 @@ +import { describe, expect, it } from "vitest"; + +import { + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, + clampLobbySelectionToCapabilities, + parseLobbyCapabilitiesBundle, + resolveModelCapabilities, + validateLobbyCreationSelection, +} from "@/lib/creationCapabilities"; + +describe("creationCapabilities", () => { + it("parses backend capability payloads with per-model caps", () => { + const bundle = parseLobbyCapabilitiesBundle({ + model_ids: ["fast-ltx2", "fast-h3"], + models: { + "fast-ltx2": { + generation_modes: ["t2va"], + resolutions: ["480p", "720p"], + duration_sec: [5, 10], + }, + "fast-h3": { + generation_modes: ["t2va", "ref2va"], + aspect_ratios: ["16:9"], + resolutions: ["720p"], + }, + }, + }); + expect(bundle.model_ids).toEqual(["fast-ltx2", "fast-h3"]); + expect(bundle.models["fast-h3"]?.aspect_ratios).toEqual(["16:9"]); + }); + + it("includes fast-h3 in default lobby models", () => { + expect(DEFAULT_LOBBY_CAPABILITIES_BUNDLE.model_ids).toContain("fast-h3"); + }); + + it("clamps unsupported lobby selections to model-specific defaults", () => { + expect( + clampLobbySelectionToCapabilities({ + capabilities: resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, "fast-h3"), + modelId: "fast-h3", + modeId: "fl2av", + aspectRatio: "9:16", + resolution: "4k", + durationSec: 99, + }), + ).toEqual({ + modelId: "fast-h3", + modeId: "t2v", + aspectRatio: "16:9", + resolution: "720p", + durationSec: 5, + }); + }); + + it("rejects unsupported generation modes with a clear message", () => { + expect( + validateLobbyCreationSelection({ + capabilities: resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, "fast-ltx23"), + modelId: "fast-ltx23", + modeId: "fl2av", + aspectRatio: "16:9", + resolution: "720p", + durationSec: 5, + }), + ).toMatch(/FL2VA/i); + }); + + it("rejects unsupported resolutions for ltx models", () => { + expect( + validateLobbyCreationSelection({ + capabilities: resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, "fast-ltx23"), + modelId: "fast-ltx23", + modeId: "t2v", + aspectRatio: "16:9", + resolution: "4k", + durationSec: 5, + }), + ).toMatch(/resolution/i); + }); +}); diff --git a/apps/dreamverse/web/src/lib/creationCapabilities.ts b/apps/dreamverse/web/src/lib/creationCapabilities.ts new file mode 100644 index 0000000000..696c0cd317 --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationCapabilities.ts @@ -0,0 +1,263 @@ +import type { + AspectRatioId, + CreationModeId, + CreationModelId, + ResolutionId, +} from "@/lib/creationConfig"; +import { fromGenerationMode, toGenerationMode, type GenerationMode } from "@/lib/generationMode"; + +const ALL_MODEL_IDS: CreationModelId[] = ["fast-ltx23", "fast-ltx2", "fast-h3"]; +const ALL_GENERATION_MODES: GenerationMode[] = ["t2va", "fl2va", "ref2va"]; +const ALL_ASPECT_RATIOS: AspectRatioId[] = ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]; +const ALL_RESOLUTIONS: ResolutionId[] = ["480p", "720p", "1080p", "4k"]; + +export interface ModelCreationCapabilities { + generation_modes: GenerationMode[]; + aspect_ratios: AspectRatioId[]; + resolutions: ResolutionId[]; + duration_sec: number[]; + unsupported_generation_modes: Record; + reference_assets: { + mime_types: string[]; + max_bytes: number; + }; +} + +export interface LobbyCreationCapabilities extends ModelCreationCapabilities { + model_ids: CreationModelId[]; +} + +export interface LobbyCapabilitiesBundle { + model_ids: CreationModelId[]; + models: Partial>; + generation_modes: GenerationMode[]; + aspect_ratios: AspectRatioId[]; + resolutions: ResolutionId[]; + duration_sec: number[]; + unsupported_generation_modes: Record; + reference_assets: { + mime_types: string[]; + max_bytes: number; + }; +} + +const DEFAULT_LTX_MODEL_CAPABILITIES: ModelCreationCapabilities = { + generation_modes: ["t2va", "ref2va"], + aspect_ratios: ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"], + resolutions: ["480p", "720p", "1080p"], + duration_sec: [5, 10, 15], + unsupported_generation_modes: { + fl2va: "First/last frame mode (FL2VA) is not supported yet.", + }, + reference_assets: { + mime_types: ["image/png", "image/jpeg", "image/webp"], + max_bytes: 15 * 1024 * 1024, + }, +}; + +const DEFAULT_H3_MODEL_CAPABILITIES: ModelCreationCapabilities = { + generation_modes: ["t2va", "ref2va"], + aspect_ratios: ["16:9"], + resolutions: ["720p"], + duration_sec: [5, 10, 15], + unsupported_generation_modes: { + fl2va: "First/last frame mode (FL2VA) is not supported yet.", + }, + reference_assets: DEFAULT_LTX_MODEL_CAPABILITIES.reference_assets, +}; + +export const DEFAULT_LOBBY_CAPABILITIES_BUNDLE: LobbyCapabilitiesBundle = { + model_ids: ALL_MODEL_IDS, + models: { + "fast-ltx2": DEFAULT_LTX_MODEL_CAPABILITIES, + "fast-ltx23": DEFAULT_LTX_MODEL_CAPABILITIES, + "fast-h3": DEFAULT_H3_MODEL_CAPABILITIES, + }, + generation_modes: ["t2va", "ref2va"], + aspect_ratios: ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"], + resolutions: ["480p", "720p", "1080p"], + duration_sec: [5, 10, 15], + unsupported_generation_modes: DEFAULT_LTX_MODEL_CAPABILITIES.unsupported_generation_modes, + reference_assets: DEFAULT_LTX_MODEL_CAPABILITIES.reference_assets, +}; + +function pickStrings(value: unknown, allowed: readonly T[], fallback: readonly T[]): T[] { + if (!Array.isArray(value)) return [...fallback]; + return value.filter((item): item is T => typeof item === "string" && allowed.includes(item as T)); +} + +function parseReferenceAssets( + value: unknown, + fallback: ModelCreationCapabilities["reference_assets"], +): ModelCreationCapabilities["reference_assets"] { + if (!value || typeof value !== "object") return fallback; + const data = value as Record; + return { + mime_types: Array.isArray(data.mime_types) + ? (data.mime_types as string[]) + : fallback.mime_types, + max_bytes: typeof data.max_bytes === "number" ? data.max_bytes : fallback.max_bytes, + }; +} + +function parseModelCreationCapabilities( + value: unknown, + fallback: ModelCreationCapabilities, +): ModelCreationCapabilities { + if (!value || typeof value !== "object") return fallback; + const data = value as Record; + return { + generation_modes: pickStrings(data.generation_modes, ALL_GENERATION_MODES, fallback.generation_modes), + aspect_ratios: pickStrings(data.aspect_ratios, ALL_ASPECT_RATIOS, fallback.aspect_ratios), + resolutions: pickStrings(data.resolutions, ALL_RESOLUTIONS, fallback.resolutions), + duration_sec: Array.isArray(data.duration_sec) + ? data.duration_sec.filter((item): item is number => typeof item === "number") + : fallback.duration_sec, + unsupported_generation_modes: + typeof data.unsupported_generation_modes === "object" && data.unsupported_generation_modes + ? (data.unsupported_generation_modes as Record) + : fallback.unsupported_generation_modes, + reference_assets: parseReferenceAssets(data.reference_assets, fallback.reference_assets), + }; +} + +export function parseLobbyCapabilitiesBundle(payload: unknown): LobbyCapabilitiesBundle { + if (!payload || typeof payload !== "object") { + return DEFAULT_LOBBY_CAPABILITIES_BUNDLE; + } + const data = payload as Record; + const modelIds = pickStrings(data.model_ids, ALL_MODEL_IDS, DEFAULT_LOBBY_CAPABILITIES_BUNDLE.model_ids); + const rawModels = typeof data.models === "object" && data.models ? (data.models as Record) : {}; + const models: Partial> = {}; + for (const modelId of modelIds) { + const fallback = + DEFAULT_LOBBY_CAPABILITIES_BUNDLE.models[modelId] ?? + (modelId === "fast-h3" ? DEFAULT_H3_MODEL_CAPABILITIES : DEFAULT_LTX_MODEL_CAPABILITIES); + models[modelId] = parseModelCreationCapabilities(rawModels[modelId], fallback); + } + const unionFallback = parseModelCreationCapabilities(payload, DEFAULT_LTX_MODEL_CAPABILITIES); + return { + model_ids: modelIds, + models, + generation_modes: unionFallback.generation_modes, + aspect_ratios: unionFallback.aspect_ratios, + resolutions: unionFallback.resolutions, + duration_sec: unionFallback.duration_sec, + unsupported_generation_modes: unionFallback.unsupported_generation_modes, + reference_assets: unionFallback.reference_assets, + }; +} + +export function resolveModelCapabilities( + bundle: LobbyCapabilitiesBundle, + modelId: CreationModelId, +): LobbyCreationCapabilities { + const modelCaps = + bundle.models[modelId] ?? + (modelId === "fast-h3" ? DEFAULT_H3_MODEL_CAPABILITIES : DEFAULT_LTX_MODEL_CAPABILITIES); + return { + model_ids: bundle.model_ids, + ...modelCaps, + }; +} + +export function supportedCreationModes(capabilities: LobbyCreationCapabilities) { + return capabilities.generation_modes.map((wireMode) => ({ + wireMode, + modeId: fromGenerationMode(wireMode), + })); +} + +export function isSupportedCreationMode(modeId: CreationModeId, capabilities: LobbyCreationCapabilities): boolean { + return capabilities.generation_modes.includes(toGenerationMode(modeId)); +} + +export function isSupportedResolution(resolution: ResolutionId, capabilities: LobbyCreationCapabilities): boolean { + return capabilities.resolutions.includes(resolution); +} + +export function isSupportedReferenceImage(file: File, capabilities: LobbyCreationCapabilities): boolean { + return capabilities.reference_assets.mime_types.includes(file.type); +} + +export function unsupportedModeNotice(modeId: CreationModeId, capabilities: LobbyCreationCapabilities): string | null { + const wireMode = toGenerationMode(modeId); + return capabilities.unsupported_generation_modes[wireMode] ?? null; +} + +export function clampLobbySelectionToCapabilities(input: { + capabilities: LobbyCreationCapabilities; + modelId: CreationModelId; + modeId: CreationModeId; + aspectRatio: AspectRatioId; + resolution: ResolutionId; + durationSec: number; +}): { + modelId: CreationModelId; + modeId: CreationModeId; + aspectRatio: AspectRatioId; + resolution: ResolutionId; + durationSec: number; +} { + const { capabilities } = input; + const modelId = capabilities.model_ids.includes(input.modelId) + ? input.modelId + : (capabilities.model_ids[0] ?? "fast-ltx23"); + const supportedModes = supportedCreationModes(capabilities); + const modeId = isSupportedCreationMode(input.modeId, capabilities) + ? input.modeId + : (supportedModes[0]?.modeId ?? "t2v"); + const aspectRatio = capabilities.aspect_ratios.includes(input.aspectRatio) + ? input.aspectRatio + : (capabilities.aspect_ratios[0] ?? "16:9"); + const resolution = isSupportedResolution(input.resolution, capabilities) + ? input.resolution + : (capabilities.resolutions[0] ?? "720p"); + const durationSec = capabilities.duration_sec.includes(input.durationSec) + ? input.durationSec + : (capabilities.duration_sec[0] ?? 5); + return { modelId, modeId, aspectRatio, resolution, durationSec }; +} + +export function validateLobbyCreationSelection(input: { + capabilities: LobbyCreationCapabilities; + modelId: CreationModelId; + modeId: CreationModeId; + aspectRatio: AspectRatioId; + resolution: ResolutionId; + durationSec: number; + referenceFile?: File | null; + firstFrameFile?: File | null; + lastFrameFile?: File | null; +}): string | null { + const unsupportedMode = unsupportedModeNotice(input.modeId, input.capabilities); + if (unsupportedMode) return unsupportedMode; + if (!input.capabilities.model_ids.includes(input.modelId)) { + return "Selected model is not supported yet."; + } + if (!isSupportedCreationMode(input.modeId, input.capabilities)) { + return "Selected mode is not supported yet."; + } + if (!input.capabilities.aspect_ratios.includes(input.aspectRatio)) { + return "Selected aspect ratio is not supported for this model yet."; + } + if (!isSupportedResolution(input.resolution, input.capabilities)) { + return "Selected resolution is not supported for this model yet."; + } + if (!input.capabilities.duration_sec.includes(input.durationSec)) { + return "Selected duration is not supported yet."; + } + if (input.modeId === "ref2av" && !input.referenceFile) { + return "Upload a reference image to use reference-guided mode."; + } + if (input.referenceFile && !isSupportedReferenceImage(input.referenceFile, input.capabilities)) { + return "Reference assets must be PNG, JPEG, or WebP images."; + } + if (input.firstFrameFile && !isSupportedReferenceImage(input.firstFrameFile, input.capabilities)) { + return "First frame must be a PNG, JPEG, or WebP image."; + } + if (input.lastFrameFile && !isSupportedReferenceImage(input.lastFrameFile, input.capabilities)) { + return "Last frame must be a PNG, JPEG, or WebP image."; + } + return null; +} diff --git a/apps/dreamverse/web/src/lib/creationConfig.test.ts b/apps/dreamverse/web/src/lib/creationConfig.test.ts index bff8d40ee2..c8ff618746 100644 --- a/apps/dreamverse/web/src/lib/creationConfig.test.ts +++ b/apps/dreamverse/web/src/lib/creationConfig.test.ts @@ -21,8 +21,8 @@ describe("creationConfig", () => { expect(formatDurationLabel(5)).toBe("5s"); }); - it("excludes H3 from lobby models", () => { - expect(CREATION_MODELS.map((model) => model.id)).toEqual(["fast-ltx23", "fast-ltx2"]); + it("includes all Dreamverse lobby models", () => { + expect(CREATION_MODELS.map((model) => model.id)).toEqual(["fast-ltx23", "fast-ltx2", "fast-h3"]); }); it("builds mention options from presets", () => { @@ -54,9 +54,9 @@ describe("creationConfig", () => { expect(modeUsesDualFrames("t2v")).toBe(false); }); - it("accepts image and video reference files", () => { + it("accepts image reference files only", () => { expect(isReferenceMediaFile(new File(["x"], "a.png", { type: "image/png" }))).toBe(true); - expect(isReferenceMediaFile(new File(["x"], "a.mp4", { type: "video/mp4" }))).toBe(true); + expect(isReferenceMediaFile(new File(["x"], "a.mp4", { type: "video/mp4" }))).toBe(false); expect(isReferenceMediaFile(new File(["x"], "a.txt", { type: "text/plain" }))).toBe(false); }); }); diff --git a/apps/dreamverse/web/src/lib/creationConfig.ts b/apps/dreamverse/web/src/lib/creationConfig.ts index dbe41c2a97..2223bd284a 100644 --- a/apps/dreamverse/web/src/lib/creationConfig.ts +++ b/apps/dreamverse/web/src/lib/creationConfig.ts @@ -1,6 +1,6 @@ export type CreationModeId = "t2v" | "fl2av" | "ref2av"; -export type CreationModelId = "fast-ltx2" | "fast-ltx23"; +export type CreationModelId = "fast-ltx2" | "fast-ltx23" | "fast-h3"; export type AspectRatioId = "21:9" | "16:9" | "4:3" | "1:1" | "3:4" | "9:16"; @@ -28,8 +28,11 @@ export interface MentionOption { export const CREATION_MODES: CreationModeOption[] = [ { id: "t2v", label: "Text to video", description: "Generate from a text prompt" }, - { id: "fl2av", label: "First and last frame", description: "Upload two assets as keyframes" }, - { id: "ref2av", label: "Omni reference", description: "Guide generation with a reference asset" }, + { id: "ref2av", label: "Image to video", description: "Guide the first segment with a reference image" }, +]; + +export const UNSUPPORTED_CREATION_MODES: CreationModeOption[] = [ + { id: "fl2av", label: "First and last frame", description: "Coming soon on FastLTX models" }, ]; export const CREATION_MODELS: CreationModelOption[] = [ @@ -44,15 +47,22 @@ export const CREATION_MODELS: CreationModelOption[] = [ label: "FastLTX 2", description: "FastLTX 2 for streaming", }, + { + id: "fast-h3", + label: "FastH3", + description: "MiniMax H3 with VSA data-free adapter", + }, ]; export const ASPECT_RATIOS: AspectRatioId[] = ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]; -export const RESOLUTIONS: ResolutionId[] = ["480p", "720p", "1080p", "4k"]; +export const RESOLUTIONS: ResolutionId[] = ["480p", "720p", "1080p"]; + +export const UNSUPPORTED_RESOLUTIONS: ResolutionId[] = ["4k"]; export const DURATION_MARKS = [5, 10, 15] as const; -export const REFERENCE_ACCEPT = "image/*,video/*"; +export const REFERENCE_ACCEPT = "image/png,image/jpeg,image/webp"; export function formatResolutionLabel(resolution: ResolutionId): string { return resolution === "4k" ? "4K" : resolution.toUpperCase(); @@ -71,7 +81,7 @@ export function modeUsesDualFrames(modeId: CreationModeId): boolean { } export function isReferenceMediaFile(file: File): boolean { - return file.type.startsWith("image/") || file.type.startsWith("video/"); + return file.type === "image/png" || file.type === "image/jpeg" || file.type === "image/webp"; } export function buildMentionOptions(storyPresets: Array<{ id?: string; label?: string; description?: string }>): MentionOption[] { diff --git a/apps/dreamverse/web/src/lib/creationPayload.test.ts b/apps/dreamverse/web/src/lib/creationPayload.test.ts new file mode 100644 index 0000000000..222ade77a2 --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationPayload.test.ts @@ -0,0 +1,66 @@ +import { describe, expect, it } from "vitest"; + +import { parseEchoedCreationConfig, validateCreationInputs } from "@/lib/creationPayload"; + +describe("creationPayload", () => { + it("requires a reference asset for omni reference mode", () => { + expect( + validateCreationInputs({ + modeId: "ref2av", + referenceFile: null, + }), + ).toMatch(/reference asset/i); + }); + + it("requires both frames for first and last frame mode", () => { + expect( + validateCreationInputs({ + modeId: "fl2av", + firstFrameFile: new File(["a"], "first.png", { type: "image/png" }), + lastFrameFile: null, + }), + ).toMatch(/both first and last/i); + }); + + it("accepts text to video without references", () => { + expect( + validateCreationInputs({ + modeId: "t2v", + }), + ).toBeNull(); + }); + + it("parses echoed creation config from server payloads", () => { + expect( + parseEchoedCreationConfig({ + type: "gpu_assigned", + creation_config: { + model_id: "fast-ltx2", + generation_mode: "ref2va", + aspect_ratio: "9:16", + resolution: "480p", + duration_sec: 10, + }, + }), + ).toEqual({ + modelId: "fast-ltx2", + modeId: "ref2av", + aspectRatio: "9:16", + resolution: "480p", + durationSec: 10, + }); + }); + + it("ignores invalid echoed creation config", () => { + expect(parseEchoedCreationConfig({ creation_config: { model_id: "unknown" } })).toBeNull(); + }); + + it("rejects unsupported reference mime types", () => { + expect( + validateCreationInputs({ + modeId: "t2v", + referenceFile: new File(["a"], "clip.mp4", { type: "video/mp4" }), + }), + ).toMatch(/PNG, JPEG, or WebP/i); + }); +}); diff --git a/apps/dreamverse/web/src/lib/creationPayload.ts b/apps/dreamverse/web/src/lib/creationPayload.ts new file mode 100644 index 0000000000..e5dbcab289 --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationPayload.ts @@ -0,0 +1,172 @@ +import type { + AspectRatioId, + CreationModeId, + CreationModelId, + ResolutionId, +} from "@/lib/creationConfig"; +import { fromGenerationMode, type GenerationMode } from "@/lib/generationMode"; + +const LOBBY_MODEL_IDS = new Set(["fast-ltx2", "fast-ltx23", "fast-h3"]); +const ASPECT_RATIO_IDS = new Set(["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]); +const RESOLUTION_IDS = new Set(["480p", "720p", "1080p", "4k"]); +const DURATION_SEC_VALUES = new Set([5, 10, 15]); + +export interface EchoedSessionCreationConfig { + modelId: CreationModelId; + modeId: CreationModeId; + aspectRatio: AspectRatioId; + resolution: ResolutionId; + durationSec: number; +} + +const MAX_IMAGE_BYTES = 15 * 1024 * 1024; +const SUPPORTED_IMAGE_TYPES = new Set(["image/png", "image/jpeg", "image/webp"]); + +export interface InitialImagePayload { + name: string; + mime_type: string; + data_url: string; +} + +export interface CreationInitPayload { + model_id: string; + aspect_ratio: string; + resolution: string; + duration_sec: number; + initial_image: InitialImagePayload | null; + last_frame_image: InitialImagePayload | null; +} + +function readFileAsDataUrl(file: File): Promise { + return new Promise((resolve, reject) => { + const reader = new FileReader(); + reader.onload = () => { + if (typeof reader.result === "string") { + resolve(reader.result); + return; + } + reject(new Error("Failed to read reference image.")); + }; + reader.onerror = () => reject(new Error("Failed to read reference image.")); + reader.readAsDataURL(file); + }); +} + +export async function fileToInitialImagePayload(file: File): Promise { + if (!SUPPORTED_IMAGE_TYPES.has(file.type)) { + throw new Error("Reference assets must be PNG, JPEG, or WebP images."); + } + if (file.size > MAX_IMAGE_BYTES) { + throw new Error("Reference image must be 15 MB or smaller."); + } + return { + name: file.name, + mime_type: file.type, + data_url: await readFileAsDataUrl(file), + }; +} + +export async function resolveCreationImages(input: { + modeId: CreationModeId; + referenceFile?: File | null; + firstFrameFile?: File | null; + lastFrameFile?: File | null; +}): Promise> { + if (input.modeId === "fl2av") { + const firstFrame = input.firstFrameFile ? await fileToInitialImagePayload(input.firstFrameFile) : null; + const lastFrame = input.lastFrameFile ? await fileToInitialImagePayload(input.lastFrameFile) : null; + return { + initial_image: firstFrame, + last_frame_image: lastFrame, + }; + } + + const reference = input.referenceFile ? await fileToInitialImagePayload(input.referenceFile) : null; + return { + initial_image: reference, + last_frame_image: null, + }; +} + +export function validateCreationInputs(input: { + modeId: CreationModeId; + referenceFile?: File | null; + firstFrameFile?: File | null; + lastFrameFile?: File | null; +}): string | null { + if (input.modeId === "ref2av" && !input.referenceFile) { + return "Upload a reference asset to use Omni reference mode."; + } + if (input.modeId === "fl2av") { + if (!input.firstFrameFile || !input.lastFrameFile) { + return "Upload both first and last frame assets."; + } + } + if (input.referenceFile && !SUPPORTED_IMAGE_TYPES.has(input.referenceFile.type)) { + return "Reference assets must be PNG, JPEG, or WebP images."; + } + if (input.firstFrameFile && !SUPPORTED_IMAGE_TYPES.has(input.firstFrameFile.type)) { + return "First frame must be a PNG, JPEG, or WebP image."; + } + if (input.lastFrameFile && !SUPPORTED_IMAGE_TYPES.has(input.lastFrameFile.type)) { + return "Last frame must be a PNG, JPEG, or WebP image."; + } + return null; +} + +export function parseEchoedCreationConfig(data: unknown): EchoedSessionCreationConfig | null { + if (!data || typeof data !== "object") { + return null; + } + const creationConfig = (data as Record).creation_config; + if (!creationConfig || typeof creationConfig !== "object") { + return null; + } + const config = creationConfig as Record; + const modelId = typeof config.model_id === "string" && LOBBY_MODEL_IDS.has(config.model_id as CreationModelId) + ? (config.model_id as CreationModelId) + : null; + const generationMode = typeof config.generation_mode === "string" ? config.generation_mode as GenerationMode : null; + const modeId = generationMode === "t2va" || generationMode === "fl2va" || generationMode === "ref2va" + ? fromGenerationMode(generationMode) + : null; + const aspectRatio = typeof config.aspect_ratio === "string" && ASPECT_RATIO_IDS.has(config.aspect_ratio as AspectRatioId) + ? (config.aspect_ratio as AspectRatioId) + : null; + const resolution = typeof config.resolution === "string" && RESOLUTION_IDS.has(config.resolution as ResolutionId) + ? (config.resolution as ResolutionId) + : null; + const durationSec = typeof config.duration_sec === "number" && DURATION_SEC_VALUES.has(config.duration_sec) + ? config.duration_sec + : null; + if (modelId === null || modeId === null || aspectRatio === null || resolution === null || durationSec === null) { + return null; + } + return { + modelId, + modeId, + aspectRatio, + resolution, + durationSec, + }; +} + +export async function buildCreationInitPayload(input: { + modelId: string; + modeId: CreationModeId; + aspectRatio: string; + resolution: string; + durationSec: number; + referenceFile?: File | null; + firstFrameFile?: File | null; + lastFrameFile?: File | null; +}): Promise { + const images = await resolveCreationImages(input); + return { + model_id: input.modelId, + aspect_ratio: input.aspectRatio, + resolution: input.resolution, + duration_sec: input.durationSec, + ...images, + }; +} diff --git a/apps/dreamverse/web/src/lib/generationMode.test.ts b/apps/dreamverse/web/src/lib/generationMode.test.ts index c35b551970..1721d547c6 100644 --- a/apps/dreamverse/web/src/lib/generationMode.test.ts +++ b/apps/dreamverse/web/src/lib/generationMode.test.ts @@ -3,6 +3,7 @@ import { describe, expect, it } from "vitest"; import { DEFAULT_GENERATION_MODE, GENERATION_MODES, + fromGenerationMode, getGenerationMode, isGenerationMode, toGenerationMode, @@ -29,4 +30,10 @@ describe("generation modes", () => { expect(toGenerationMode("fl2av")).toBe("fl2va"); expect(toGenerationMode("ref2av")).toBe("ref2va"); }); + + it("maps upstream wire values back to creation studio mode IDs", () => { + expect(fromGenerationMode("t2va")).toBe("t2v"); + expect(fromGenerationMode("fl2va")).toBe("fl2av"); + expect(fromGenerationMode("ref2va")).toBe("ref2av"); + }); }); diff --git a/apps/dreamverse/web/src/lib/generationMode.ts b/apps/dreamverse/web/src/lib/generationMode.ts index 0d491eee5f..c7902e350d 100644 --- a/apps/dreamverse/web/src/lib/generationMode.ts +++ b/apps/dreamverse/web/src/lib/generationMode.ts @@ -39,6 +39,16 @@ export function getGenerationMode(value: GenerationMode) { return GENERATION_MODES.find((mode) => mode.id === value) ?? GENERATION_MODES[0]; } +const GENERATION_MODE_TO_CREATION_MODE: Record = { + t2va: "t2v", + fl2va: "fl2av", + ref2va: "ref2av", +}; + +export function fromGenerationMode(mode: GenerationMode): CreationModeId { + return GENERATION_MODE_TO_CREATION_MODE[mode]; +} + export function toGenerationMode(modeId: CreationModeId): GenerationMode { return CREATION_MODE_TO_GENERATION_MODE[modeId]; }