[feat] Streaming WebSocket server skeleton (single generator + fMP4) (#1251)
This commit is contained in:
@@ -0,0 +1,177 @@
|
||||
# Streaming WebSocket Server Contract
|
||||
|
||||
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
|
||||
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
|
||||
document is the authoritative spec for the message catalogue and the
|
||||
session state machine. Any change to either must update this document
|
||||
in the same PR that touches `protocol.py` or `session.py`.
|
||||
|
||||
## Endpoint
|
||||
|
||||
| Path | Protocol | Purpose |
|
||||
|---|---|---|
|
||||
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
|
||||
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
|
||||
|
||||
The server is launched by `fastvideo serve --config <serve.yaml>` when
|
||||
the config carries a `streaming:` block. Without that block the same CLI
|
||||
launches the OpenAI stateless HTTP server instead.
|
||||
|
||||
## Connection lifecycle
|
||||
|
||||
Every WebSocket connection holds exactly one `Session`. Sessions move
|
||||
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
|
||||
|
||||
```
|
||||
┌──────────────┐
|
||||
│ INITIALIZING │ ← WebSocket accepted, before init frame
|
||||
└──────┬───────┘
|
||||
│ session_init_v2 received
|
||||
┌──────────────┼──────────────┐
|
||||
▼ ▼ ▼
|
||||
QUEUED GPU_BINDING REJECTED
|
||||
│ │ ↑
|
||||
│ slot ready │ │ max-sessions hit
|
||||
▼ ▼ │ or invalid init
|
||||
┌────────┐ │
|
||||
│ ACTIVE │ ────────┘
|
||||
└────┬───┘
|
||||
segment loop │
|
||||
│
|
||||
┌───────────┼───────────┐
|
||||
▼ ▼ ▼
|
||||
COMPLETE ERROR TIMEOUT
|
||||
(clean leave) (any failure) (idle / segment_cap reached)
|
||||
```
|
||||
|
||||
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
|
||||
no transitions out. The transition matrix is enforced in
|
||||
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
|
||||
|
||||
`SessionManager` enforces the per-process budgets pulled from
|
||||
`StreamingConfig`:
|
||||
|
||||
- `session_timeout_seconds` — idle reaper drops sessions that haven't
|
||||
advanced; non-terminal sessions transition to `TIMEOUT`.
|
||||
- `generation_segment_cap` — a session that hits the cap transitions to
|
||||
`COMPLETE` after the last segment ships.
|
||||
|
||||
## Message catalogue
|
||||
|
||||
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
|
||||
`protocol.py` are the source of truth; this table is the human-readable
|
||||
view.
|
||||
|
||||
### Client → server
|
||||
|
||||
| `type` | Required fields | Purpose |
|
||||
|---|---|---|
|
||||
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
|
||||
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
|
||||
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
|
||||
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
|
||||
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
|
||||
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
|
||||
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
|
||||
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
|
||||
|
||||
The opening frame must be `session_init_v2`. Any other first frame is
|
||||
rejected with an `error` (code `invalid_message`) and the WebSocket is
|
||||
closed.
|
||||
|
||||
### Server → client
|
||||
|
||||
| `type` | Carries | When emitted |
|
||||
|---|---|---|
|
||||
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
|
||||
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
|
||||
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
|
||||
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
|
||||
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
|
||||
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
|
||||
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
|
||||
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
|
||||
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
|
||||
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
|
||||
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
|
||||
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
|
||||
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
|
||||
|
||||
## Continuation state
|
||||
|
||||
The session optionally accepts a `continuation_state` dict inside the
|
||||
opening `session_init_v2` frame. When present, the server hydrates it
|
||||
into a `ContinuationState(kind, payload)` envelope and feeds it as the
|
||||
`request.state` on the first segment's `GenerationRequest` — letting a
|
||||
client resume after a disconnect, migrate sessions across processes,
|
||||
or replay a prior session.
|
||||
|
||||
After every segment, if the runtime returns a fresh state, the server
|
||||
persists it to the `SessionStore` so a `snapshot_state` request can
|
||||
export it. The store and serialization contracts live with the model
|
||||
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
|
||||
|
||||
## Example flow
|
||||
|
||||
```
|
||||
client server
|
||||
────── ──────
|
||||
WS /v1/stream ─────── connect ─────────────────────────►
|
||||
◄────── (accept)
|
||||
|
||||
{"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox in snow", "the fox jumps"],
|
||||
"initial_image": {...},
|
||||
"stream_mode": "av_fmp4"} ─────────────────────────────►
|
||||
|
||||
(validate, queue, bind)
|
||||
◄──── {"type": "queue_status",
|
||||
"position": 0, "queue_depth": 0}
|
||||
◄──── {"type": "gpu_assigned",
|
||||
"gpu_id": 0, "model_id": "..."}
|
||||
◄──── {"type": "ltx2_stream_start", ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "a fox in snow",
|
||||
"source": "curated"} ───────────────────────────────────►
|
||||
(run pipeline)
|
||||
◄──── {"type": "ltx2_segment_start",
|
||||
"segment_idx": 1, ...}
|
||||
◄──── {"type": "step_complete",
|
||||
"segment_idx": 1, "timings": {...}}
|
||||
◄──── {"type": "media_init",
|
||||
"segment_idx": 1,
|
||||
"mime": "video/mp4", ...}
|
||||
◄──── <binary fMP4 init segment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── <binary fMP4 fragment>
|
||||
◄──── {"type": "media_segment_complete",
|
||||
"segment_idx": 1, "chunks": 12}
|
||||
◄──── {"type": "ltx2_segment_complete",
|
||||
"segment_idx": 1, ...}
|
||||
|
||||
{"type": "segment_prompt_source",
|
||||
"prompt": "the fox jumps"} ─────────────────────────────►
|
||||
(segment 2 …)
|
||||
|
||||
{"type": "snapshot_state"} ──────────────────────────────►
|
||||
◄──── {"type": "continuation_state_snapshot",
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, ...}}
|
||||
|
||||
(close) ──────────────────────────────────────────────────►
|
||||
(session → COMPLETE)
|
||||
```
|
||||
|
||||
## Backward / forward compatibility
|
||||
|
||||
- Adding a new client message: append a Pydantic model to `protocol.py`
|
||||
with a unique `type`; add the discriminator entry to `ClientMessage`;
|
||||
add a row to the table above. Old clients that don't send the new
|
||||
message remain compatible.
|
||||
- Adding a new server message: emit only when a new feature flag is
|
||||
enabled (or always emit, since clients ignore unknown types).
|
||||
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
|
||||
→ `session_init_v3`) and accept both for one release cycle. Never
|
||||
silently change field semantics under the same `type`.
|
||||
@@ -1,16 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.entrypoints.streaming.server import run_server
|
||||
from fastvideo.entrypoints.streaming.server import build_app, run_server
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
BlobStore,
|
||||
InMemoryBlobStore,
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BlobStore",
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionState",
|
||||
"SessionStore",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""JSON WebSocket protocol schemas for the streaming server.
|
||||
|
||||
Every control message shares the envelope ``{"type": <str>, ...}``.
|
||||
Pydantic models live here so the server can parse / validate incoming
|
||||
frames and emit well-typed outgoing frames without hand-rolled dicts.
|
||||
|
||||
The message catalogue matches the contract in
|
||||
``docs/design/server_contracts/streaming.md``; additions must land in
|
||||
both places in the same PR.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated, Any, Literal, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client → server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SessionInitV2(BaseModel):
|
||||
"""Opening frame the client sends after the WebSocket handshake."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
type: Literal["session_init_v2"]
|
||||
client_id: str | None = None
|
||||
preset: str | None = None
|
||||
preset_label: str | None = None
|
||||
curated_prompts: list[str] = Field(default_factory=list)
|
||||
initial_image: dict[str, Any] | None = None
|
||||
enhancement_enabled: bool = False
|
||||
auto_extension_enabled: bool = False
|
||||
loop_generation_enabled: bool = False
|
||||
single_clip_mode: bool = False
|
||||
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
|
||||
continuation_state: dict[str, Any] | None = None
|
||||
"""Optional ``{kind, payload}`` dict; hydrated into
|
||||
:class:`fastvideo.api.ContinuationState` server-side."""
|
||||
|
||||
|
||||
class SegmentPromptSource(BaseModel):
|
||||
"""Request a new segment using a specific prompt."""
|
||||
|
||||
type: Literal["segment_prompt_source"]
|
||||
prompt: str
|
||||
negative_prompt: str | None = None
|
||||
source: Literal["curated", "enhanced", "user", "auto_extension"] = "user"
|
||||
seed: int | None = None
|
||||
num_inference_steps: int | None = None
|
||||
guidance_scale: float | None = None
|
||||
|
||||
|
||||
class SeedPromptsUpdated(BaseModel):
|
||||
type: Literal["seed_prompts_updated"]
|
||||
seed_prompts: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class EnhancementUpdated(BaseModel):
|
||||
type: Literal["enhancement_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class AutoExtensionUpdated(BaseModel):
|
||||
type: Literal["auto_extension_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class LoopGenerationUpdated(BaseModel):
|
||||
type: Literal["loop_generation_updated"]
|
||||
enabled: bool
|
||||
|
||||
|
||||
class GenerationPausedUpdated(BaseModel):
|
||||
type: Literal["generation_paused_updated"]
|
||||
paused: bool
|
||||
|
||||
|
||||
class SnapshotState(BaseModel):
|
||||
"""Request the current ``ContinuationState`` for export."""
|
||||
|
||||
type: Literal["snapshot_state"]
|
||||
|
||||
|
||||
ClientMessage = Annotated[
|
||||
Union[ # noqa: UP007 - Annotated requires Union for discriminator
|
||||
SessionInitV2,
|
||||
SegmentPromptSource,
|
||||
SeedPromptsUpdated,
|
||||
EnhancementUpdated,
|
||||
AutoExtensionUpdated,
|
||||
LoopGenerationUpdated,
|
||||
GenerationPausedUpdated,
|
||||
SnapshotState,
|
||||
],
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Server → client
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class QueueStatus(BaseModel):
|
||||
type: Literal["queue_status"] = "queue_status"
|
||||
position: int
|
||||
queue_depth: int
|
||||
|
||||
|
||||
class GpuAssigned(BaseModel):
|
||||
type: Literal["gpu_assigned"] = "gpu_assigned"
|
||||
gpu_id: int
|
||||
session_timeout: int
|
||||
|
||||
|
||||
class Ltx2StreamStart(BaseModel):
|
||||
type: Literal["ltx2_stream_start"] = "ltx2_stream_start"
|
||||
preset: str | None = None
|
||||
width: int
|
||||
height: int
|
||||
fps: int
|
||||
num_frames: int
|
||||
|
||||
|
||||
class Ltx2SegmentStart(BaseModel):
|
||||
type: Literal["ltx2_segment_start"] = "ltx2_segment_start"
|
||||
segment_idx: int
|
||||
prompt: str
|
||||
total_steps: int
|
||||
|
||||
|
||||
class StepComplete(BaseModel):
|
||||
type: Literal["step_complete"] = "step_complete"
|
||||
segment_idx: int
|
||||
step: int
|
||||
total_steps: int
|
||||
stage: str = "denoise"
|
||||
|
||||
|
||||
class MediaInit(BaseModel):
|
||||
"""Descriptor for the fMP4 initialization segment that follows."""
|
||||
|
||||
type: Literal["media_init"] = "media_init"
|
||||
segment_idx: int
|
||||
mime: str = "video/mp4; codecs=\"avc1.64001f, mp4a.40.2\""
|
||||
stream_id: str
|
||||
mode: Literal["av_fmp4"] = "av_fmp4"
|
||||
|
||||
|
||||
class MediaSegmentComplete(BaseModel):
|
||||
type: Literal["media_segment_complete"] = "media_segment_complete"
|
||||
segment_idx: int
|
||||
stream_id: str
|
||||
chunks: int
|
||||
duration_ms: float | None = None
|
||||
pts_base_ms: float | None = None
|
||||
|
||||
|
||||
class Ltx2SegmentComplete(BaseModel):
|
||||
type: Literal["ltx2_segment_complete"] = "ltx2_segment_complete"
|
||||
segment_idx: int
|
||||
generation_time_ms: float
|
||||
e2e_latency_ms: float | None = None
|
||||
|
||||
|
||||
class Ltx2StreamComplete(BaseModel):
|
||||
type: Literal["ltx2_stream_complete"] = "ltx2_stream_complete"
|
||||
reason: Literal["segment_cap", "stop_requested", "error"] = "stop_requested"
|
||||
|
||||
|
||||
class SessionTimeout(BaseModel):
|
||||
type: Literal["session_timeout"] = "session_timeout"
|
||||
timeout_seconds: int
|
||||
|
||||
|
||||
class ContinuationStateSnapshot(BaseModel):
|
||||
type: Literal["continuation_state_snapshot"] = "continuation_state_snapshot"
|
||||
state: dict[str, Any]
|
||||
"""``{kind, payload}`` dict matching
|
||||
:class:`fastvideo.api.ContinuationState`."""
|
||||
|
||||
|
||||
class ErrorMessage(BaseModel):
|
||||
type: Literal["error"] = "error"
|
||||
code: Literal[
|
||||
"session_rejected",
|
||||
"invalid_message",
|
||||
"preset_mismatch",
|
||||
"gpu_unavailable",
|
||||
"worker_failed",
|
||||
"upstream_timeout",
|
||||
"internal_error",
|
||||
] = "internal_error"
|
||||
message: str
|
||||
retryable: bool = False
|
||||
|
||||
|
||||
ServerMessage = Union[ # noqa: UP007 - pydantic Union handling
|
||||
QueueStatus,
|
||||
GpuAssigned,
|
||||
Ltx2StreamStart,
|
||||
Ltx2SegmentStart,
|
||||
StepComplete,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2StreamComplete,
|
||||
SessionTimeout,
|
||||
ContinuationStateSnapshot,
|
||||
ErrorMessage,
|
||||
]
|
||||
|
||||
|
||||
def parse_client_message(raw: dict[str, Any]) -> ClientMessage:
|
||||
"""Parse an incoming WebSocket dict into a typed client message.
|
||||
|
||||
Unknown ``type`` values raise :class:`pydantic.ValidationError`; the
|
||||
server handler turns that into an ``error`` frame with
|
||||
``code="invalid_message"``.
|
||||
"""
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
return TypeAdapter(ClientMessage).validate_python(raw)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AutoExtensionUpdated",
|
||||
"ClientMessage",
|
||||
"ContinuationStateSnapshot",
|
||||
"EnhancementUpdated",
|
||||
"ErrorMessage",
|
||||
"GenerationPausedUpdated",
|
||||
"GpuAssigned",
|
||||
"Ltx2SegmentComplete",
|
||||
"Ltx2SegmentStart",
|
||||
"Ltx2StreamComplete",
|
||||
"Ltx2StreamStart",
|
||||
"LoopGenerationUpdated",
|
||||
"MediaInit",
|
||||
"MediaSegmentComplete",
|
||||
"QueueStatus",
|
||||
"SeedPromptsUpdated",
|
||||
"SegmentPromptSource",
|
||||
"ServerMessage",
|
||||
"SessionInitV2",
|
||||
"SessionTimeout",
|
||||
"SnapshotState",
|
||||
"StepComplete",
|
||||
"parse_client_message",
|
||||
]
|
||||
@@ -1,15 +1,531 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-generator FastAPI + WebSocket streaming server."""
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.api.schema import ServeConfig
|
||||
import asyncio
|
||||
import contextlib
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
AutoExtensionUpdated,
|
||||
ContinuationStateSnapshot,
|
||||
EnhancementUpdated,
|
||||
ErrorMessage,
|
||||
GenerationPausedUpdated,
|
||||
GpuAssigned,
|
||||
LoopGenerationUpdated,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamComplete,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SeedPromptsUpdated,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image, )
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# RFC 6455 WebSocket close codes used by the server.
|
||||
_WS_CLOSE_UNSUPPORTED_DATA = 1003
|
||||
_WS_CLOSE_TRY_AGAIN_LATER = 1013
|
||||
|
||||
def run_server(serve_config: ServeConfig) -> None:
|
||||
"""Launch the streaming (WebSocket / Dynamo) server."""
|
||||
|
||||
class _GeneratorProto(Protocol):
|
||||
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
|
||||
|
||||
def generate(self, request: GenerationRequest) -> Any:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerState:
|
||||
serve_config: ServeConfig
|
||||
generator: _GeneratorProto
|
||||
sessions: SessionManager
|
||||
session_store: SessionStore
|
||||
|
||||
|
||||
def build_app(
|
||||
serve_config: ServeConfig,
|
||||
generator: _GeneratorProto,
|
||||
*,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI app used by :func:`run_server`.
|
||||
|
||||
Exposed so tests can drive the WebSocket endpoint in-process via
|
||||
``starlette.testclient.TestClient(app).websocket_connect(...)``.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
|
||||
"server; got None. Add a `streaming:` block to your serve config.")
|
||||
|
||||
sessions = SessionManager(
|
||||
segment_cap=serve_config.streaming.generation_segment_cap,
|
||||
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
|
||||
)
|
||||
state = ServerState(
|
||||
serve_config=serve_config,
|
||||
generator=generator,
|
||||
sessions=sessions,
|
||||
session_store=session_store or InMemorySessionStore(),
|
||||
)
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming")
|
||||
|
||||
@app.get("/health")
|
||||
async def _health() -> JSONResponse:
|
||||
return JSONResponse({
|
||||
"status": "ok",
|
||||
"sessions": len(state.sessions),
|
||||
"stream_mode": state.serve_config.streaming.stream_mode,
|
||||
})
|
||||
|
||||
@app.websocket("/v1/stream")
|
||||
async def _stream(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
try:
|
||||
session = state.sessions.create()
|
||||
except SessionRejected as exc:
|
||||
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
|
||||
return
|
||||
|
||||
try:
|
||||
await _handle_session(websocket, session, state)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("session %s: client disconnected", session.id[:8])
|
||||
except Exception: # pragma: no cover - defensive catch-all
|
||||
logger.exception("session %s: unhandled error", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
finally:
|
||||
_cleanup_session(session, state)
|
||||
|
||||
app.state.server_state = state
|
||||
return app
|
||||
|
||||
|
||||
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
|
||||
"""Launch the streaming server.
|
||||
|
||||
Boots a :class:`fastvideo.VideoGenerator` from
|
||||
``serve_config.generator`` unless ``generator`` is provided, then
|
||||
serves ``build_app(...)`` via uvicorn.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
|
||||
"got None. Add a `streaming:` block to your serve config.")
|
||||
raise NotImplementedError("streaming server is not implemented yet")
|
||||
|
||||
import uvicorn
|
||||
|
||||
if generator is None:
|
||||
from fastvideo import VideoGenerator # lazy to avoid boot cost
|
||||
|
||||
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
|
||||
app = build_app(serve_config, generator)
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=serve_config.server.host,
|
||||
port=serve_config.server.port,
|
||||
)
|
||||
|
||||
|
||||
async def _handle_session(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
init = await _read_init_message(websocket, session, state)
|
||||
if init is None:
|
||||
return
|
||||
|
||||
await _apply_session_init(session, init, state)
|
||||
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
|
||||
session.transition(SessionState.GPU_BINDING)
|
||||
await _send_json(websocket, GpuAssigned(
|
||||
gpu_id=0,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
session.transition(SessionState.ACTIVE)
|
||||
await _send_json(websocket, _build_stream_start(session, state))
|
||||
|
||||
try:
|
||||
await _run_segment_loop(websocket, session, state)
|
||||
finally:
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
|
||||
|
||||
|
||||
async def _read_init_message(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> SessionInitV2 | None:
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: init timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return None
|
||||
except WebSocketDisconnect:
|
||||
return None
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
|
||||
return None
|
||||
if not isinstance(parsed, SessionInitV2):
|
||||
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
|
||||
return None
|
||||
return parsed
|
||||
|
||||
|
||||
async def _reject_init(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
message: str,
|
||||
close_reason: str,
|
||||
) -> None:
|
||||
await _send_error(websocket, "invalid_message", message, retryable=False)
|
||||
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.REJECTED)
|
||||
|
||||
|
||||
async def _apply_session_init(
|
||||
session: Session,
|
||||
init: SessionInitV2,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
session.client_id = init.client_id
|
||||
session.preset = init.preset
|
||||
session.preset_label = init.preset_label
|
||||
session.curated_prompts = list(init.curated_prompts)
|
||||
session.enhancement_enabled = init.enhancement_enabled
|
||||
session.auto_extension_enabled = init.auto_extension_enabled
|
||||
session.loop_generation_enabled = init.loop_generation_enabled
|
||||
session.single_clip_mode = init.single_clip_mode
|
||||
session.stream_mode = init.stream_mode
|
||||
|
||||
if init.initial_image is not None:
|
||||
# Decode + disk write off the event loop; payload is up to 32 MiB.
|
||||
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
|
||||
if image is not None:
|
||||
session.metadata["session_init_image"] = image.path
|
||||
|
||||
if init.continuation_state is not None:
|
||||
session.continuation_state = _coerce_state(init.continuation_state)
|
||||
if session.continuation_state is not None:
|
||||
state.session_store.store(session.id, session.continuation_state)
|
||||
|
||||
|
||||
async def _run_segment_loop(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> None:
|
||||
cap = state.sessions.segment_cap
|
||||
while True:
|
||||
if session.segment_cap_reached(cap):
|
||||
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
|
||||
return
|
||||
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_json(),
|
||||
timeout=state.sessions.session_timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.info("session %s: idle timeout", session.id[:8])
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return
|
||||
except WebSocketDisconnect:
|
||||
return
|
||||
session.touch()
|
||||
|
||||
try:
|
||||
parsed = parse_client_message(raw)
|
||||
except Exception as exc:
|
||||
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SnapshotState):
|
||||
snap = state.session_store.snapshot(session.id)
|
||||
if snap is None:
|
||||
await _send_error(websocket,
|
||||
"internal_error",
|
||||
"no continuation state available for session",
|
||||
retryable=False)
|
||||
continue
|
||||
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
|
||||
continue
|
||||
|
||||
if isinstance(parsed, SegmentPromptSource):
|
||||
await _run_segment(websocket, session, state, parsed)
|
||||
continue
|
||||
|
||||
# Silently ignore unknown-but-valid types (additive-evolution
|
||||
# rule in streaming.md).
|
||||
_apply_toggle(session, parsed)
|
||||
|
||||
|
||||
async def _run_segment(
|
||||
websocket: WebSocket,
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
message: SegmentPromptSource,
|
||||
) -> None:
|
||||
request = _build_generation_request(session, message, state)
|
||||
segment_idx = session.segment_idx
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentStart(
|
||||
segment_idx=segment_idx,
|
||||
prompt=message.prompt,
|
||||
total_steps=request.sampling.num_inference_steps,
|
||||
))
|
||||
|
||||
start = time.perf_counter()
|
||||
loop = asyncio.get_running_loop()
|
||||
# TODO: executor-wrapped generate() cannot be cancelled, so a
|
||||
# client disconnect mid-segment leaves the GPU work running to
|
||||
# completion. Real cancellation needs the generate_async API.
|
||||
try:
|
||||
result = await loop.run_in_executor(None, state.generator.generate, request)
|
||||
except Exception as exc:
|
||||
logger.exception("session %s: generator failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
elapsed_ms = (time.perf_counter() - start) * 1000.0
|
||||
|
||||
frames = _extract_frames(result)
|
||||
if not frames:
|
||||
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
|
||||
# Synchronous generator call has no per-step hook; emit one
|
||||
# terminal StepComplete so observability wiring still sees the
|
||||
# segment finish.
|
||||
total = request.sampling.num_inference_steps
|
||||
await _send_json(websocket, StepComplete(
|
||||
segment_idx=segment_idx,
|
||||
step=total,
|
||||
total_steps=total,
|
||||
stage="denoise",
|
||||
))
|
||||
|
||||
encoder = FragmentedMP4Encoder(
|
||||
width=request.sampling.width,
|
||||
height=request.sampling.height,
|
||||
fps=request.sampling.fps,
|
||||
segment_idx=segment_idx,
|
||||
)
|
||||
chunks_relayed = 0
|
||||
async with encoder:
|
||||
init_sent = False
|
||||
async for chunk in encoder.encode(frames):
|
||||
if chunk.kind == "init":
|
||||
await _send_json(websocket, MediaInit(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=chunk.stream_id,
|
||||
))
|
||||
init_sent = True
|
||||
await websocket.send_bytes(chunk.data)
|
||||
if init_sent and chunk.kind == "media":
|
||||
chunks_relayed += 1
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
MediaSegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
stream_id=encoder.stream_id,
|
||||
chunks=chunks_relayed,
|
||||
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
|
||||
))
|
||||
|
||||
new_state = _extract_state(result)
|
||||
if new_state is not None:
|
||||
session.continuation_state = new_state
|
||||
state.session_store.store(session.id, new_state)
|
||||
|
||||
session.segment_idx += 1
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ACTIVE)
|
||||
|
||||
await _send_json(
|
||||
websocket,
|
||||
Ltx2SegmentComplete(
|
||||
segment_idx=segment_idx,
|
||||
generation_time_ms=elapsed_ms,
|
||||
e2e_latency_ms=elapsed_ms,
|
||||
))
|
||||
|
||||
|
||||
def _build_stream_start(
|
||||
session: Session,
|
||||
state: ServerState,
|
||||
) -> Ltx2StreamStart:
|
||||
default = state.serve_config.default_request
|
||||
return Ltx2StreamStart(
|
||||
preset=session.preset,
|
||||
width=default.sampling.width,
|
||||
height=default.sampling.height,
|
||||
fps=default.sampling.fps,
|
||||
num_frames=default.sampling.num_frames,
|
||||
)
|
||||
|
||||
|
||||
def _build_generation_request(
|
||||
session: Session,
|
||||
message: SegmentPromptSource,
|
||||
state: ServerState,
|
||||
) -> GenerationRequest:
|
||||
# Start from the operator-pinned default_request to pick up the
|
||||
# preset-selected sampling knobs; override with per-message values.
|
||||
base = state.serve_config.default_request
|
||||
sampling_kwargs: dict[str, Any] = {
|
||||
"num_videos_per_prompt":
|
||||
base.sampling.num_videos_per_prompt,
|
||||
"seed":
|
||||
message.seed if message.seed is not None else base.sampling.seed,
|
||||
"num_frames":
|
||||
base.sampling.num_frames,
|
||||
"height":
|
||||
base.sampling.height,
|
||||
"width":
|
||||
base.sampling.width,
|
||||
"fps":
|
||||
base.sampling.fps,
|
||||
"num_inference_steps":
|
||||
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
|
||||
"guidance_scale":
|
||||
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
|
||||
}
|
||||
request = GenerationRequest(
|
||||
prompt=message.prompt,
|
||||
negative_prompt=message.negative_prompt or base.negative_prompt,
|
||||
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
|
||||
sampling=SamplingConfig(**sampling_kwargs),
|
||||
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
|
||||
state=session.continuation_state,
|
||||
)
|
||||
return request
|
||||
|
||||
|
||||
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
|
||||
kind = raw.get("kind")
|
||||
payload = raw.get("payload")
|
||||
if not isinstance(kind, str) or not isinstance(payload, dict):
|
||||
return None
|
||||
return ContinuationState(kind=kind, payload=payload)
|
||||
|
||||
|
||||
def _apply_toggle(session: Session, message: Any) -> None:
|
||||
if isinstance(message, EnhancementUpdated):
|
||||
session.enhancement_enabled = message.enabled
|
||||
elif isinstance(message, AutoExtensionUpdated):
|
||||
session.auto_extension_enabled = message.enabled
|
||||
elif isinstance(message, LoopGenerationUpdated):
|
||||
session.loop_generation_enabled = message.enabled
|
||||
elif isinstance(message, GenerationPausedUpdated):
|
||||
session.generation_paused = message.paused
|
||||
elif isinstance(message, SeedPromptsUpdated):
|
||||
session.curated_prompts = list(message.seed_prompts)
|
||||
|
||||
|
||||
def _extract_frames(result: Any) -> list:
|
||||
if hasattr(result, "frames"):
|
||||
return list(result.frames or [])
|
||||
if isinstance(result, dict):
|
||||
return list(result.get("frames") or [])
|
||||
return []
|
||||
|
||||
|
||||
def _extract_state(result: Any) -> ContinuationState | None:
|
||||
state = getattr(result, "state", None)
|
||||
if state is None and isinstance(result, dict):
|
||||
state = result.get("state")
|
||||
if isinstance(state, ContinuationState):
|
||||
return state
|
||||
if isinstance(state, dict):
|
||||
return _coerce_state(state)
|
||||
return None
|
||||
|
||||
|
||||
async def _send_json(websocket: WebSocket, message: Any) -> None:
|
||||
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
|
||||
await websocket.send_json(payload)
|
||||
|
||||
|
||||
async def _send_error(
|
||||
websocket: WebSocket,
|
||||
code: str,
|
||||
message: str,
|
||||
*,
|
||||
retryable: bool,
|
||||
) -> None:
|
||||
await _send_json(
|
||||
websocket,
|
||||
ErrorMessage(code=code, message=message, retryable=retryable),
|
||||
)
|
||||
|
||||
|
||||
def _cleanup_session(session: Session, state: ServerState) -> None:
|
||||
state.sessions.close(session.id)
|
||||
state.session_store.drop(session.id)
|
||||
init_image_path = session.metadata.get("session_init_image")
|
||||
if isinstance(init_image_path, str):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(init_image_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ServerState",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-connection session lifecycle for the streaming server.
|
||||
|
||||
Each WebSocket opens exactly one :class:`Session`. :class:`SessionManager`
|
||||
enforces the ``generation_segment_cap`` and ``session_timeout_seconds``
|
||||
budgets from :class:`fastvideo.api.StreamingConfig`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.schema import ContinuationState
|
||||
|
||||
|
||||
class SessionState(enum.Enum):
|
||||
"""State-machine positions for a streaming session.
|
||||
|
||||
Transitions are server-owned. See
|
||||
``docs/design/server_contracts/streaming.md`` for the full diagram.
|
||||
"""
|
||||
|
||||
INITIALIZING = "initializing"
|
||||
QUEUED = "queued"
|
||||
GPU_BINDING = "gpu_binding"
|
||||
ACTIVE = "active"
|
||||
COMPLETE = "complete"
|
||||
ERROR = "error"
|
||||
TIMEOUT = "timeout"
|
||||
REJECTED = "rejected"
|
||||
|
||||
|
||||
_VALID_TRANSITIONS: dict[SessionState, frozenset[SessionState]] = {
|
||||
SessionState.INITIALIZING:
|
||||
frozenset({
|
||||
SessionState.QUEUED,
|
||||
SessionState.GPU_BINDING,
|
||||
SessionState.REJECTED,
|
||||
SessionState.ERROR,
|
||||
}),
|
||||
SessionState.QUEUED:
|
||||
frozenset({
|
||||
SessionState.GPU_BINDING,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
SessionState.REJECTED,
|
||||
}),
|
||||
SessionState.GPU_BINDING:
|
||||
frozenset({
|
||||
SessionState.ACTIVE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
}),
|
||||
SessionState.ACTIVE:
|
||||
frozenset({
|
||||
SessionState.ACTIVE,
|
||||
SessionState.COMPLETE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
}),
|
||||
SessionState.COMPLETE:
|
||||
frozenset(),
|
||||
SessionState.ERROR:
|
||||
frozenset(),
|
||||
SessionState.TIMEOUT:
|
||||
frozenset(),
|
||||
SessionState.REJECTED:
|
||||
frozenset(),
|
||||
}
|
||||
|
||||
|
||||
class InvalidSessionTransition(RuntimeError):
|
||||
"""Raised when a session is asked to transition along an illegal edge."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Session:
|
||||
id: str = field(default_factory=lambda: uuid.uuid4().hex)
|
||||
state: SessionState = SessionState.INITIALIZING
|
||||
created_at: float = field(default_factory=time.monotonic)
|
||||
last_activity: float = field(default_factory=time.monotonic)
|
||||
|
||||
client_id: str | None = None
|
||||
preset: str | None = None
|
||||
preset_label: str | None = None
|
||||
|
||||
curated_prompts: list[str] = field(default_factory=list)
|
||||
|
||||
segment_idx: int = 0
|
||||
|
||||
enhancement_enabled: bool = False
|
||||
auto_extension_enabled: bool = False
|
||||
loop_generation_enabled: bool = False
|
||||
single_clip_mode: bool = False
|
||||
generation_paused: bool = False
|
||||
|
||||
stream_mode: str = "av_fmp4"
|
||||
gpu_id: int | None = None
|
||||
|
||||
continuation_state: ContinuationState | None = None
|
||||
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def transition(self, target: SessionState) -> None:
|
||||
"""Move to ``target`` if the edge is allowed.
|
||||
|
||||
Raises :class:`InvalidSessionTransition` on illegal moves. The
|
||||
self-loop on ``ACTIVE`` is legal so the server can re-assert
|
||||
ACTIVE on segment completion without special casing.
|
||||
"""
|
||||
allowed = _VALID_TRANSITIONS.get(self.state, frozenset())
|
||||
if target not in allowed and target is not self.state:
|
||||
raise InvalidSessionTransition(f"{self.state.value} -> {target.value} is not a valid "
|
||||
f"session transition")
|
||||
self.state = target
|
||||
self.last_activity = time.monotonic()
|
||||
|
||||
def touch(self) -> None:
|
||||
self.last_activity = time.monotonic()
|
||||
|
||||
def is_active(self) -> bool:
|
||||
return self.state is SessionState.ACTIVE
|
||||
|
||||
def segment_cap_reached(self, cap: int) -> bool:
|
||||
return self.segment_idx >= cap
|
||||
|
||||
|
||||
class SessionManager:
|
||||
"""Registers sessions and enforces per-server session limits."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
segment_cap: int,
|
||||
session_timeout_seconds: int,
|
||||
max_sessions: int = 1,
|
||||
) -> None:
|
||||
self._segment_cap = segment_cap
|
||||
self._session_timeout_seconds = session_timeout_seconds
|
||||
self._max_sessions = max_sessions
|
||||
self._sessions: dict[str, Session] = {}
|
||||
|
||||
@property
|
||||
def segment_cap(self) -> int:
|
||||
return self._segment_cap
|
||||
|
||||
@property
|
||||
def session_timeout_seconds(self) -> int:
|
||||
return self._session_timeout_seconds
|
||||
|
||||
def create(self) -> Session:
|
||||
if len(self._sessions) >= self._max_sessions:
|
||||
raise SessionRejected(f"max sessions reached ({self._max_sessions})")
|
||||
session = Session()
|
||||
self._sessions[session.id] = session
|
||||
return session
|
||||
|
||||
def get(self, session_id: str) -> Session | None:
|
||||
return self._sessions.get(session_id)
|
||||
|
||||
def close(self, session_id: str) -> None:
|
||||
self._sessions.pop(session_id, None)
|
||||
|
||||
def __contains__(self, session_id: str) -> bool:
|
||||
return session_id in self._sessions
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._sessions)
|
||||
|
||||
def active_sessions(self) -> list[Session]:
|
||||
return [s for s in self._sessions.values() if s.is_active()]
|
||||
|
||||
def reap_timed_out(self, now: float | None = None) -> list[str]:
|
||||
"""Return the ids of sessions that have exceeded the idle timeout.
|
||||
|
||||
The caller is responsible for actually closing them — this
|
||||
method only *identifies* dead sessions so the server can emit
|
||||
``session_timeout`` frames before dropping the WebSocket.
|
||||
|
||||
TODO: unused until a background driver calls it. Per-connection
|
||||
idle enforcement currently happens via asyncio.wait_for on
|
||||
receive_json; this helper catches sessions stuck before any
|
||||
receive (e.g. future QUEUED state) and is expected to be wired
|
||||
into the GPU-pool reaper.
|
||||
"""
|
||||
now = now if now is not None else time.monotonic()
|
||||
dead: list[str] = []
|
||||
for sid, session in self._sessions.items():
|
||||
if session.state in {
|
||||
SessionState.COMPLETE,
|
||||
SessionState.ERROR,
|
||||
SessionState.TIMEOUT,
|
||||
SessionState.REJECTED,
|
||||
}:
|
||||
continue
|
||||
if now - session.last_activity > self._session_timeout_seconds:
|
||||
dead.append(sid)
|
||||
return dead
|
||||
|
||||
|
||||
class SessionRejected(RuntimeError):
|
||||
"""Raised when session creation fails (queue full, auth, etc.)."""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InvalidSessionTransition",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionRejected",
|
||||
"SessionState",
|
||||
]
|
||||
@@ -0,0 +1,103 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Persist the initial-image blob attached to a streaming session."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import contextlib
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
_ACCEPTED_MIMES = {
|
||||
"image/png": ".png",
|
||||
"image/jpeg": ".jpg",
|
||||
"image/jpg": ".jpg",
|
||||
"image/webp": ".webp",
|
||||
}
|
||||
|
||||
_MAX_IMAGE_BYTES = 32 * 1024 * 1024 # 32 MiB cap
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SessionInitImage:
|
||||
"""Location of the persisted init image.
|
||||
|
||||
Callers pass ``path`` to ``InputConfig.image_path``; ``display_name``
|
||||
is only used for logs.
|
||||
"""
|
||||
|
||||
path: str
|
||||
display_name: str
|
||||
mime: str
|
||||
|
||||
|
||||
def persist_session_init_image(
|
||||
payload: Any,
|
||||
*,
|
||||
output_dir: str | None = None,
|
||||
) -> SessionInitImage | None:
|
||||
"""Decode a client init-image blob and persist it to disk.
|
||||
|
||||
``payload`` shape (matches the internal UI protocol)::
|
||||
|
||||
{
|
||||
"mime": "image/png",
|
||||
"name": "ref.png",
|
||||
"data": "<base64 bytes>",
|
||||
}
|
||||
|
||||
Returns ``None`` when ``payload`` is falsy (no init image). Raises
|
||||
:class:`ValueError` on schema / size / decode errors so the caller
|
||||
can surface a user-facing ``error`` frame.
|
||||
"""
|
||||
if not payload:
|
||||
return None
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("session init image must be an object")
|
||||
|
||||
mime = payload.get("mime")
|
||||
if mime not in _ACCEPTED_MIMES:
|
||||
raise ValueError(f"session init image mime {mime!r} is not one of "
|
||||
f"{sorted(_ACCEPTED_MIMES)}")
|
||||
data_b64 = payload.get("data")
|
||||
if not isinstance(data_b64, str):
|
||||
raise ValueError("session init image data must be a base64 string")
|
||||
try:
|
||||
data = base64.b64decode(data_b64, validate=True)
|
||||
except (binascii.Error, ValueError) as exc:
|
||||
raise ValueError(f"session init image data is not valid base64: {exc}") from exc
|
||||
if len(data) > _MAX_IMAGE_BYTES:
|
||||
raise ValueError(f"session init image is {len(data)} bytes; limit is "
|
||||
f"{_MAX_IMAGE_BYTES}")
|
||||
if len(data) == 0:
|
||||
raise ValueError("session init image data is empty")
|
||||
|
||||
ext = _ACCEPTED_MIMES[mime]
|
||||
display_name = _sanitize_display_name(payload.get("name")) or f"init{ext}"
|
||||
fd, path = tempfile.mkstemp(prefix="fastvideo-init-", suffix=ext, dir=output_dir)
|
||||
try:
|
||||
with os.fdopen(fd, "wb") as f:
|
||||
f.write(data)
|
||||
except Exception:
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(path)
|
||||
raise
|
||||
return SessionInitImage(path=path, display_name=display_name, mime=mime)
|
||||
|
||||
|
||||
def _sanitize_display_name(name: Any) -> str | None:
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
name = name.strip()
|
||||
if not name:
|
||||
return None
|
||||
# Strip any path components — we only keep the leaf for logging.
|
||||
return os.path.basename(name)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"SessionInitImage",
|
||||
"persist_session_init_image",
|
||||
]
|
||||
@@ -0,0 +1,213 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""fMP4 stream encoder used by the streaming server.
|
||||
|
||||
The client's Media Source Extensions player needs a continuous fMP4
|
||||
byte stream: first an *initialization segment* (``ftyp`` + ``moov``),
|
||||
then one or more *media segments* (``moof`` + ``mdat``). We pipe raw
|
||||
RGB frames into an ffmpeg subprocess configured for fragmented output
|
||||
via ``-movflags empty_moov+default_base_moof+frag_keyframe+faststart``
|
||||
and stream the bytes back out.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import subprocess
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import numpy as np
|
||||
|
||||
|
||||
@dataclass
|
||||
class FragmentedMP4Chunk:
|
||||
"""A single fMP4 byte chunk emitted by :class:`FragmentedMP4Encoder`.
|
||||
|
||||
``kind`` identifies whether the chunk is the init segment (must be
|
||||
fed into the client's ``SourceBuffer`` first) or a media fragment.
|
||||
"""
|
||||
|
||||
kind: Literal["init", "media"]
|
||||
data: bytes
|
||||
stream_id: str
|
||||
segment_idx: int
|
||||
|
||||
|
||||
class FragmentedMP4Encoder:
|
||||
"""Stream RGB frames in, fMP4 chunks out.
|
||||
|
||||
One encoder covers one segment. The server creates a new encoder
|
||||
per :class:`ltx2_segment_start`` boundary so each segment becomes
|
||||
one media fragment the client can append independently.
|
||||
|
||||
Example::
|
||||
|
||||
encoder = FragmentedMP4Encoder(width=1024, height=576, fps=24,
|
||||
segment_idx=0)
|
||||
async with encoder:
|
||||
async for chunk in encoder.encode(frames):
|
||||
await websocket.send_bytes(chunk.data)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
width: int,
|
||||
height: int,
|
||||
fps: int,
|
||||
segment_idx: int,
|
||||
stream_id: str | None = None,
|
||||
ffmpeg_path: str = "ffmpeg",
|
||||
preset: str = "ultrafast",
|
||||
pixel_format_out: str = "yuv420p",
|
||||
extra_args: list[str] | None = None,
|
||||
) -> None:
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.fps = fps
|
||||
self.segment_idx = segment_idx
|
||||
self.stream_id = stream_id or uuid.uuid4().hex
|
||||
self._ffmpeg_path = ffmpeg_path
|
||||
self._preset = preset
|
||||
self._pixel_format_out = pixel_format_out
|
||||
self._extra_args = list(extra_args or [])
|
||||
self._proc: subprocess.Popen | None = None
|
||||
self._init_emitted = False
|
||||
|
||||
async def __aenter__(self) -> FragmentedMP4Encoder:
|
||||
self._spawn()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None:
|
||||
await self.close()
|
||||
|
||||
def _spawn(self) -> None:
|
||||
args = [
|
||||
self._ffmpeg_path,
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"error",
|
||||
"-f",
|
||||
"rawvideo",
|
||||
"-pix_fmt",
|
||||
"rgb24",
|
||||
"-s",
|
||||
f"{self.width}x{self.height}",
|
||||
"-r",
|
||||
str(self.fps),
|
||||
"-i",
|
||||
"-",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
self._preset,
|
||||
"-tune",
|
||||
"zerolatency",
|
||||
"-pix_fmt",
|
||||
self._pixel_format_out,
|
||||
"-movflags",
|
||||
"empty_moov+default_base_moof+frag_keyframe+faststart",
|
||||
"-f",
|
||||
"mp4",
|
||||
*self._extra_args,
|
||||
"-",
|
||||
]
|
||||
# stderr → DEVNULL: with -loglevel error on, the only thing
|
||||
# stderr would carry is unsolicited warnings. Piping without a
|
||||
# reader deadlocks ffmpeg once the pipe buffer (~64 KiB) fills.
|
||||
self._proc = subprocess.Popen( # noqa: S603
|
||||
args,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
bufsize=0,
|
||||
)
|
||||
|
||||
async def encode(
|
||||
self,
|
||||
frames: list[np.ndarray] | AsyncIterator[np.ndarray],
|
||||
) -> AsyncIterator[FragmentedMP4Chunk]:
|
||||
"""Feed frames into ffmpeg and yield fMP4 chunks as they appear."""
|
||||
if self._proc is None:
|
||||
self._spawn()
|
||||
assert self._proc is not None and self._proc.stdin is not None
|
||||
proc = self._proc
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
async def _writer() -> None:
|
||||
try:
|
||||
if hasattr(frames, "__aiter__"):
|
||||
async for frame in frames: # type: ignore[union-attr]
|
||||
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
|
||||
else:
|
||||
for frame in frames: # type: ignore[assignment]
|
||||
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
|
||||
finally:
|
||||
with contextlib.suppress(BrokenPipeError):
|
||||
proc.stdin.close()
|
||||
|
||||
writer_task = asyncio.create_task(_writer())
|
||||
try:
|
||||
reader = proc.stdout
|
||||
assert reader is not None
|
||||
# Read in reasonably-sized chunks; MSE tolerates any size
|
||||
# but we don't want to starve the event loop.
|
||||
chunk_size = 64 * 1024
|
||||
while True:
|
||||
data = await loop.run_in_executor(None, reader.read, chunk_size)
|
||||
if not data:
|
||||
break
|
||||
kind: Literal["init", "media"] = "init" if not self._init_emitted else "media"
|
||||
self._init_emitted = True
|
||||
yield FragmentedMP4Chunk(
|
||||
kind=kind,
|
||||
data=bytes(data),
|
||||
stream_id=self.stream_id,
|
||||
segment_idx=self.segment_idx,
|
||||
)
|
||||
finally:
|
||||
await writer_task
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._proc is None:
|
||||
return
|
||||
proc = self._proc
|
||||
self._proc = None
|
||||
try:
|
||||
if proc.stdin and not proc.stdin.closed:
|
||||
proc.stdin.close()
|
||||
except BrokenPipeError:
|
||||
pass
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
loop.run_in_executor(None, proc.wait),
|
||||
timeout=5.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
proc.kill()
|
||||
await loop.run_in_executor(None, proc.wait)
|
||||
|
||||
|
||||
def _write_frame(stdin, frame: np.ndarray) -> None:
|
||||
import numpy as np
|
||||
|
||||
if not isinstance(frame, np.ndarray):
|
||||
raise TypeError("fMP4 encoder frames must be numpy.ndarray")
|
||||
if frame.dtype != np.uint8:
|
||||
frame = frame.astype(np.uint8)
|
||||
if frame.ndim != 3 or frame.shape[-1] != 3:
|
||||
raise ValueError("fMP4 encoder frames must be HxWx3 uint8 RGB; got "
|
||||
f"shape={frame.shape}, dtype={frame.dtype}")
|
||||
with contextlib.suppress(BrokenPipeError):
|
||||
stdin.write(frame.tobytes())
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
]
|
||||
@@ -519,7 +519,7 @@ def test_main_rejects_top_level_config_without_subcommand(tmp_path, monkeypatch)
|
||||
cli_main.main()
|
||||
|
||||
|
||||
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
|
||||
def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path, monkeypatch):
|
||||
config_path = tmp_path / "serve-streaming.yaml"
|
||||
config_path.write_text(
|
||||
"generator:\n"
|
||||
@@ -530,9 +530,21 @@ def test_serve_cmd_dispatches_to_streaming_when_streaming_block_set(tmp_path):
|
||||
)
|
||||
args, _ = _parse_serve_args(["--config", str(config_path)])
|
||||
|
||||
with pytest.raises(NotImplementedError,
|
||||
match="streaming server is not implemented"):
|
||||
ServeSubcommand().cmd(args)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_run_server(serve_config, *, generator=None):
|
||||
captured["serve_config"] = serve_config
|
||||
|
||||
def fail_if_called(*_args, **_kwargs):
|
||||
raise AssertionError("OpenAI server must not run when streaming is set")
|
||||
|
||||
monkeypatch.setattr(streaming_server, "run_server", fake_run_server)
|
||||
monkeypatch.setattr(api_server, "run_server", fail_if_called)
|
||||
ServeSubcommand().cmd(args)
|
||||
|
||||
serve_config = captured["serve_config"]
|
||||
assert serve_config.streaming is not None
|
||||
assert serve_config.streaming.stream_mode == "av_fmp4"
|
||||
|
||||
|
||||
def test_streaming_run_server_rejects_missing_streaming_block():
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Protocol schema tests for the streaming server.
|
||||
|
||||
Covers:
|
||||
|
||||
* accepted client messages parse into the correct discriminated model
|
||||
* unknown ``type`` values raise validation errors
|
||||
* server-side messages serialize to the expected wire shape
|
||||
* continuation_state field on session_init_v2 carries through
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from fastvideo.entrypoints.streaming.protocol import (
|
||||
ContinuationStateSnapshot,
|
||||
ErrorMessage,
|
||||
GpuAssigned,
|
||||
Ltx2SegmentComplete,
|
||||
Ltx2SegmentStart,
|
||||
Ltx2StreamStart,
|
||||
MediaInit,
|
||||
MediaSegmentComplete,
|
||||
QueueStatus,
|
||||
SegmentPromptSource,
|
||||
SessionInitV2,
|
||||
SnapshotState,
|
||||
StepComplete,
|
||||
parse_client_message,
|
||||
)
|
||||
|
||||
|
||||
class TestClientMessageParsing:
|
||||
|
||||
def test_session_init_v2_minimal(self):
|
||||
parsed = parse_client_message({"type": "session_init_v2"})
|
||||
assert isinstance(parsed, SessionInitV2)
|
||||
assert parsed.curated_prompts == []
|
||||
assert parsed.stream_mode == "av_fmp4"
|
||||
|
||||
def test_session_init_v2_full(self):
|
||||
raw = {
|
||||
"type": "session_init_v2",
|
||||
"client_id": "client-1",
|
||||
"preset": "ltx2_two_stage",
|
||||
"preset_label": "2x refine",
|
||||
"curated_prompts": ["a fox", "a deer"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
"single_clip_mode": True,
|
||||
"stream_mode": "av_fmp4",
|
||||
"continuation_state": {
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, "segment_index": 2},
|
||||
},
|
||||
}
|
||||
parsed = parse_client_message(raw)
|
||||
assert isinstance(parsed, SessionInitV2)
|
||||
assert parsed.preset == "ltx2_two_stage"
|
||||
assert parsed.curated_prompts == ["a fox", "a deer"]
|
||||
assert parsed.continuation_state["kind"] == "ltx2.v1"
|
||||
|
||||
def test_segment_prompt_source(self):
|
||||
parsed = parse_client_message({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "hello world",
|
||||
"source": "curated",
|
||||
"seed": 7,
|
||||
})
|
||||
assert isinstance(parsed, SegmentPromptSource)
|
||||
assert parsed.source == "curated"
|
||||
assert parsed.seed == 7
|
||||
|
||||
def test_snapshot_state(self):
|
||||
parsed = parse_client_message({"type": "snapshot_state"})
|
||||
assert isinstance(parsed, SnapshotState)
|
||||
|
||||
def test_unknown_type_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"type": "not_a_real_message"})
|
||||
|
||||
def test_missing_type_rejected(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"prompt": "x"})
|
||||
|
||||
def test_segment_prompt_source_requires_prompt(self):
|
||||
with pytest.raises(ValidationError):
|
||||
parse_client_message({"type": "segment_prompt_source"})
|
||||
|
||||
|
||||
class TestServerMessageSerialization:
|
||||
|
||||
def test_queue_status(self):
|
||||
msg = QueueStatus(position=3, queue_depth=5)
|
||||
assert msg.model_dump() == {
|
||||
"type": "queue_status",
|
||||
"position": 3,
|
||||
"queue_depth": 5,
|
||||
}
|
||||
|
||||
def test_gpu_assigned(self):
|
||||
msg = GpuAssigned(gpu_id=1, session_timeout=300)
|
||||
assert msg.model_dump()["type"] == "gpu_assigned"
|
||||
|
||||
def test_ltx2_stream_start(self):
|
||||
msg = Ltx2StreamStart(
|
||||
preset="ltx2_two_stage",
|
||||
width=1024, height=1536, fps=24, num_frames=121,
|
||||
)
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["type"] == "ltx2_stream_start"
|
||||
assert dumped["width"] == 1024
|
||||
|
||||
def test_ltx2_segment_start(self):
|
||||
msg = Ltx2SegmentStart(
|
||||
segment_idx=0,
|
||||
prompt="a fox",
|
||||
total_steps=8,
|
||||
)
|
||||
assert msg.model_dump()["segment_idx"] == 0
|
||||
|
||||
def test_step_complete(self):
|
||||
msg = StepComplete(segment_idx=0, step=1, total_steps=8)
|
||||
assert msg.model_dump()["stage"] == "denoise"
|
||||
|
||||
def test_media_init_has_mode(self):
|
||||
msg = MediaInit(segment_idx=0, stream_id="abc")
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["mode"] == "av_fmp4"
|
||||
assert "avc1" in dumped["mime"]
|
||||
|
||||
def test_media_segment_complete(self):
|
||||
msg = MediaSegmentComplete(
|
||||
segment_idx=0, stream_id="abc", chunks=4,
|
||||
)
|
||||
dumped = msg.model_dump()
|
||||
assert dumped["chunks"] == 4
|
||||
|
||||
def test_ltx2_segment_complete(self):
|
||||
msg = Ltx2SegmentComplete(segment_idx=0, generation_time_ms=1234.5)
|
||||
assert msg.model_dump()["generation_time_ms"] == 1234.5
|
||||
|
||||
def test_error_message_code_restricted(self):
|
||||
with pytest.raises(ValidationError):
|
||||
ErrorMessage(code="not_a_code", message="x")
|
||||
|
||||
def test_continuation_state_snapshot(self):
|
||||
msg = ContinuationStateSnapshot(state={
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1},
|
||||
})
|
||||
assert msg.model_dump()["state"]["kind"] == "ltx2.v1"
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""End-to-end WebSocket smoke for the streaming server skeleton.
|
||||
|
||||
Uses a mock generator so these tests run CPU-only (no GPU, no model
|
||||
weights). Skips the fMP4 assertions when ``ffmpeg`` is missing.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("starlette")
|
||||
from starlette.testclient import TestClient # noqa: E402
|
||||
|
||||
from fastvideo.api.schema import ( # noqa: E402
|
||||
ContinuationState,
|
||||
GeneratorConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
StreamingConfig,
|
||||
GenerationRequest,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.server import build_app # noqa: E402
|
||||
|
||||
|
||||
_FFMPEG_AVAILABLE = shutil.which("ffmpeg") is not None
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MockGenerator:
|
||||
width: int = 64
|
||||
height: int = 64
|
||||
fps: int = 12
|
||||
num_frames: int = 12
|
||||
return_state: bool = True
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
frames = [
|
||||
np.full((self.height, self.width, 3), i * 5, dtype=np.uint8)
|
||||
for i in range(self.num_frames)
|
||||
]
|
||||
state = (ContinuationState(
|
||||
kind="ltx2.v1",
|
||||
payload={
|
||||
"schema_version": 1,
|
||||
"segment_index": 0,
|
||||
"source_prompt": request.prompt,
|
||||
},
|
||||
) if self.return_state else None)
|
||||
return {
|
||||
"frames": frames,
|
||||
"audio_sample_rate": 24000,
|
||||
"state": state,
|
||||
}
|
||||
|
||||
|
||||
def _build_serve_config() -> ServeConfig:
|
||||
return ServeConfig(
|
||||
generator=GeneratorConfig(model_path="/models/fake"),
|
||||
default_request=GenerationRequest(
|
||||
sampling=SamplingConfig(
|
||||
num_frames=12,
|
||||
height=64,
|
||||
width=64,
|
||||
fps=12,
|
||||
num_inference_steps=1,
|
||||
),
|
||||
),
|
||||
streaming=StreamingConfig(
|
||||
session_timeout_seconds=60,
|
||||
generation_segment_cap=2,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _build_client() -> tuple[TestClient, _MockGenerator]:
|
||||
generator = _MockGenerator()
|
||||
app = build_app(_build_serve_config(), generator)
|
||||
return TestClient(app), generator
|
||||
|
||||
|
||||
class TestHealth:
|
||||
|
||||
def test_health_endpoint_reports_stream_mode(self):
|
||||
client, _ = _build_client()
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["stream_mode"] == "av_fmp4"
|
||||
assert body["sessions"] == 0
|
||||
|
||||
|
||||
class TestSessionHandshake:
|
||||
|
||||
def test_rejects_non_session_init_opening_frame(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "segment_prompt_source", "prompt": "x"})
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
assert err["code"] == "invalid_message"
|
||||
|
||||
def test_rejects_unknown_message_on_init(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "not_a_message"})
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
|
||||
def test_emits_queue_and_gpu_assigned_on_valid_init(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({
|
||||
"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"curated_prompts": ["a fox"],
|
||||
})
|
||||
assert ws.receive_json()["type"] == "queue_status"
|
||||
assert ws.receive_json()["type"] == "gpu_assigned"
|
||||
assert ws.receive_json()["type"] == "ltx2_stream_start"
|
||||
|
||||
def test_init_hydrates_continuation_state(self):
|
||||
client, _ = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({
|
||||
"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage",
|
||||
"continuation_state": {
|
||||
"kind": "ltx2.v1",
|
||||
"payload": {"schema_version": 1, "segment_index": 3},
|
||||
},
|
||||
})
|
||||
# Drain handshake frames
|
||||
ws.receive_json() # queue_status
|
||||
ws.receive_json() # gpu_assigned
|
||||
ws.receive_json() # ltx2_stream_start
|
||||
# Ask the server for the state back; it should echo what we sent.
|
||||
ws.send_json({"type": "snapshot_state"})
|
||||
snap = ws.receive_json()
|
||||
assert snap["type"] == "continuation_state_snapshot"
|
||||
assert snap["state"]["kind"] == "ltx2.v1"
|
||||
assert snap["state"]["payload"]["segment_index"] == 3
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _FFMPEG_AVAILABLE, reason="ffmpeg not installed")
|
||||
class TestSegmentFlow:
|
||||
|
||||
def test_segment_generates_media_init_plus_complete(self):
|
||||
client, generator = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage"})
|
||||
for _ in range(3):
|
||||
ws.receive_json() # queue_status + gpu_assigned + stream_start
|
||||
|
||||
ws.send_json({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "a test segment",
|
||||
"num_inference_steps": 1,
|
||||
})
|
||||
start = ws.receive_json()
|
||||
assert start["type"] == "ltx2_segment_start"
|
||||
assert start["segment_idx"] == 0
|
||||
step = ws.receive_json()
|
||||
assert step["type"] == "step_complete"
|
||||
media_init = ws.receive_json()
|
||||
assert media_init["type"] == "media_init"
|
||||
# Then one or more binary frames until media_segment_complete.
|
||||
saw_binary = False
|
||||
while True:
|
||||
msg = ws.receive()
|
||||
if "bytes" in msg and msg["bytes"]:
|
||||
saw_binary = True
|
||||
continue
|
||||
parsed = _as_json(msg)
|
||||
if parsed is None:
|
||||
continue
|
||||
if parsed["type"] == "media_segment_complete":
|
||||
break
|
||||
assert saw_binary
|
||||
final = ws.receive_json()
|
||||
assert final["type"] == "ltx2_segment_complete"
|
||||
assert final["segment_idx"] == 0
|
||||
|
||||
|
||||
class TestContinuationStatePersistence:
|
||||
|
||||
def test_snapshot_after_segment_carries_generator_state(self):
|
||||
if not _FFMPEG_AVAILABLE:
|
||||
pytest.skip("ffmpeg not installed")
|
||||
client, generator = _build_client()
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2",
|
||||
"preset": "ltx2_two_stage"})
|
||||
for _ in range(3):
|
||||
ws.receive_json()
|
||||
ws.send_json({
|
||||
"type": "segment_prompt_source",
|
||||
"prompt": "a cat",
|
||||
"num_inference_steps": 1,
|
||||
})
|
||||
_drain_until(ws, "ltx2_segment_complete")
|
||||
ws.send_json({"type": "snapshot_state"})
|
||||
snap = ws.receive_json()
|
||||
assert snap["type"] == "continuation_state_snapshot"
|
||||
assert snap["state"]["kind"] == "ltx2.v1"
|
||||
assert snap["state"]["payload"]["source_prompt"] == "a cat"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
def _drain_until(ws, target_type: str) -> dict[str, Any]:
|
||||
while True:
|
||||
msg = ws.receive()
|
||||
if "text" in msg and msg["text"]:
|
||||
import json
|
||||
|
||||
parsed = json.loads(msg["text"])
|
||||
if parsed.get("type") == target_type:
|
||||
return parsed
|
||||
# skip binary / other
|
||||
|
||||
|
||||
def _as_json(msg: dict[str, Any]) -> dict[str, Any] | None:
|
||||
if "text" not in msg or not msg["text"]:
|
||||
return None
|
||||
import json
|
||||
|
||||
return json.loads(msg["text"])
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Session lifecycle tests."""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.session import (
|
||||
InvalidSessionTransition,
|
||||
Session,
|
||||
SessionManager,
|
||||
SessionRejected,
|
||||
SessionState,
|
||||
)
|
||||
|
||||
|
||||
class TestSessionStateMachine:
|
||||
|
||||
def test_starts_initializing(self):
|
||||
s = Session()
|
||||
assert s.state is SessionState.INITIALIZING
|
||||
|
||||
def test_legal_sequence(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.COMPLETE)
|
||||
assert s.state is SessionState.COMPLETE
|
||||
|
||||
def test_active_self_loop_allowed(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.ACTIVE) # re-asserting is fine
|
||||
assert s.state is SessionState.ACTIVE
|
||||
|
||||
def test_illegal_backwards_transition(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.INITIALIZING)
|
||||
|
||||
def test_cannot_leave_terminal_state(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.transition(SessionState.COMPLETE)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.ACTIVE)
|
||||
|
||||
def test_error_terminal(self):
|
||||
s = Session()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.ERROR)
|
||||
with pytest.raises(InvalidSessionTransition):
|
||||
s.transition(SessionState.ACTIVE)
|
||||
|
||||
def test_transition_updates_activity(self):
|
||||
s = Session()
|
||||
prior = s.last_activity
|
||||
time.sleep(0.001)
|
||||
s.transition(SessionState.QUEUED)
|
||||
assert s.last_activity > prior
|
||||
|
||||
def test_segment_cap(self):
|
||||
s = Session()
|
||||
s.segment_idx = 5
|
||||
assert s.segment_cap_reached(5) is True
|
||||
assert s.segment_cap_reached(6) is False
|
||||
|
||||
|
||||
class TestSessionManager:
|
||||
|
||||
def test_create_assigns_unique_ids(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=2)
|
||||
a = mgr.create()
|
||||
b = mgr.create()
|
||||
assert a.id != b.id
|
||||
assert len(mgr) == 2
|
||||
|
||||
def test_max_sessions_enforced(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
|
||||
mgr.create()
|
||||
with pytest.raises(SessionRejected):
|
||||
mgr.create()
|
||||
|
||||
def test_close_releases_slot(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=1)
|
||||
s = mgr.create()
|
||||
mgr.close(s.id)
|
||||
assert len(mgr) == 0
|
||||
# Now can create again.
|
||||
mgr.create()
|
||||
|
||||
def test_reap_timed_out_flags_idle_sessions(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
|
||||
s = mgr.create()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.GPU_BINDING)
|
||||
s.transition(SessionState.ACTIVE)
|
||||
s.last_activity = time.monotonic() - 10 # 10s ago, past the 1s budget
|
||||
dead = mgr.reap_timed_out()
|
||||
assert s.id in dead
|
||||
|
||||
def test_reap_skips_terminal_states(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=1, max_sessions=4)
|
||||
s = mgr.create()
|
||||
s.transition(SessionState.QUEUED)
|
||||
s.transition(SessionState.ERROR)
|
||||
s.last_activity = time.monotonic() - 10
|
||||
assert s.id not in mgr.reap_timed_out()
|
||||
|
||||
def test_active_sessions_filter(self):
|
||||
mgr = SessionManager(
|
||||
segment_cap=4, session_timeout_seconds=60, max_sessions=4)
|
||||
a = mgr.create()
|
||||
a.transition(SessionState.QUEUED)
|
||||
a.transition(SessionState.GPU_BINDING)
|
||||
a.transition(SessionState.ACTIVE)
|
||||
b = mgr.create() # INITIALIZING
|
||||
assert mgr.active_sessions() == [a]
|
||||
assert b not in mgr.active_sessions()
|
||||
@@ -0,0 +1,90 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the session init-image persistence helper."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image,
|
||||
)
|
||||
|
||||
|
||||
def _png_bytes(size: tuple[int, int] = (64, 64)) -> bytes:
|
||||
buffer = io.BytesIO()
|
||||
Image.new("RGB", size, color=(10, 20, 30)).save(buffer, format="PNG")
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
class TestPersistSessionInitImage:
|
||||
|
||||
def test_none_payload_returns_none(self):
|
||||
assert persist_session_init_image(None) is None
|
||||
assert persist_session_init_image({}) is None
|
||||
|
||||
def test_non_object_payload_rejected(self):
|
||||
with pytest.raises(ValueError):
|
||||
persist_session_init_image("not-a-dict")
|
||||
|
||||
def test_png_payload_persists(self, tmp_path):
|
||||
data = _png_bytes()
|
||||
image = persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"name": "ref.png",
|
||||
"data": base64.b64encode(data).decode("ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
assert image is not None
|
||||
assert os.path.exists(image.path)
|
||||
assert image.mime == "image/png"
|
||||
assert image.path.endswith(".png")
|
||||
with open(image.path, "rb") as f:
|
||||
assert f.read() == data
|
||||
|
||||
def test_unknown_mime_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="mime"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/bmp",
|
||||
"data": "ignored",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_bad_base64_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="base64"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": "not!base64!",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_empty_data_rejected(self, tmp_path):
|
||||
with pytest.raises(ValueError, match="empty"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": "",
|
||||
}, output_dir=str(tmp_path))
|
||||
|
||||
def test_display_name_sanitized(self, tmp_path):
|
||||
image = persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"name": "../evil/../name.png",
|
||||
"data": base64.b64encode(_png_bytes()).decode("ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
assert image is not None
|
||||
assert image.display_name == "name.png"
|
||||
|
||||
def test_oversize_rejected(self, tmp_path):
|
||||
from fastvideo.entrypoints.streaming import session_init_image as mod
|
||||
|
||||
original = mod._MAX_IMAGE_BYTES
|
||||
mod._MAX_IMAGE_BYTES = 100
|
||||
try:
|
||||
with pytest.raises(ValueError, match="limit"):
|
||||
persist_session_init_image({
|
||||
"mime": "image/png",
|
||||
"data": base64.b64encode(_png_bytes((512, 512))).decode(
|
||||
"ascii"),
|
||||
}, output_dir=str(tmp_path))
|
||||
finally:
|
||||
mod._MAX_IMAGE_BYTES = original
|
||||
@@ -0,0 +1,99 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the fMP4 encoder.
|
||||
|
||||
These tests require ``ffmpeg`` on PATH. Skip when missing so the suite
|
||||
stays CPU/CI friendly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import shutil
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
shutil.which("ffmpeg") is None,
|
||||
reason="ffmpeg not installed",
|
||||
)
|
||||
|
||||
|
||||
def _frame(width: int, height: int, value: int = 128) -> np.ndarray:
|
||||
return np.full((height, width, 3), value, dtype=np.uint8)
|
||||
|
||||
|
||||
def test_encoder_emits_init_then_media_chunks():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
chunks: list[FragmentedMP4Chunk] = []
|
||||
async with enc:
|
||||
frames = [_frame(64, 64, v) for v in range(4, 28)]
|
||||
async for chunk in enc.encode(frames):
|
||||
chunks.append(chunk)
|
||||
assert len(chunks) > 0
|
||||
assert chunks[0].kind == "init"
|
||||
assert all(c.stream_id == enc.stream_id for c in chunks)
|
||||
assert all(c.segment_idx == 0 for c in chunks)
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_init_chunk_is_fmp4():
|
||||
"""The first chunk must contain the ``ftyp`` box (fMP4 init segment)."""
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
first_chunk = None
|
||||
async with enc:
|
||||
async for chunk in enc.encode([_frame(64, 64, 20)] * 24):
|
||||
first_chunk = chunk
|
||||
break
|
||||
assert first_chunk is not None
|
||||
assert first_chunk.kind == "init"
|
||||
# Box header: 4 bytes length, 4 bytes type. "ftyp" should appear
|
||||
# near the start of the init segment.
|
||||
assert b"ftyp" in first_chunk.data[:32]
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_rejects_non_ndarray_frames():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
async with enc:
|
||||
with pytest.raises(TypeError):
|
||||
async for _ in enc.encode(["not-a-frame"]):
|
||||
pass
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_rejects_wrong_shape():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
async with enc:
|
||||
with pytest.raises(ValueError):
|
||||
async for _ in enc.encode(
|
||||
[np.zeros((64, 64, 4), dtype=np.uint8)]):
|
||||
pass
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_encoder_close_is_idempotent():
|
||||
async def run():
|
||||
enc = FragmentedMP4Encoder(
|
||||
width=64, height=64, fps=24, segment_idx=0)
|
||||
await enc.__aenter__()
|
||||
await enc.close()
|
||||
await enc.close() # no raise
|
||||
|
||||
asyncio.run(run())
|
||||
Reference in New Issue
Block a user