[feat] Improve API: streaming auxiliaries (safety, rewrite, logger, mock) (#1284)

Co-authored-by: Junda (David) Su <90978028+Davids048@users.noreply.github.com>
Co-authored-by: Matthew Noto <99706358+RandNMR73@users.noreply.github.com>
Co-authored-by: XOR-op <17672363+XOR-op@users.noreply.github.com>
Co-authored-by: Zhang Peiyuan <42993249+jzhang38@users.noreply.github.com>
This commit is contained in:
William Lin
2026-05-05 00:14:34 -07:00
committed by GitHub
co-authored by Junda Su Matthew Noto XOR-op Zhang Peiyuan
parent f673423b51
commit eb3a394224
7 changed files with 755 additions and 0 deletions
@@ -17,10 +17,22 @@ from fastvideo.entrypoints.streaming.gpu_pool import (
PoolAcquireTimeout,
SubprocessGpuPool,
)
from fastvideo.entrypoints.streaming.mock_server import (
MockGenerator,
build_mock_app,
)
from fastvideo.entrypoints.streaming.prompt import (
LLMProvider,
PromptEnhancer,
)
from fastvideo.entrypoints.streaming.prompt.safety import (
PromptSafetyFilter,
SafetyDecision,
)
from fastvideo.entrypoints.streaming.session_logger import (
SessionLogEvent,
SessionLogger,
)
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
@@ -35,8 +47,14 @@ __all__ = [
"InMemorySessionStore",
"InProcessGpuPool",
"LLMProvider",
"MockGenerator",
"PoolAcquireTimeout",
"PromptEnhancer",
"PromptSafetyFilter",
"SafetyDecision",
"SessionLogEvent",
"SessionLogger",
"build_mock_app",
"Session",
"SessionManager",
"SessionState",
@@ -0,0 +1,122 @@
# SPDX-License-Identifier: Apache-2.0
"""Mock streaming server — a frontend dev aid.
Boots the same FastAPI app the real streaming server uses, but backs
it with :class:`InProcessGpuPool` wrapping a synthetic generator that
emits pre-baked RGB frames. No GPU or model weights required.
Use cases:
* Frontend development without a real model loaded.
* Integration tests that exercise the WS protocol end-to-end.
* Reproducing protocol bugs locally.
Launch: ``python -m fastvideo.entrypoints.streaming.mock_server``.
"""
from __future__ import annotations
import argparse
import time
from dataclasses import dataclass
from typing import Any
import numpy as np
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
GeneratorConfig,
SamplingConfig,
ServeConfig,
StreamingConfig,
)
from fastvideo.entrypoints.streaming.server import build_app
@dataclass
class MockGenerator:
"""Generator stand-in that returns synthetic gradient frames.
Each call produces one segment worth of frames whose pixels vary by
a constant derived from the request seed and segment index. Latency
is configurable via ``sleep_ms`` so the caller can exercise slow-
generate scenarios without spinning a GPU.
"""
sleep_ms: float = 0.0
def generate(self, request: GenerationRequest) -> dict[str, Any]:
if self.sleep_ms:
time.sleep(self.sleep_ms / 1000.0)
width = max(16, request.sampling.width)
height = max(16, request.sampling.height)
num_frames = max(1, request.sampling.num_frames)
frames = [_gradient_frame(height, width, idx, seed=request.sampling.seed) for idx in range(num_frames)]
state = ContinuationState(
kind="ltx2.v1",
payload={
"schema_version": 1,
"segment_index": 0,
"source_prompt": request.prompt,
},
)
return {
"frames": frames,
"audio_sample_rate": 24000,
"state": state,
}
def _gradient_frame(height: int, width: int, idx: int, *, seed: int) -> np.ndarray:
base = (idx * 17 + seed * 3) % 256
row = np.linspace(base, (base + 64) % 256, width, dtype=np.uint8)
frame = np.tile(row, (height, 1))
stacked = np.stack([frame, np.roll(frame, 8, axis=1), np.roll(frame, 16, axis=1)], axis=-1)
return stacked.astype(np.uint8)
def build_mock_app(*, sleep_ms: float = 0.0):
"""Build a FastAPI app backed by :class:`MockGenerator`."""
serve_config = ServeConfig(
generator=GeneratorConfig(model_path="/models/mock"),
streaming=StreamingConfig(
session_timeout_seconds=120,
generation_segment_cap=6,
),
)
serve_config.default_request.sampling = SamplingConfig(
num_frames=24,
height=256,
width=256,
fps=24,
num_inference_steps=1,
)
return build_app(serve_config, MockGenerator(sleep_ms=sleep_ms))
def main() -> None: # pragma: no cover - CLI entry
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8000)
parser.add_argument(
"--sleep-ms",
type=float,
default=0.0,
help="Per-segment artificial latency for testing slow paths",
)
args = parser.parse_args()
import uvicorn
app = build_mock_app(sleep_ms=args.sleep_ms)
uvicorn.run(app, host=args.host, port=args.port)
__all__ = [
"MockGenerator",
"build_mock_app",
"main",
]
if __name__ == "__main__": # pragma: no cover - CLI entry
main()
@@ -0,0 +1,82 @@
# SPDX-License-Identifier: Apache-2.0
"""Rewrite payload builder.
The UI's "rewrite seed prompts" flow asks the enhancer to produce a
batch of alternative prompts given one seed. This module packages the
seed + options into the payload the enhancer expects and unpacks the
response back into a typed :class:`RewriteResult`.
Separating this from :mod:`enhancer` keeps the enhancer provider-
agnostic; anything UI-specific (how many alternatives to request, how
to split the response, temperature) lives here.
"""
from __future__ import annotations
import re
from dataclasses import dataclass
from fastvideo.entrypoints.streaming.prompt.enhancer import PromptEnhancer
_LEADING_MARKER_RE = re.compile(r"^(?:[-*•]\s*|\d+\s*[.)]\s*)+")
@dataclass
class RewriteOptions:
count: int = 3
"""Number of alternative prompts to request."""
temperature: float | None = None
@dataclass
class RewriteResult:
seed_prompt: str
alternatives: list[str]
provider: str
model: str
latency_ms: float
fallback_used: bool = False
async def build_rewrite(
enhancer: PromptEnhancer,
seed_prompt: str,
*,
options: RewriteOptions | None = None,
) -> RewriteResult:
"""Run a rewrite op through the enhancer and return a typed result."""
if not seed_prompt.strip():
raise ValueError("rewrite seed prompt must be non-empty")
options = options or RewriteOptions()
response = await enhancer.rewrite(seed_prompt)
alternatives = _split_response(response.content, limit=options.count)
return RewriteResult(
seed_prompt=seed_prompt,
alternatives=alternatives,
provider=response.provider,
model=response.model,
latency_ms=response.latency_ms,
fallback_used=response.fallback_used,
)
def _split_response(content: str, *, limit: int) -> list[str]:
"""Split the LLM response into discrete prompt candidates.
The shipped system prompt instructs the model to emit one prompt
per line; this function is forgiving about numbered lists or
leading bullets so user-supplied system prompts don't break it.
"""
lines = [line.strip() for line in content.splitlines() if line.strip()]
cleaned: list[str] = []
for line in lines:
stripped = _LEADING_MARKER_RE.sub("", line).strip()
if stripped:
cleaned.append(stripped)
return cleaned[:max(1, limit)]
__all__ = [
"RewriteOptions",
"RewriteResult",
"build_rewrite",
]
@@ -0,0 +1,146 @@
# SPDX-License-Identifier: Apache-2.0
"""Optional prompt safety filter.
Uses a fastText classifier to score prompts against a banned-content
rubric. Only loaded when ``ServeConfig.streaming.safety.enabled`` is
True and fastText is installed — users who don't need it see no
runtime cost.
Install: ``pip install fastvideo[prompt-safety]`` (ships fasttext as an
optional extra) or install fasttext directly.
"""
from __future__ import annotations
import enum
import threading
from dataclasses import dataclass
from typing import Any
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class SafetyDecision(enum.Enum):
ALLOW = "allow"
BLOCK = "block"
UNAVAILABLE = "unavailable"
"""Returned when the classifier can't run (not configured, fastText
missing). Safety is opt-in; the server treats ``UNAVAILABLE`` as
``ALLOW`` but logs it so operators know the filter is off."""
@dataclass
class SafetyResult:
prompt: str
decision: SafetyDecision
score: float = 0.0
label: str | None = None
reason: str | None = None
class PromptSafetyFilter:
"""Minimal fastText-backed prompt safety filter.
Loads the classifier lazily on first use so the streaming server
can construct the filter eagerly at startup without paying the
model-load cost when safety is disabled.
"""
def __init__(
self,
*,
classifier_path: str | None,
enabled: bool = True,
block_threshold: float = 0.5,
) -> None:
self._classifier_path = classifier_path
self._enabled = enabled
self._block_threshold = block_threshold
self._model: Any | None = None
self._load_attempted = False
self._load_lock = threading.Lock()
@property
def enabled(self) -> bool:
return self._enabled and self._classifier_path is not None
def classify(self, prompt: str) -> SafetyResult:
if not self.enabled:
return SafetyResult(
prompt=prompt,
decision=SafetyDecision.UNAVAILABLE,
reason="safety filter not enabled",
)
model = self._ensure_loaded()
if model is None:
return SafetyResult(
prompt=prompt,
decision=SafetyDecision.UNAVAILABLE,
reason="fastText model unavailable",
)
try:
labels, probs = model.predict(prompt.replace("\n", " "), k=1)
except Exception as exc: # pragma: no cover - defensive
logger.warning("safety: classifier failed: %s", exc)
return SafetyResult(
prompt=prompt,
decision=SafetyDecision.UNAVAILABLE,
reason=f"classifier error: {exc}",
)
label = labels[0].removeprefix("__label__") if labels else None
score = float(probs[0]) if len(probs) else 0.0
decision = (SafetyDecision.BLOCK if
(label == "unsafe" and score >= self._block_threshold) else SafetyDecision.ALLOW)
return SafetyResult(
prompt=prompt,
decision=decision,
score=score,
label=label,
)
def _ensure_loaded(self) -> Any | None:
if self._model is not None:
return self._model
if self._load_attempted:
return None
with self._load_lock:
if self._model is not None:
return self._model
if self._load_attempted:
return None
self._load_attempted = True
if self._classifier_path is None:
return None
try:
import fasttext # type: ignore[import-not-found]
except ImportError:
logger.warning("safety: fasttext not installed; safety filter disabled. "
"Install fastvideo[prompt-safety] to enable.")
return None
try:
self._model = fasttext.load_model(self._classifier_path)
except Exception as exc: # pragma: no cover - requires real model
logger.warning("safety: failed to load %s: %s", self._classifier_path, exc)
return None
return self._model
def first_blocked(
filter_: PromptSafetyFilter,
prompts: list[str],
) -> SafetyResult | None:
"""Return the first prompt the filter blocks, or ``None``."""
for prompt in prompts:
result = filter_.classify(prompt)
if result.decision is SafetyDecision.BLOCK:
return result
return None
__all__ = [
"PromptSafetyFilter",
"SafetyDecision",
"SafetyResult",
"first_blocked",
]
@@ -0,0 +1,113 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-session JSONL event logger.
Each session gets its own JSONL file under the configured log root so
post-hoc analytics (enhancer latency, GPU assignment, segment timings)
can be recovered without a tracing backend. The internal UI uses this
format; keeping the same shape makes log tooling portable.
"""
from __future__ import annotations
import contextlib
import json
import os
import re
import threading
import time
from dataclasses import dataclass, field
from typing import Any, TextIO
_FILENAME_SANITIZE_RE = re.compile(r"[^A-Za-z0-9._-]")
@dataclass
class SessionLogEvent:
"""One line in the session JSONL file."""
session_id: str
event: str
payload: dict[str, Any] = field(default_factory=dict)
ts: float = field(default_factory=time.time)
class SessionLogger:
"""Append-only JSONL logger keyed by session id.
Thread-safe; the server may be writing from multiple asyncio tasks
(fMP4 encoder thread + control-frame handler) for the same session.
"""
def __init__(self, log_dir: str | None) -> None:
self._log_dir = log_dir
self._files: dict[str, TextIO] = {}
self._locks: dict[str, threading.Lock] = {}
self._registry_lock = threading.Lock()
self._ensure_dir()
def log(self, event: SessionLogEvent) -> None:
if self._log_dir is None:
return
opened = self._get_file(event.session_id)
if opened is None:
return
handle, lock = opened
line = json.dumps({
"session_id": event.session_id,
"event": event.event,
"ts": event.ts,
"payload": event.payload,
})
with lock, contextlib.suppress(ValueError):
handle.write(line + "\n")
handle.flush()
def close(self, session_id: str) -> None:
with self._registry_lock:
handle = self._files.pop(session_id, None)
lock = self._locks.pop(session_id, None)
if handle is None or lock is None:
return
with lock, contextlib.suppress(Exception):
handle.close()
def close_all(self) -> None:
with self._registry_lock:
sids = list(self._files)
for sid in sids:
self.close(sid)
def _ensure_dir(self) -> None:
if self._log_dir is None:
return
os.makedirs(self._log_dir, exist_ok=True)
def _get_file(self, session_id: str) -> tuple[TextIO, threading.Lock] | None:
if self._log_dir is None:
return None
with self._registry_lock:
handle = self._files.get(session_id)
lock = self._locks.get(session_id)
if handle is not None and lock is not None:
return handle, lock
# Defense-in-depth: session_id is server-generated UUID today,
# but sanitize against path traversal in case future code paths
# allow client-supplied ids.
safe_id = _FILENAME_SANITIZE_RE.sub("_", session_id) or "unknown"
path = os.path.join(
self._log_dir,
f"session-{safe_id}.jsonl",
)
try:
handle = open(path, "a", encoding="utf-8") # noqa: SIM115
except OSError:
return None
lock = threading.Lock()
self._files[session_id] = handle
self._locks[session_id] = lock
return handle, lock
__all__ = [
"SessionLogEvent",
"SessionLogger",
]
@@ -0,0 +1,261 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for PR 7.8 streaming auxiliaries.
Covers:
* PromptSafetyFilter gracefully disables when fastText isn't installed
* SafetyResult semantics (allow, block, unavailable)
* RewriteOptions + _split_response parsing behavior
* SessionLogger JSONL append semantics + close lifecycle
* MockServer builds an app that drives the WS protocol end-to-end
"""
from __future__ import annotations
import json
import sys
import types
from pathlib import Path
from typing import Any
import pytest
from fastvideo.entrypoints.streaming.prompt.safety import (
PromptSafetyFilter,
SafetyDecision,
first_blocked,
)
from fastvideo.entrypoints.streaming.prompt.rewrite import (
RewriteOptions,
_split_response,
build_rewrite,
)
from fastvideo.entrypoints.streaming.session_logger import (
SessionLogEvent,
SessionLogger,
)
# ----------------------------------------------------------------------
# Safety
# ----------------------------------------------------------------------
class TestPromptSafetyFilter:
def test_disabled_by_default_when_no_path(self):
f = PromptSafetyFilter(classifier_path=None)
assert f.enabled is False
result = f.classify("hi")
assert result.decision is SafetyDecision.UNAVAILABLE
def test_disabled_when_enabled_false(self):
f = PromptSafetyFilter(classifier_path="/tmp/m.bin", enabled=False)
assert f.enabled is False
def test_unavailable_when_fasttext_missing(self, monkeypatch):
# Force `import fasttext` inside _ensure_loaded to fail.
monkeypatch.setitem(sys.modules, "fasttext", None)
f = PromptSafetyFilter(classifier_path="/tmp/m.bin", enabled=True)
result = f.classify("hi")
assert result.decision is SafetyDecision.UNAVAILABLE
def test_block_when_classifier_flags_unsafe(self, monkeypatch, tmp_path):
fake_model = types.SimpleNamespace(
predict=lambda text, k=1: (["__label__unsafe"], [0.95]))
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
monkeypatch.setitem(sys.modules, "fasttext", stub)
model_path = str(tmp_path / "m.bin")
Path(model_path).write_text("")
f = PromptSafetyFilter(classifier_path=model_path, enabled=True)
result = f.classify("please")
assert result.decision is SafetyDecision.BLOCK
assert result.label == "unsafe"
assert result.score == pytest.approx(0.95)
def test_allow_when_classifier_flags_safe(self, monkeypatch, tmp_path):
fake_model = types.SimpleNamespace(
predict=lambda text, k=1: (["__label__safe"], [0.99]))
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
monkeypatch.setitem(sys.modules, "fasttext", stub)
f = PromptSafetyFilter(classifier_path="ignored", enabled=True)
result = f.classify("hello")
assert result.decision is SafetyDecision.ALLOW
def test_below_threshold_allows_even_if_unsafe_label(self, monkeypatch):
fake_model = types.SimpleNamespace(
predict=lambda text, k=1: (["__label__unsafe"], [0.3]))
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
monkeypatch.setitem(sys.modules, "fasttext", stub)
f = PromptSafetyFilter(
classifier_path="m", enabled=True, block_threshold=0.5)
assert f.classify("x").decision is SafetyDecision.ALLOW
def test_first_blocked_returns_first_hit(self, monkeypatch):
responses = iter([
(["__label__safe"], [0.9]),
(["__label__unsafe"], [0.9]),
(["__label__safe"], [0.9]),
])
fake_model = types.SimpleNamespace(
predict=lambda text, k=1: next(responses))
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
monkeypatch.setitem(sys.modules, "fasttext", stub)
f = PromptSafetyFilter(classifier_path="m", enabled=True)
blocked = first_blocked(f, ["ok", "bad", "also ok"])
assert blocked is not None
assert blocked.prompt == "bad"
# ----------------------------------------------------------------------
# Rewrite
# ----------------------------------------------------------------------
class TestRewriteSplit:
def test_plain_lines(self):
assert _split_response("one\ntwo\nthree", limit=3) == [
"one", "two", "three",
]
def test_numbered_list(self):
assert _split_response("1. first\n2. second", limit=3) == [
"first", "second",
]
def test_bulleted_list(self):
assert _split_response("- one\n* two\n• three", limit=3) == [
"one", "two", "three",
]
def test_respects_limit(self):
assert _split_response("a\nb\nc\nd", limit=2) == ["a", "b"]
def test_limit_min_one(self):
assert _split_response("only one", limit=0) == ["only one"]
class _StubEnhancer:
async def rewrite(self, seed):
from fastvideo.entrypoints.streaming.prompt.providers.base import LLMResponse
return LLMResponse(
content="1. alpha\n2. beta\n3. gamma",
provider="stub",
model="m",
latency_ms=1.0,
)
class TestBuildRewrite:
def test_empty_seed_rejected(self):
import asyncio
with pytest.raises(ValueError):
asyncio.run(build_rewrite(_StubEnhancer(), " "))
def test_returns_limited_alternatives(self):
import asyncio
result = asyncio.run(build_rewrite(
_StubEnhancer(), "seed",
options=RewriteOptions(count=2)))
assert result.seed_prompt == "seed"
assert result.alternatives == ["alpha", "beta"]
assert result.provider == "stub"
# ----------------------------------------------------------------------
# Session logger
# ----------------------------------------------------------------------
class TestSessionLogger:
def test_no_log_dir_is_noop(self):
logger = SessionLogger(None)
logger.log(SessionLogEvent(session_id="s", event="x")) # no raise
def test_appends_jsonl(self, tmp_path):
logger = SessionLogger(str(tmp_path))
logger.log(SessionLogEvent(
session_id="s1",
event="start",
payload={"preset": "ltx2"},
ts=1.0,
))
logger.log(SessionLogEvent(
session_id="s1",
event="segment",
payload={"idx": 0},
ts=2.0,
))
logger.close("s1")
path = tmp_path / "session-s1.jsonl"
lines = path.read_text().splitlines()
assert len(lines) == 2
first = json.loads(lines[0])
assert first["event"] == "start"
assert first["payload"]["preset"] == "ltx2"
def test_separate_files_per_session(self, tmp_path):
logger = SessionLogger(str(tmp_path))
logger.log(SessionLogEvent(session_id="a", event="e"))
logger.log(SessionLogEvent(session_id="b", event="e"))
logger.close_all()
assert (tmp_path / "session-a.jsonl").exists()
assert (tmp_path / "session-b.jsonl").exists()
# ----------------------------------------------------------------------
# Mock server
# ----------------------------------------------------------------------
class TestMockServer:
def test_build_mock_app_returns_fastapi(self):
from fastapi import FastAPI
from fastvideo.entrypoints.streaming.mock_server import build_mock_app
app = build_mock_app()
assert isinstance(app, FastAPI)
def test_mock_generator_produces_frames(self):
from fastvideo.api.schema import GenerationRequest, SamplingConfig
from fastvideo.entrypoints.streaming.mock_server import MockGenerator
gen = MockGenerator()
result = gen.generate(GenerationRequest(
prompt="x",
sampling=SamplingConfig(
num_frames=3, height=32, width=32, num_inference_steps=1),
))
assert len(result["frames"]) == 3
assert result["frames"][0].shape == (32, 32, 3)
assert result["state"].kind == "ltx2.v1"
def test_mock_app_health_endpoint(self):
from starlette.testclient import TestClient
from fastvideo.entrypoints.streaming.mock_server import build_mock_app
app = build_mock_app()
client = TestClient(app)
assert client.get("/health").json()["status"] == "ok"
def test_mock_app_ws_handshake(self):
from starlette.testclient import TestClient
from fastvideo.entrypoints.streaming.mock_server import build_mock_app
app = build_mock_app()
client = TestClient(app)
with client.websocket_connect("/v1/stream") as ws:
ws.send_json({"type": "session_init_v2"})
assert ws.receive_json()["type"] == "queue_status"
assert ws.receive_json()["type"] == "gpu_assigned"
assert ws.receive_json()["type"] == "ltx2_stream_start"
+13
View File
@@ -134,6 +134,19 @@ test = [
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
prompt-safety = [
"fasttext",
]
prompt-enhancer = [
"httpx",
]
streaming = [
"fastvideo[prompt-enhancer]",
"fastvideo[prompt-safety]",
]
rocm = [
"amdsmi",
]