From 811e1c0a253098f400656d8892742b35c211c0e0 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 8 Sep 2026 04:51:12 +0000 Subject: [PATCH 1/5] [feat]: wire creation studio config to session init and generation Connect lobby model/mode/aspect/resolution/duration and reference images to session_init_v2 and project_init_v1. Backend parses creation config, validates mode assets, resolves frame size and segment cap, and passes per-step sampling overrides to the GPU worker. Frontend builds init payloads from creation state and applies echoed creation_config. Co-authored-by: Aryan Kumar --- apps/dreamverse/arch.md | 35 +++- .../dreamverse/generation_contracts.py | 4 + .../dreamverse/generation_worker.py | 7 + apps/dreamverse/dreamverse/gpu_pool.py | 10 + apps/dreamverse/dreamverse/ltx2_generation.py | 11 +- .../dreamverse/minimax_h3_generation.py | 5 + apps/dreamverse/dreamverse/mock_server.py | 64 ++++++- .../dreamverse/session/controller.py | 149 ++++++++++----- .../dreamverse/session_creation_config.py | 133 ++++++++++++++ .../tests/test_session_creation_config.py | 52 ++++++ apps/dreamverse/dreamverse/worker_ipc.py | 3 + apps/dreamverse/web/src/app/page.tsx | 76 +++++++- .../web/src/lib/creationPayload.test.ts | 66 +++++++ .../dreamverse/web/src/lib/creationPayload.ts | 173 ++++++++++++++++++ 14 files changed, 724 insertions(+), 64 deletions(-) create mode 100644 apps/dreamverse/dreamverse/session_creation_config.py create mode 100644 apps/dreamverse/dreamverse/tests/test_session_creation_config.py create mode 100644 apps/dreamverse/web/src/lib/creationPayload.test.ts create mode 100644 apps/dreamverse/web/src/lib/creationPayload.ts diff --git a/apps/dreamverse/arch.md b/apps/dreamverse/arch.md index eb532f17d8..87d264848e 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`: + +- `model_id`: `fast-ltx2` or `fast-ltx23` +- `creation_mode`: `t2v`, `fl2av`, or `ref2av` +- `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: + +- `ref2av` requires `initial_image` +- `fl2av` 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/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..cce9ce3274 100644 --- a/apps/dreamverse/dreamverse/ltx2_generation.py +++ b/apps/dreamverse/dreamverse/ltx2_generation.py @@ -454,8 +454,13 @@ 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.""" + del frame_width, frame_height, num_frames timings: dict = {} prompt = self._inject_style_trigger(prompt) @@ -464,9 +469,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/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..127f79c33a 100644 --- a/apps/dreamverse/dreamverse/mock_server.py +++ b/apps/dreamverse/dreamverse/mock_server.py @@ -31,6 +31,7 @@ from dreamverse._deps import require_dreamverse_runtime_deps from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP +from dreamverse.session_creation_config import parse_session_creation_config, validate_creation_mode_assets from dreamverse.session_init_image import cleanup_session_init_image, persist_session_init_image LATENCY_MS = 200 @@ -290,6 +291,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 +351,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 +360,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_creation_mode_assets( + session_creation_config.creation_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 +413,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 +430,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 +546,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 +558,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_creation_mode_assets( + session_creation_config.creation_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 +1235,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/session/controller.py b/apps/dreamverse/dreamverse/session/controller.py index e3985938d8..8eae5295f9 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_creation_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_creation_mode_assets( + session_creation_config.creation_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.creation_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_creation_mode_assets( + session_creation_config.creation_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..ad1d81b61f --- /dev/null +++ b/apps/dreamverse/dreamverse/session_creation_config.py @@ -0,0 +1,133 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from dreamverse.config import FRAME_HEIGHT, FRAME_WIDTH, GENERATION_SEGMENT_CAP, MODEL_REGISTRY, NUM_FRAMES + +LOBBY_MODEL_IDS = frozenset({"fast-ltx2", "fast-ltx23"}) +SUPPORTED_CREATION_MODES = frozenset({"t2v", "fl2av", "ref2av"}) +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 + creation_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, + "creation_mode": self.creation_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 LOBBY_MODEL_IDS else "fast-ltx23" + + creation_mode = str(payload.get("creation_mode") or "t2v").strip() + if creation_mode not in SUPPORTED_CREATION_MODES: + raise ValueError(f"Unsupported creation_mode: {creation_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.") + + 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, + creation_mode=creation_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_creation_mode_assets( + creation_mode: str, + *, + has_initial_image: bool, + has_last_frame_image: bool, +) -> None: + if creation_mode == "ref2av" and not has_initial_image: + raise ValueError("Omni reference mode requires a reference image.") + if creation_mode == "fl2av" and (not has_initial_image or not has_last_frame_image): + raise ValueError("First and last frame mode requires both frame images.") 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..78ea9e6841 --- /dev/null +++ b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py @@ -0,0 +1,52 @@ +import pytest + +from dreamverse.session_creation_config import ( + duration_sec_to_segment_cap, + parse_session_creation_config, + resolve_frame_size, + validate_creation_mode_assets, +) + + +def test_parse_session_creation_config_defaults(): + config = parse_session_creation_config({}) + assert config.model_id == "fast-ltx23" + assert config.creation_mode == "t2v" + 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", + "creation_mode": "ref2av", + "aspect_ratio": "9:16", + "resolution": "480p", + "duration_sec": 15, + }, + ) + assert config.model_id == "fast-ltx2" + assert config.creation_mode == "ref2av" + 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_validate_creation_mode_assets(): + validate_creation_mode_assets("t2v", has_initial_image=False, has_last_frame_image=False) + with pytest.raises(ValueError, match="Omni reference"): + validate_creation_mode_assets("ref2av", has_initial_image=False, has_last_frame_image=False) + with pytest.raises(ValueError, match="First and last frame"): + validate_creation_mode_assets("fl2av", has_initial_image=True, 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/src/app/page.tsx b/apps/dreamverse/web/src/app/page.tsx index 65c1172ec2..be45a053cd 100644 --- a/apps/dreamverse/web/src/app/page.tsx +++ b/apps/dreamverse/web/src/app/page.tsx @@ -31,6 +31,11 @@ import { buildRewritePromptWindowSnapshotFromPrompts, normalizePromptWindowSnapshot, } from "@/lib/prompts/promptWindowSnapshot"; +import { + buildCreationInitPayload, + parseEchoedCreationConfig, + validateCreationInputs, +} from "@/lib/creationPayload"; import rawPresets from "@/lib/storyPresetsData"; import { cn } from "@/lib/utils"; import { createWebSocketConnection, detachAndCloseWebSocket } from "@/lib/ws/client"; @@ -240,6 +245,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"); @@ -299,14 +307,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); } @@ -1739,33 +1750,43 @@ 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, preset_id: getInitialPresetId(), 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() { @@ -1802,6 +1823,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, @@ -1844,7 +1868,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 @@ -1927,6 +1956,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; @@ -1990,13 +2027,29 @@ export default function Page() { } async function joinSession({ force = false } = {}) { + const validationError = validateCreationInputs({ + modeId: creationModeId, + 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(""); @@ -2018,7 +2071,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(); 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..faab083003 --- /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", + creation_mode: "ref2av", + 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..0ea13edb6a --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationPayload.ts @@ -0,0 +1,173 @@ +import type { + AspectRatioId, + CreationModeId, + CreationModelId, + ResolutionId, +} from "@/lib/creationConfig"; + +const LOBBY_MODEL_IDS = new Set(["fast-ltx2", "fast-ltx23"]); +const CREATION_MODE_IDS = new Set(["t2v", "fl2av", "ref2av"]); +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; + creation_mode: CreationModeId; + 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 modeId = typeof config.creation_mode === "string" && CREATION_MODE_IDS.has(config.creation_mode as CreationModeId) + ? (config.creation_mode as CreationModeId) + : 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, + creation_mode: input.modeId, + aspect_ratio: input.aspectRatio, + resolution: input.resolution, + duration_sec: input.durationSec, + ...images, + }; +} From 4121de9a9f1119185e9548e9517c1a7fb086f278 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Fri, 11 Sep 2026 03:06:56 +0000 Subject: [PATCH 2/5] [feat]: align backend wiring with upstream generation_mode contract Use canonical t2va/fl2va/ref2va from #1834/#1835 instead of a parallel creation_mode field. Wire model, resolution, duration, and image assets through session_creation_config while leaving H3 mode routing to upstream. Co-authored-by: Aryan Kumar --- apps/dreamverse/arch.md | 6 ++--- apps/dreamverse/dreamverse/mock_server.py | 10 +++---- .../dreamverse/session/controller.py | 12 ++++----- .../dreamverse/session_creation_config.py | 26 +++++++++---------- .../tests/test_session_creation_config.py | 20 +++++++------- .../web/src/lib/creationPayload.test.ts | 2 +- .../dreamverse/web/src/lib/creationPayload.ts | 9 +++---- .../web/src/lib/generationMode.test.ts | 7 +++++ apps/dreamverse/web/src/lib/generationMode.ts | 10 +++++++ 9 files changed, 59 insertions(+), 43 deletions(-) diff --git a/apps/dreamverse/arch.md b/apps/dreamverse/arch.md index 87d264848e..24860c43d2 100644 --- a/apps/dreamverse/arch.md +++ b/apps/dreamverse/arch.md @@ -321,8 +321,8 @@ creation config and echoes the resolved values back on `gpu_assigned` and 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` or `fast-ltx23` -- `creation_mode`: `t2v`, `fl2av`, or `ref2av` - `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` @@ -334,8 +334,8 @@ Echoed `creation_config` includes the resolved frame size, Mode validation: -- `ref2av` requires `initial_image` -- `fl2av` requires both `initial_image` and `last_frame_image` +- `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. diff --git a/apps/dreamverse/dreamverse/mock_server.py b/apps/dreamverse/dreamverse/mock_server.py index 127f79c33a..e2e85badbc 100644 --- a/apps/dreamverse/dreamverse/mock_server.py +++ b/apps/dreamverse/dreamverse/mock_server.py @@ -31,7 +31,7 @@ from dreamverse._deps import require_dreamverse_runtime_deps from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP -from dreamverse.session_creation_config import parse_session_creation_config, validate_creation_mode_assets +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 @@ -362,8 +362,8 @@ async def session_timeout() -> None: try: session_creation_config = parse_session_creation_config(init_data) - validate_creation_mode_assets( - session_creation_config.creation_mode, + 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, ) @@ -568,8 +568,8 @@ async def apply_project_init_payload(payload: dict[str, object], ) -> bool: try: session_creation_config = parse_session_creation_config(payload) - validate_creation_mode_assets( - session_creation_config.creation_mode, + 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, ) diff --git a/apps/dreamverse/dreamverse/session/controller.py b/apps/dreamverse/dreamverse/session/controller.py index 8eae5295f9..461fa19b22 100644 --- a/apps/dreamverse/dreamverse/session/controller.py +++ b/apps/dreamverse/dreamverse/session/controller.py @@ -27,7 +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_creation_mode_assets +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 ( @@ -253,8 +253,8 @@ async def cancel_task(task: asyncio.Task | None): try: session_creation_config = parse_session_creation_config(init_data) session_generation_segment_cap = session_creation_config.generation_segment_cap - validate_creation_mode_assets( - session_creation_config.creation_mode, + 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, ) @@ -283,7 +283,7 @@ async def cancel_task(task: asyncio.Task | None): 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.creation_mode}, " + 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}") @@ -556,8 +556,8 @@ async def apply_project_init_payload(payload: dict[str, object]) -> bool: try: session_creation_config = parse_session_creation_config(payload) session_generation_segment_cap = session_creation_config.generation_segment_cap - validate_creation_mode_assets( - session_creation_config.creation_mode, + 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, ) diff --git a/apps/dreamverse/dreamverse/session_creation_config.py b/apps/dreamverse/dreamverse/session_creation_config.py index ad1d81b61f..9207b571f4 100644 --- a/apps/dreamverse/dreamverse/session_creation_config.py +++ b/apps/dreamverse/dreamverse/session_creation_config.py @@ -5,7 +5,7 @@ from dreamverse.config import FRAME_HEIGHT, FRAME_WIDTH, GENERATION_SEGMENT_CAP, MODEL_REGISTRY, NUM_FRAMES LOBBY_MODEL_IDS = frozenset({"fast-ltx2", "fast-ltx23"}) -SUPPORTED_CREATION_MODES = frozenset({"t2v", "fl2av", "ref2av"}) +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 @@ -14,7 +14,7 @@ @dataclass(frozen=True) class SessionCreationConfig: model_id: str - creation_mode: str + generation_mode: str aspect_ratio: str resolution: str duration_sec: int @@ -26,7 +26,7 @@ class SessionCreationConfig: def as_dict(self) -> dict[str, object]: return { "model_id": self.model_id, - "creation_mode": self.creation_mode, + "generation_mode": self.generation_mode, "aspect_ratio": self.aspect_ratio, "resolution": self.resolution, "duration_sec": self.duration_sec, @@ -85,9 +85,9 @@ def parse_session_creation_config(payload: dict[str, object]) -> SessionCreation raw_model_id = str(payload.get("model_id") or "").strip() model_id = raw_model_id if raw_model_id in LOBBY_MODEL_IDS else "fast-ltx23" - creation_mode = str(payload.get("creation_mode") or "t2v").strip() - if creation_mode not in SUPPORTED_CREATION_MODES: - raise ValueError(f"Unsupported creation_mode: {creation_mode}") + 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: @@ -110,7 +110,7 @@ def parse_session_creation_config(payload: dict[str, object]) -> SessionCreation frame_width, frame_height = resolve_frame_size(aspect_ratio, resolution) return SessionCreationConfig( model_id=model_id, - creation_mode=creation_mode, + generation_mode=generation_mode, aspect_ratio=aspect_ratio, resolution=resolution, duration_sec=duration_sec, @@ -121,13 +121,13 @@ def parse_session_creation_config(payload: dict[str, object]) -> SessionCreation ) -def validate_creation_mode_assets( - creation_mode: str, +def validate_generation_mode_assets( + generation_mode: str, *, has_initial_image: bool, has_last_frame_image: bool, ) -> None: - if creation_mode == "ref2av" and not has_initial_image: - raise ValueError("Omni reference mode requires a reference image.") - if creation_mode == "fl2av" and (not has_initial_image or not has_last_frame_image): - raise ValueError("First and last frame mode requires both frame images.") + if generation_mode == "ref2va" and not has_initial_image: + raise ValueError("Ref2VA mode requires a reference image.") + if generation_mode == "fl2va" and (not has_initial_image or not has_last_frame_image): + raise ValueError("FL2VA mode requires both frame images.") diff --git a/apps/dreamverse/dreamverse/tests/test_session_creation_config.py b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py index 78ea9e6841..fa09a09649 100644 --- a/apps/dreamverse/dreamverse/tests/test_session_creation_config.py +++ b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py @@ -4,14 +4,14 @@ duration_sec_to_segment_cap, parse_session_creation_config, resolve_frame_size, - validate_creation_mode_assets, + validate_generation_mode_assets, ) def test_parse_session_creation_config_defaults(): config = parse_session_creation_config({}) assert config.model_id == "fast-ltx23" - assert config.creation_mode == "t2v" + assert config.generation_mode == "t2va" assert config.aspect_ratio == "16:9" assert config.resolution == "720p" assert config.duration_sec == 5 @@ -22,14 +22,14 @@ def test_parse_session_creation_config_maps_duration_to_segment_cap(): config = parse_session_creation_config( { "model_id": "fast-ltx2", - "creation_mode": "ref2av", + "generation_mode": "ref2va", "aspect_ratio": "9:16", "resolution": "480p", "duration_sec": 15, }, ) assert config.model_id == "fast-ltx2" - assert config.creation_mode == "ref2av" + assert config.generation_mode == "ref2va" assert config.generation_segment_cap == 3 assert config.frame_width >= 480 assert config.frame_height >= 480 @@ -44,9 +44,9 @@ def test_duration_sec_to_segment_cap_respects_global_cap(): assert duration_sec_to_segment_cap(15, global_cap=2) == 2 -def test_validate_creation_mode_assets(): - validate_creation_mode_assets("t2v", has_initial_image=False, has_last_frame_image=False) - with pytest.raises(ValueError, match="Omni reference"): - validate_creation_mode_assets("ref2av", has_initial_image=False, has_last_frame_image=False) - with pytest.raises(ValueError, match="First and last frame"): - validate_creation_mode_assets("fl2av", has_initial_image=True, has_last_frame_image=False) +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) + with pytest.raises(ValueError, match="FL2VA"): + validate_generation_mode_assets("fl2va", has_initial_image=True, has_last_frame_image=False) diff --git a/apps/dreamverse/web/src/lib/creationPayload.test.ts b/apps/dreamverse/web/src/lib/creationPayload.test.ts index faab083003..222ade77a2 100644 --- a/apps/dreamverse/web/src/lib/creationPayload.test.ts +++ b/apps/dreamverse/web/src/lib/creationPayload.test.ts @@ -36,7 +36,7 @@ describe("creationPayload", () => { type: "gpu_assigned", creation_config: { model_id: "fast-ltx2", - creation_mode: "ref2av", + generation_mode: "ref2va", aspect_ratio: "9:16", resolution: "480p", duration_sec: 10, diff --git a/apps/dreamverse/web/src/lib/creationPayload.ts b/apps/dreamverse/web/src/lib/creationPayload.ts index 0ea13edb6a..a03c7e9f1e 100644 --- a/apps/dreamverse/web/src/lib/creationPayload.ts +++ b/apps/dreamverse/web/src/lib/creationPayload.ts @@ -4,9 +4,9 @@ import type { CreationModelId, ResolutionId, } from "@/lib/creationConfig"; +import { fromGenerationMode, type GenerationMode } from "@/lib/generationMode"; const LOBBY_MODEL_IDS = new Set(["fast-ltx2", "fast-ltx23"]); -const CREATION_MODE_IDS = new Set(["t2v", "fl2av", "ref2av"]); 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]); @@ -30,7 +30,6 @@ export interface InitialImagePayload { export interface CreationInitPayload { model_id: string; - creation_mode: CreationModeId; aspect_ratio: string; resolution: string; duration_sec: number; @@ -127,8 +126,9 @@ export function parseEchoedCreationConfig(data: unknown): EchoedSessionCreationC const modelId = typeof config.model_id === "string" && LOBBY_MODEL_IDS.has(config.model_id as CreationModelId) ? (config.model_id as CreationModelId) : null; - const modeId = typeof config.creation_mode === "string" && CREATION_MODE_IDS.has(config.creation_mode as CreationModeId) - ? (config.creation_mode as CreationModeId) + 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) @@ -164,7 +164,6 @@ export async function buildCreationInitPayload(input: { const images = await resolveCreationImages(input); return { model_id: input.modelId, - creation_mode: input.modeId, aspect_ratio: input.aspectRatio, resolution: input.resolution, duration_sec: input.durationSec, 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]; } From fee44648c0327a7df92fc1700a088a1f74a76960 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Fri, 11 Sep 2026 03:28:18 +0000 Subject: [PATCH 3/5] [test]: assert creation studio placeholder instead of invisible text Playwright was matching placeholder copy as visible text; the lobby composer only exposes it via the textarea placeholder attribute. Co-authored-by: Aryan Kumar --- apps/dreamverse/web/e2e/frontend-shell.spec.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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); }); }); From e28b9efa427e19fce2a407cfaeef4ddd0cbcfbaa Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Fri, 11 Sep 2026 03:32:03 +0000 Subject: [PATCH 4/5] [feat]: gate creation studio on FastLTX lobby capabilities Expose GET /creation-capabilities from backend and mirror it in the frontend so only supported models, modes (t2va/ref2va), resolutions (480p-1080p), and image references are selectable. Reject FL2VA, 4K, and invalid assets client-side and in session_creation_config. Lock in-session config pills read-only until live updates exist. Co-authored-by: Aryan Kumar --- .../dreamverse/creation_capabilities.py | 82 ++++++++ apps/dreamverse/dreamverse/ltx2_generation.py | 1 - apps/dreamverse/dreamverse/main.py | 2 + apps/dreamverse/dreamverse/mock_server.py | 6 + apps/dreamverse/dreamverse/routes/creation.py | 14 ++ .../dreamverse/session_creation_config.py | 15 +- .../tests/test_creation_capabilities.py | 44 +++++ .../tests/test_session_creation_config.py | 26 ++- apps/dreamverse/web/next.config.ts | 4 + apps/dreamverse/web/src/app/page.tsx | 63 +++++- .../components/creation/CreationComposer.tsx | 79 +++++++- .../components/creation/CreationStudio.tsx | 5 +- .../web/src/lib/creationCapabilities.test.ts | 71 +++++++ .../web/src/lib/creationCapabilities.ts | 184 ++++++++++++++++++ .../web/src/lib/creationConfig.test.ts | 4 +- apps/dreamverse/web/src/lib/creationConfig.ts | 15 +- 16 files changed, 589 insertions(+), 26 deletions(-) create mode 100644 apps/dreamverse/dreamverse/creation_capabilities.py create mode 100644 apps/dreamverse/dreamverse/routes/creation.py create mode 100644 apps/dreamverse/dreamverse/tests/test_creation_capabilities.py create mode 100644 apps/dreamverse/web/src/lib/creationCapabilities.test.ts create mode 100644 apps/dreamverse/web/src/lib/creationCapabilities.ts diff --git a/apps/dreamverse/dreamverse/creation_capabilities.py b/apps/dreamverse/dreamverse/creation_capabilities.py new file mode 100644 index 0000000000..02ce1d399d --- /dev/null +++ b/apps/dreamverse/dreamverse/creation_capabilities.py @@ -0,0 +1,82 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from dreamverse.config import MODEL_REGISTRY + +# FastLTX lobby models share the same creation-studio surface today. +LTX_LOBBY_MODEL_IDS = frozenset({"fast-ltx2", "fast-ltx23"}) + +# Canonical upstream wire IDs. FL2VA is tracked in #1834 but not implemented on +# FastLTX streaming yet (last-frame conditioning is not wired). +LTX_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"}) + +LTX_LOBBY_DURATION_SEC = frozenset({5, 10, 15}) + +UNSUPPORTED_GENERATION_MODE_MESSAGES = { + "fl2va": "First/last frame mode (FL2VA) is not supported on FastLTX models yet.", +} + + +@dataclass(frozen=True) +class LobbyCreationCapabilities: + model_ids: frozenset[str] + generation_modes: frozenset[str] + aspect_ratios: frozenset[str] + resolutions: frozenset[str] + duration_sec: frozenset[int] + + def as_dict(self) -> dict[str, object]: + return { + "model_ids": sorted(self.model_ids), + "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": dict(UNSUPPORTED_GENERATION_MODE_MESSAGES), + "reference_assets": { + "mime_types": ["image/png", "image/jpeg", "image/webp"], + "max_bytes": 15 * 1024 * 1024, + }, + } + + +LOBBY_CREATION_CAPABILITIES = LobbyCreationCapabilities( + model_ids=LTX_LOBBY_MODEL_IDS, + generation_modes=LTX_LOBBY_GENERATION_MODES, + aspect_ratios=LTX_LOBBY_ASPECT_RATIOS, + resolutions=LTX_LOBBY_RESOLUTIONS, + duration_sec=LTX_LOBBY_DURATION_SEC, +) + + +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 LOBBY_CREATION_CAPABILITIES.model_ids: + raise ValueError(f"Unsupported model_id: {model_id}") + if model_id not in MODEL_REGISTRY: + raise ValueError(f"Unknown model_id: {model_id}") + + if generation_mode in UNSUPPORTED_GENERATION_MODE_MESSAGES: + raise ValueError(UNSUPPORTED_GENERATION_MODE_MESSAGES[generation_mode]) + if generation_mode not in LOBBY_CREATION_CAPABILITIES.generation_modes: + raise ValueError(f"Unsupported generation_mode: {generation_mode}") + + if aspect_ratio not in LOBBY_CREATION_CAPABILITIES.aspect_ratios: + raise ValueError(f"Unsupported aspect_ratio: {aspect_ratio}") + if resolution not in LOBBY_CREATION_CAPABILITIES.resolutions: + raise ValueError(f"Unsupported resolution: {resolution}") + if duration_sec not in LOBBY_CREATION_CAPABILITIES.duration_sec: + raise ValueError("duration_sec must be 5, 10, or 15.") diff --git a/apps/dreamverse/dreamverse/ltx2_generation.py b/apps/dreamverse/dreamverse/ltx2_generation.py index cce9ce3274..ecd9caf5ed 100644 --- a/apps/dreamverse/dreamverse/ltx2_generation.py +++ b/apps/dreamverse/dreamverse/ltx2_generation.py @@ -460,7 +460,6 @@ def generate_step( num_frames: int | None = None, ) -> StepResult: """Execute one generation step; snapshot state for the next segment.""" - del frame_width, frame_height, num_frames timings: dict = {} prompt = self._inject_style_trigger(prompt) 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/mock_server.py b/apps/dreamverse/dreamverse/mock_server.py index e2e85badbc..329dc334cf 100644 --- a/apps/dreamverse/dreamverse/mock_server.py +++ b/apps/dreamverse/dreamverse/mock_server.py @@ -31,6 +31,7 @@ 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_CREATION_CAPABILITIES 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 @@ -226,6 +227,11 @@ async def prompt_system_config(): } +@app.get("/creation-capabilities") +async def creation_capabilities(): + return LOBBY_CREATION_CAPABILITIES.as_dict() + + @app.get("/curated-presets") async def curated_presets(): presets = [ diff --git a/apps/dreamverse/dreamverse/routes/creation.py b/apps/dreamverse/dreamverse/routes/creation.py new file mode 100644 index 0000000000..30434b4b36 --- /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_CREATION_CAPABILITIES + +creation_router = APIRouter(tags=["creation"]) + + +@creation_router.get("/creation-capabilities") +async def creation_capabilities() -> dict[str, object]: + return LOBBY_CREATION_CAPABILITIES.as_dict() diff --git a/apps/dreamverse/dreamverse/session_creation_config.py b/apps/dreamverse/dreamverse/session_creation_config.py index 9207b571f4..dfd52b8e9a 100644 --- a/apps/dreamverse/dreamverse/session_creation_config.py +++ b/apps/dreamverse/dreamverse/session_creation_config.py @@ -3,8 +3,9 @@ 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 -LOBBY_MODEL_IDS = frozenset({"fast-ltx2", "fast-ltx23"}) +LTX_LOBBY_MODEL_IDS = frozenset({"fast-ltx2", "fast-ltx23"}) 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"}) @@ -83,7 +84,7 @@ def duration_sec_to_segment_cap(duration_sec: int, *, global_cap: int = GENERATI 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 LOBBY_MODEL_IDS else "fast-ltx23" + 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: @@ -104,6 +105,14 @@ def parse_session_creation_config(payload: dict[str, object]) -> SessionCreation 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}") @@ -129,5 +138,3 @@ def validate_generation_mode_assets( ) -> None: if generation_mode == "ref2va" and not has_initial_image: raise ValueError("Ref2VA mode requires a reference image.") - if generation_mode == "fl2va" and (not has_initial_image or not has_last_frame_image): - raise ValueError("FL2VA mode requires both frame images.") 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..374d990d9a --- /dev/null +++ b/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py @@ -0,0 +1,44 @@ +import pytest + +from dreamverse.creation_capabilities import ( + LOBBY_CREATION_CAPABILITIES, + validate_lobby_creation_config, +) + + +def test_lobby_capabilities_exclude_fl2va_and_4k(): + assert "fl2va" not in LOBBY_CREATION_CAPABILITIES.generation_modes + assert "4k" not in LOBBY_CREATION_CAPABILITIES.resolutions + assert LOBBY_CREATION_CAPABILITIES.generation_modes == frozenset({"t2va", "ref2va"}) + + +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_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, + ) diff --git a/apps/dreamverse/dreamverse/tests/test_session_creation_config.py b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py index fa09a09649..cba7741661 100644 --- a/apps/dreamverse/dreamverse/tests/test_session_creation_config.py +++ b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py @@ -44,9 +44,31 @@ 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_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) - with pytest.raises(ValueError, match="FL2VA"): - validate_generation_mode_assets("fl2va", has_initial_image=True, has_last_frame_image=False) 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 0887fa04e9..782d399f40 100644 --- a/apps/dreamverse/web/src/app/page.tsx +++ b/apps/dreamverse/web/src/app/page.tsx @@ -32,10 +32,16 @@ import { buildRewritePromptWindowSnapshotFromPrompts, normalizePromptWindowSnapshot, } from "@/lib/prompts/promptWindowSnapshot"; +import { + DEFAULT_LOBBY_CREATION_CAPABILITIES, + clampLobbySelectionToCapabilities, + parseLobbyCreationCapabilities, + validateLobbyCreationSelection, + type LobbyCreationCapabilities, +} from "@/lib/creationCapabilities"; import { buildCreationInitPayload, parseEchoedCreationConfig, - validateCreationInputs, } from "@/lib/creationPayload"; import rawPresets from "@/lib/storyPresetsData"; import { cn } from "@/lib/utils"; @@ -256,6 +262,9 @@ export default function Page() { const [creationAspectRatio, setCreationAspectRatio] = useState("16:9"); const [creationResolution, setCreationResolution] = useState("720p"); const [creationDurationSec, setCreationDurationSec] = useState(5); + const [lobbyCapabilities, setLobbyCapabilities] = useState( + DEFAULT_LOBBY_CREATION_CAPABILITIES, + ); const [sessionCreationConfig, setSessionCreationConfig] = useState({ modelId: "fast-ltx23", modeId: "t2v", @@ -516,6 +525,49 @@ export default function Page() { setRuntimeReady(true); }, []); + function applyLobbyCapabilities(capabilities: LobbyCreationCapabilities) { + setLobbyCapabilities(capabilities); + const clamped = clampLobbySelectionToCapabilities({ + capabilities, + 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); + } + + 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_CREATION_CAPABILITIES; + return parseLobbyCreationCapabilities(await response.json()); + }) + .then((capabilities) => { + if (!cancelled) { + applyLobbyCapabilities(capabilities); + } + }) + .catch(() => { + if (!cancelled) { + applyLobbyCapabilities(DEFAULT_LOBBY_CREATION_CAPABILITIES); + } + }); + return () => { + cancelled = true; + }; + }, [runtimeReady]); + useEffect(() => { if (!runtimeReady || initializedRef.current) return; initializedRef.current = true; @@ -2029,8 +2081,13 @@ export default function Page() { } async function joinSession({ force = false } = {}) { - const validationError = validateCreationInputs({ + const validationError = validateLobbyCreationSelection({ + capabilities: lobbyCapabilities, + modelId: creationModelId, modeId: creationModeId, + aspectRatio: creationAspectRatio, + resolution: creationResolution, + durationSec: creationDurationSec, referenceFile: referenceFileRef.current, firstFrameFile: firstFrameFileRef.current, lastFrameFile: lastFrameFileRef.current, @@ -2825,6 +2882,7 @@ export default function Page() { lastFramePreviewUrl={lastFramePreviewUrl} mentionOptions={mentionOptions} storyPresets={lobbyStoryPresets} + capabilities={lobbyCapabilities} onValueChange={(value) => sessionStore.patch({ livePromptDraft: value })} onSubmit={() => void joinSession()} onKeyDown={handleLivePromptKeydown} @@ -2855,6 +2913,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..11ffeddd79 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,13 @@ import { formatDurationLabel, formatResolutionLabel, } from "@/lib/creationConfig"; +import { + DEFAULT_LOBBY_CREATION_CAPABILITIES, + isSupportedCreationMode, + isSupportedResolution, + unsupportedModeNotice, + type LobbyCreationCapabilities, +} from "@/lib/creationCapabilities"; import { cn } from "@/lib/utils"; const PROMPT_MAX_LENGTH = 500; @@ -64,6 +73,7 @@ interface CreationComposerProps { onLastFrameSelect?: (file: File | null) => void; onSpeechTranscript?: (text: string) => void; onSpeechInterimChange?: (text: string) => void; + capabilities?: LobbyCreationCapabilities; } export default function CreationComposer({ @@ -93,6 +103,7 @@ export default function CreationComposer({ onLastFrameSelect, onSpeechTranscript, onSpeechInterimChange, + capabilities = DEFAULT_LOBBY_CREATION_CAPABILITIES, }: CreationComposerProps) { const inputRef = useRef(null); const [sttBusy, setSttBusy] = useState(false); @@ -100,8 +111,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 +319,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 +342,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 +370,7 @@ export default function CreationComposer({

Aspect ratio

- {ASPECT_RATIOS.map((ratio) => ( + {availableAspectRatios.map((ratio) => ( + ))}
@@ -363,11 +424,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..80584d95bd --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationCapabilities.test.ts @@ -0,0 +1,71 @@ +import { describe, expect, it } from "vitest"; + +import { + DEFAULT_LOBBY_CREATION_CAPABILITIES, + clampLobbySelectionToCapabilities, + parseLobbyCreationCapabilities, + validateLobbyCreationSelection, +} from "@/lib/creationCapabilities"; + +describe("creationCapabilities", () => { + it("parses backend capability payloads", () => { + expect( + parseLobbyCreationCapabilities({ + model_ids: ["fast-ltx2"], + generation_modes: ["t2va"], + resolutions: ["480p", "720p"], + duration_sec: [5, 10], + }), + ).toMatchObject({ + model_ids: ["fast-ltx2"], + generation_modes: ["t2va"], + resolutions: ["480p", "720p"], + duration_sec: [5, 10], + }); + }); + + it("clamps unsupported lobby selections to supported defaults", () => { + expect( + clampLobbySelectionToCapabilities({ + capabilities: DEFAULT_LOBBY_CREATION_CAPABILITIES, + modelId: "fast-ltx23", + modeId: "fl2av", + aspectRatio: "16:9", + resolution: "4k", + durationSec: 99, + }), + ).toEqual({ + modelId: "fast-ltx23", + modeId: "t2v", + aspectRatio: "16:9", + resolution: "480p", + durationSec: 5, + }); + }); + + it("rejects unsupported generation modes with a clear message", () => { + expect( + validateLobbyCreationSelection({ + capabilities: DEFAULT_LOBBY_CREATION_CAPABILITIES, + modelId: "fast-ltx23", + modeId: "fl2av", + aspectRatio: "16:9", + resolution: "720p", + durationSec: 5, + }), + ).toMatch(/FL2VA/i); + }); + + it("rejects unsupported resolutions", () => { + expect( + validateLobbyCreationSelection({ + capabilities: DEFAULT_LOBBY_CREATION_CAPABILITIES, + 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..3a13aae0a7 --- /dev/null +++ b/apps/dreamverse/web/src/lib/creationCapabilities.ts @@ -0,0 +1,184 @@ +import type { + AspectRatioId, + CreationModeId, + CreationModelId, + ResolutionId, +} from "@/lib/creationConfig"; +import { fromGenerationMode, toGenerationMode, type GenerationMode } from "@/lib/generationMode"; + +export interface LobbyCreationCapabilities { + model_ids: CreationModelId[]; + generation_modes: GenerationMode[]; + aspect_ratios: AspectRatioId[]; + resolutions: ResolutionId[]; + duration_sec: number[]; + unsupported_generation_modes: Record; + reference_assets: { + mime_types: string[]; + max_bytes: number; + }; +} + +export const DEFAULT_LOBBY_CREATION_CAPABILITIES: LobbyCreationCapabilities = { + model_ids: ["fast-ltx23", "fast-ltx2"], + 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 on FastLTX models yet.", + }, + reference_assets: { + mime_types: ["image/png", "image/jpeg", "image/webp"], + max_bytes: 15 * 1024 * 1024, + }, +}; + +export function parseLobbyCreationCapabilities(payload: unknown): LobbyCreationCapabilities { + if (!payload || typeof payload !== "object") { + return DEFAULT_LOBBY_CREATION_CAPABILITIES; + } + const data = payload as Record; + const 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)); + }; + return { + model_ids: pickStrings(data.model_ids, ["fast-ltx2", "fast-ltx23"], DEFAULT_LOBBY_CREATION_CAPABILITIES.model_ids), + generation_modes: pickStrings( + data.generation_modes, + ["t2va", "fl2va", "ref2va"], + DEFAULT_LOBBY_CREATION_CAPABILITIES.generation_modes, + ), + aspect_ratios: pickStrings( + data.aspect_ratios, + ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"], + DEFAULT_LOBBY_CREATION_CAPABILITIES.aspect_ratios, + ), + resolutions: pickStrings( + data.resolutions, + ["480p", "720p", "1080p", "4k"], + DEFAULT_LOBBY_CREATION_CAPABILITIES.resolutions, + ), + duration_sec: Array.isArray(data.duration_sec) + ? data.duration_sec.filter((item): item is number => typeof item === "number") + : DEFAULT_LOBBY_CREATION_CAPABILITIES.duration_sec, + unsupported_generation_modes: + typeof data.unsupported_generation_modes === "object" && data.unsupported_generation_modes + ? (data.unsupported_generation_modes as Record) + : DEFAULT_LOBBY_CREATION_CAPABILITIES.unsupported_generation_modes, + reference_assets: + typeof data.reference_assets === "object" && data.reference_assets + ? { + mime_types: Array.isArray((data.reference_assets as Record).mime_types) + ? ((data.reference_assets as Record).mime_types as string[]) + : DEFAULT_LOBBY_CREATION_CAPABILITIES.reference_assets.mime_types, + max_bytes: + typeof (data.reference_assets as Record).max_bytes === "number" + ? ((data.reference_assets as Record).max_bytes as number) + : DEFAULT_LOBBY_CREATION_CAPABILITIES.reference_assets.max_bytes, + } + : DEFAULT_LOBBY_CREATION_CAPABILITIES.reference_assets, + }; +} + +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 yet."; + } + if (!isSupportedResolution(input.resolution, input.capabilities)) { + return "Selected resolution is not supported on FastLTX models 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..14b6df7846 100644 --- a/apps/dreamverse/web/src/lib/creationConfig.test.ts +++ b/apps/dreamverse/web/src/lib/creationConfig.test.ts @@ -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..f715dfa657 100644 --- a/apps/dreamverse/web/src/lib/creationConfig.ts +++ b/apps/dreamverse/web/src/lib/creationConfig.ts @@ -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[] = [ @@ -48,11 +51,13 @@ export const CREATION_MODELS: CreationModelOption[] = [ 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 +76,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[] { From 8c5ee9b51cc3b75f4eba9cb3904fe08578c9dc9e Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Fri, 11 Sep 2026 03:33:51 +0000 Subject: [PATCH 5/5] [feat]: add fast-h3 to lobby model capabilities Derive lobby model_ids from MODEL_REGISTRY and expose per-model creation capabilities so FastH3 appears in the composer alongside FastLTX models. H3 keeps fixed 16:9 / 720p-class geometry while LTX models retain the full aspect and resolution surface. Co-authored-by: Aryan Kumar --- apps/dreamverse/arch.md | 2 +- .../dreamverse/creation_capabilities.py | 97 +++++++--- apps/dreamverse/dreamverse/mock_server.py | 4 +- apps/dreamverse/dreamverse/routes/creation.py | 4 +- .../dreamverse/session_creation_config.py | 2 +- .../tests/test_creation_capabilities.py | 40 ++++- .../tests/test_session_creation_config.py | 15 ++ apps/dreamverse/web/src/app/page.tsx | 53 ++++-- .../components/creation/CreationComposer.tsx | 5 +- .../web/src/lib/creationCapabilities.test.ts | 57 +++--- .../web/src/lib/creationCapabilities.ts | 167 +++++++++++++----- .../web/src/lib/creationConfig.test.ts | 4 +- apps/dreamverse/web/src/lib/creationConfig.ts | 7 +- .../dreamverse/web/src/lib/creationPayload.ts | 2 +- 14 files changed, 337 insertions(+), 122 deletions(-) diff --git a/apps/dreamverse/arch.md b/apps/dreamverse/arch.md index 24860c43d2..c94cc96255 100644 --- a/apps/dreamverse/arch.md +++ b/apps/dreamverse/arch.md @@ -322,7 +322,7 @@ creation config and echoes the resolved values back on `gpu_assigned` and 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` or `fast-ltx23` +- `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` diff --git a/apps/dreamverse/dreamverse/creation_capabilities.py b/apps/dreamverse/dreamverse/creation_capabilities.py index 02ce1d399d..afa8facb75 100644 --- a/apps/dreamverse/dreamverse/creation_capabilities.py +++ b/apps/dreamverse/dreamverse/creation_capabilities.py @@ -4,12 +4,10 @@ from dreamverse.config import MODEL_REGISTRY -# FastLTX lobby models share the same creation-studio surface today. -LTX_LOBBY_MODEL_IDS = frozenset({"fast-ltx2", "fast-ltx23"}) - -# Canonical upstream wire IDs. FL2VA is tracked in #1834 but not implemented on -# FastLTX streaming yet (last-frame conditioning is not wired). +# 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"}) @@ -17,29 +15,37 @@ # until the runtime path is tested on Dreamverse GPUs. LTX_LOBBY_RESOLUTIONS = frozenset({"480p", "720p", "1080p"}) -LTX_LOBBY_DURATION_SEC = frozenset({5, 10, 15}) +# 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 on FastLTX models yet.", + "fl2va": "First/last frame mode (FL2VA) is not supported yet.", } @dataclass(frozen=True) -class LobbyCreationCapabilities: - model_ids: frozenset[str] +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 { - "model_ids": sorted(self.model_ids), "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": dict(UNSUPPORTED_GENERATION_MODE_MESSAGES), + "unsupported_generation_modes": unsupported, "reference_assets": { "mime_types": ["image/png", "image/jpeg", "image/webp"], "max_bytes": 15 * 1024 * 1024, @@ -47,14 +53,63 @@ def as_dict(self) -> dict[str, object]: } -LOBBY_CREATION_CAPABILITIES = LobbyCreationCapabilities( - model_ids=LTX_LOBBY_MODEL_IDS, +LTX_MODEL_CREATION_CAPABILITIES = ModelCreationCapabilities( generation_modes=LTX_LOBBY_GENERATION_MODES, aspect_ratios=LTX_LOBBY_ASPECT_RATIOS, resolutions=LTX_LOBBY_RESOLUTIONS, - duration_sec=LTX_LOBBY_DURATION_SEC, + 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( *, @@ -64,19 +119,19 @@ def validate_lobby_creation_config( resolution: str, duration_sec: int, ) -> None: - if model_id not in LOBBY_CREATION_CAPABILITIES.model_ids: - raise ValueError(f"Unsupported model_id: {model_id}") if model_id not in MODEL_REGISTRY: raise ValueError(f"Unknown model_id: {model_id}") - if generation_mode in UNSUPPORTED_GENERATION_MODE_MESSAGES: + 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 LOBBY_CREATION_CAPABILITIES.generation_modes: + if generation_mode not in caps.generation_modes: raise ValueError(f"Unsupported generation_mode: {generation_mode}") - if aspect_ratio not in LOBBY_CREATION_CAPABILITIES.aspect_ratios: + if aspect_ratio not in caps.aspect_ratios: raise ValueError(f"Unsupported aspect_ratio: {aspect_ratio}") - if resolution not in LOBBY_CREATION_CAPABILITIES.resolutions: + if resolution not in caps.resolutions: raise ValueError(f"Unsupported resolution: {resolution}") - if duration_sec not in LOBBY_CREATION_CAPABILITIES.duration_sec: + if duration_sec not in caps.duration_sec: raise ValueError("duration_sec must be 5, 10, or 15.") diff --git a/apps/dreamverse/dreamverse/mock_server.py b/apps/dreamverse/dreamverse/mock_server.py index 329dc334cf..204e387602 100644 --- a/apps/dreamverse/dreamverse/mock_server.py +++ b/apps/dreamverse/dreamverse/mock_server.py @@ -31,7 +31,7 @@ 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_CREATION_CAPABILITIES +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 @@ -229,7 +229,7 @@ async def prompt_system_config(): @app.get("/creation-capabilities") async def creation_capabilities(): - return LOBBY_CREATION_CAPABILITIES.as_dict() + return lobby_capabilities_as_dict() @app.get("/curated-presets") diff --git a/apps/dreamverse/dreamverse/routes/creation.py b/apps/dreamverse/dreamverse/routes/creation.py index 30434b4b36..4c73d7046b 100644 --- a/apps/dreamverse/dreamverse/routes/creation.py +++ b/apps/dreamverse/dreamverse/routes/creation.py @@ -4,11 +4,11 @@ from fastapi import APIRouter -from dreamverse.creation_capabilities import LOBBY_CREATION_CAPABILITIES +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_CREATION_CAPABILITIES.as_dict() + return lobby_capabilities_as_dict() diff --git a/apps/dreamverse/dreamverse/session_creation_config.py b/apps/dreamverse/dreamverse/session_creation_config.py index dfd52b8e9a..0bf9cc0a1b 100644 --- a/apps/dreamverse/dreamverse/session_creation_config.py +++ b/apps/dreamverse/dreamverse/session_creation_config.py @@ -5,7 +5,7 @@ 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({"fast-ltx2", "fast-ltx23"}) +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"}) diff --git a/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py b/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py index 374d990d9a..b2253ea60e 100644 --- a/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py +++ b/apps/dreamverse/dreamverse/tests/test_creation_capabilities.py @@ -1,15 +1,24 @@ import pytest from dreamverse.creation_capabilities import ( - LOBBY_CREATION_CAPABILITIES, + capabilities_for_model, + lobby_capabilities_as_dict, validate_lobby_creation_config, ) -def test_lobby_capabilities_exclude_fl2va_and_4k(): - assert "fl2va" not in LOBBY_CREATION_CAPABILITIES.generation_modes - assert "4k" not in LOBBY_CREATION_CAPABILITIES.resolutions - assert LOBBY_CREATION_CAPABILITIES.generation_modes == frozenset({"t2va", "ref2va"}) +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(): @@ -22,6 +31,16 @@ def test_validate_lobby_creation_config_accepts_supported_t2va(): ) +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( @@ -42,3 +61,14 @@ def test_validate_lobby_creation_config_rejects_4k(): 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 index cba7741661..4295374c3e 100644 --- a/apps/dreamverse/dreamverse/tests/test_session_creation_config.py +++ b/apps/dreamverse/dreamverse/tests/test_session_creation_config.py @@ -44,6 +44,21 @@ 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( diff --git a/apps/dreamverse/web/src/app/page.tsx b/apps/dreamverse/web/src/app/page.tsx index 782d399f40..7172866be9 100644 --- a/apps/dreamverse/web/src/app/page.tsx +++ b/apps/dreamverse/web/src/app/page.tsx @@ -33,11 +33,12 @@ import { normalizePromptWindowSnapshot, } from "@/lib/prompts/promptWindowSnapshot"; import { - DEFAULT_LOBBY_CREATION_CAPABILITIES, + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, clampLobbySelectionToCapabilities, - parseLobbyCreationCapabilities, + parseLobbyCapabilitiesBundle, + resolveModelCapabilities, validateLobbyCreationSelection, - type LobbyCreationCapabilities, + type LobbyCapabilitiesBundle, } from "@/lib/creationCapabilities"; import { buildCreationInitPayload, @@ -262,8 +263,12 @@ export default function Page() { const [creationAspectRatio, setCreationAspectRatio] = useState("16:9"); const [creationResolution, setCreationResolution] = useState("720p"); const [creationDurationSec, setCreationDurationSec] = useState(5); - const [lobbyCapabilities, setLobbyCapabilities] = useState( - DEFAULT_LOBBY_CREATION_CAPABILITIES, + const [lobbyCapabilitiesBundle, setLobbyCapabilitiesBundle] = useState( + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, + ); + const activeModelCapabilities = useMemo( + () => resolveModelCapabilities(lobbyCapabilitiesBundle, creationModelId), + [lobbyCapabilitiesBundle, creationModelId], ); const [sessionCreationConfig, setSessionCreationConfig] = useState({ modelId: "fast-ltx23", @@ -525,10 +530,10 @@ export default function Page() { setRuntimeReady(true); }, []); - function applyLobbyCapabilities(capabilities: LobbyCreationCapabilities) { - setLobbyCapabilities(capabilities); + function applyLobbyCapabilitiesBundle(bundle: LobbyCapabilitiesBundle) { + setLobbyCapabilitiesBundle(bundle); const clamped = clampLobbySelectionToCapabilities({ - capabilities, + capabilities: resolveModelCapabilities(bundle, creationModelId), modelId: creationModelId, modeId: creationModeId, aspectRatio: creationAspectRatio, @@ -542,6 +547,22 @@ export default function Page() { 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; @@ -550,17 +571,17 @@ export default function Page() { cache: "no-store", }) .then(async (response) => { - if (!response.ok) return DEFAULT_LOBBY_CREATION_CAPABILITIES; - return parseLobbyCreationCapabilities(await response.json()); + if (!response.ok) return DEFAULT_LOBBY_CAPABILITIES_BUNDLE; + return parseLobbyCapabilitiesBundle(await response.json()); }) - .then((capabilities) => { + .then((bundle) => { if (!cancelled) { - applyLobbyCapabilities(capabilities); + applyLobbyCapabilitiesBundle(bundle); } }) .catch(() => { if (!cancelled) { - applyLobbyCapabilities(DEFAULT_LOBBY_CREATION_CAPABILITIES); + applyLobbyCapabilitiesBundle(DEFAULT_LOBBY_CAPABILITIES_BUNDLE); } }); return () => { @@ -2082,7 +2103,7 @@ export default function Page() { async function joinSession({ force = false } = {}) { const validationError = validateLobbyCreationSelection({ - capabilities: lobbyCapabilities, + capabilities: activeModelCapabilities, modelId: creationModelId, modeId: creationModeId, aspectRatio: creationAspectRatio, @@ -2882,11 +2903,11 @@ export default function Page() { lastFramePreviewUrl={lastFramePreviewUrl} mentionOptions={mentionOptions} storyPresets={lobbyStoryPresets} - capabilities={lobbyCapabilities} + capabilities={activeModelCapabilities} onValueChange={(value) => sessionStore.patch({ livePromptDraft: value })} onSubmit={() => void joinSession()} onKeyDown={handleLivePromptKeydown} - onModelChange={setCreationModelId} + onModelChange={handleCreationModelChange} onModeChange={setCreationModeId} onAspectRatioChange={setCreationAspectRatio} onResolutionChange={setCreationResolution} diff --git a/apps/dreamverse/web/src/components/creation/CreationComposer.tsx b/apps/dreamverse/web/src/components/creation/CreationComposer.tsx index 11ffeddd79..e2140cbf83 100644 --- a/apps/dreamverse/web/src/components/creation/CreationComposer.tsx +++ b/apps/dreamverse/web/src/components/creation/CreationComposer.tsx @@ -36,9 +36,10 @@ import { formatResolutionLabel, } from "@/lib/creationConfig"; import { - DEFAULT_LOBBY_CREATION_CAPABILITIES, + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, isSupportedCreationMode, isSupportedResolution, + resolveModelCapabilities, unsupportedModeNotice, type LobbyCreationCapabilities, } from "@/lib/creationCapabilities"; @@ -103,7 +104,7 @@ export default function CreationComposer({ onLastFrameSelect, onSpeechTranscript, onSpeechInterimChange, - capabilities = DEFAULT_LOBBY_CREATION_CAPABILITIES, + capabilities = resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, modelId), }: CreationComposerProps) { const inputRef = useRef(null); const [sttBusy, setSttBusy] = useState(false); diff --git a/apps/dreamverse/web/src/lib/creationCapabilities.test.ts b/apps/dreamverse/web/src/lib/creationCapabilities.test.ts index 80584d95bd..3ba1576d8a 100644 --- a/apps/dreamverse/web/src/lib/creationCapabilities.test.ts +++ b/apps/dreamverse/web/src/lib/creationCapabilities.test.ts @@ -1,44 +1,53 @@ import { describe, expect, it } from "vitest"; import { - DEFAULT_LOBBY_CREATION_CAPABILITIES, + DEFAULT_LOBBY_CAPABILITIES_BUNDLE, clampLobbySelectionToCapabilities, - parseLobbyCreationCapabilities, + parseLobbyCapabilitiesBundle, + resolveModelCapabilities, validateLobbyCreationSelection, } from "@/lib/creationCapabilities"; describe("creationCapabilities", () => { - it("parses backend capability payloads", () => { - expect( - parseLobbyCreationCapabilities({ - model_ids: ["fast-ltx2"], - generation_modes: ["t2va"], - resolutions: ["480p", "720p"], - duration_sec: [5, 10], - }), - ).toMatchObject({ - model_ids: ["fast-ltx2"], - generation_modes: ["t2va"], - resolutions: ["480p", "720p"], - duration_sec: [5, 10], + 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 supported defaults", () => { + it("clamps unsupported lobby selections to model-specific defaults", () => { expect( clampLobbySelectionToCapabilities({ - capabilities: DEFAULT_LOBBY_CREATION_CAPABILITIES, - modelId: "fast-ltx23", + capabilities: resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, "fast-h3"), + modelId: "fast-h3", modeId: "fl2av", - aspectRatio: "16:9", + aspectRatio: "9:16", resolution: "4k", durationSec: 99, }), ).toEqual({ - modelId: "fast-ltx23", + modelId: "fast-h3", modeId: "t2v", aspectRatio: "16:9", - resolution: "480p", + resolution: "720p", durationSec: 5, }); }); @@ -46,7 +55,7 @@ describe("creationCapabilities", () => { it("rejects unsupported generation modes with a clear message", () => { expect( validateLobbyCreationSelection({ - capabilities: DEFAULT_LOBBY_CREATION_CAPABILITIES, + capabilities: resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, "fast-ltx23"), modelId: "fast-ltx23", modeId: "fl2av", aspectRatio: "16:9", @@ -56,10 +65,10 @@ describe("creationCapabilities", () => { ).toMatch(/FL2VA/i); }); - it("rejects unsupported resolutions", () => { + it("rejects unsupported resolutions for ltx models", () => { expect( validateLobbyCreationSelection({ - capabilities: DEFAULT_LOBBY_CREATION_CAPABILITIES, + capabilities: resolveModelCapabilities(DEFAULT_LOBBY_CAPABILITIES_BUNDLE, "fast-ltx23"), modelId: "fast-ltx23", modeId: "t2v", aspectRatio: "16:9", diff --git a/apps/dreamverse/web/src/lib/creationCapabilities.ts b/apps/dreamverse/web/src/lib/creationCapabilities.ts index 3a13aae0a7..696c0cd317 100644 --- a/apps/dreamverse/web/src/lib/creationCapabilities.ts +++ b/apps/dreamverse/web/src/lib/creationCapabilities.ts @@ -6,8 +6,30 @@ import type { } from "@/lib/creationConfig"; import { fromGenerationMode, toGenerationMode, type GenerationMode } from "@/lib/generationMode"; -export interface LobbyCreationCapabilities { +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[]; @@ -19,14 +41,13 @@ export interface LobbyCreationCapabilities { }; } -export const DEFAULT_LOBBY_CREATION_CAPABILITIES: LobbyCreationCapabilities = { - model_ids: ["fast-ltx23", "fast-ltx2"], +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 on FastLTX models yet.", + fl2va: "First/last frame mode (FL2VA) is not supported yet.", }, reference_assets: { mime_types: ["image/png", "image/jpeg", "image/webp"], @@ -34,51 +55,109 @@ export const DEFAULT_LOBBY_CREATION_CAPABILITIES: LobbyCreationCapabilities = { }, }; -export function parseLobbyCreationCapabilities(payload: unknown): LobbyCreationCapabilities { - if (!payload || typeof payload !== "object") { - return DEFAULT_LOBBY_CREATION_CAPABILITIES; - } - const data = payload as Record; - const 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)); +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 { - model_ids: pickStrings(data.model_ids, ["fast-ltx2", "fast-ltx23"], DEFAULT_LOBBY_CREATION_CAPABILITIES.model_ids), - generation_modes: pickStrings( - data.generation_modes, - ["t2va", "fl2va", "ref2va"], - DEFAULT_LOBBY_CREATION_CAPABILITIES.generation_modes, - ), - aspect_ratios: pickStrings( - data.aspect_ratios, - ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"], - DEFAULT_LOBBY_CREATION_CAPABILITIES.aspect_ratios, - ), - resolutions: pickStrings( - data.resolutions, - ["480p", "720p", "1080p", "4k"], - DEFAULT_LOBBY_CREATION_CAPABILITIES.resolutions, - ), + 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") - : DEFAULT_LOBBY_CREATION_CAPABILITIES.duration_sec, + : fallback.duration_sec, unsupported_generation_modes: typeof data.unsupported_generation_modes === "object" && data.unsupported_generation_modes ? (data.unsupported_generation_modes as Record) - : DEFAULT_LOBBY_CREATION_CAPABILITIES.unsupported_generation_modes, - reference_assets: - typeof data.reference_assets === "object" && data.reference_assets - ? { - mime_types: Array.isArray((data.reference_assets as Record).mime_types) - ? ((data.reference_assets as Record).mime_types as string[]) - : DEFAULT_LOBBY_CREATION_CAPABILITIES.reference_assets.mime_types, - max_bytes: - typeof (data.reference_assets as Record).max_bytes === "number" - ? ((data.reference_assets as Record).max_bytes as number) - : DEFAULT_LOBBY_CREATION_CAPABILITIES.reference_assets.max_bytes, - } - : DEFAULT_LOBBY_CREATION_CAPABILITIES.reference_assets, + : 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, }; } @@ -160,10 +239,10 @@ export function validateLobbyCreationSelection(input: { return "Selected mode is not supported yet."; } if (!input.capabilities.aspect_ratios.includes(input.aspectRatio)) { - return "Selected aspect ratio is not supported yet."; + return "Selected aspect ratio is not supported for this model yet."; } if (!isSupportedResolution(input.resolution, input.capabilities)) { - return "Selected resolution is not supported on FastLTX models yet."; + 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."; diff --git a/apps/dreamverse/web/src/lib/creationConfig.test.ts b/apps/dreamverse/web/src/lib/creationConfig.test.ts index 14b6df7846..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", () => { diff --git a/apps/dreamverse/web/src/lib/creationConfig.ts b/apps/dreamverse/web/src/lib/creationConfig.ts index f715dfa657..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"; @@ -47,6 +47,11 @@ 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"]; diff --git a/apps/dreamverse/web/src/lib/creationPayload.ts b/apps/dreamverse/web/src/lib/creationPayload.ts index a03c7e9f1e..e5dbcab289 100644 --- a/apps/dreamverse/web/src/lib/creationPayload.ts +++ b/apps/dreamverse/web/src/lib/creationPayload.ts @@ -6,7 +6,7 @@ import type { } from "@/lib/creationConfig"; import { fromGenerationMode, type GenerationMode } from "@/lib/generationMode"; -const LOBBY_MODEL_IDS = new Set(["fast-ltx2", "fast-ltx23"]); +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]);