From e76e314901c1c0e86247848291e61629e4a16394 Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 10:44:12 +0000 Subject: [PATCH 01/11] feat(streaming): add LiveKit session runtime Add LiveKit configuration, admission scheduling, token issuance, room workers, media bridging, and session lifecycle APIs. Route stream-serve through the LiveKit runtime while keeping the legacy transport available for the later removal commit. Verification: - .venv/bin/python -m pytest tests/unit/service/livekit --ignore=tests/unit/service/livekit/test_demo.py -q - 42 passed - git diff --cached --check --- pyproject.toml | 2 + telefuser/entrypoints/cli/main.py | 66 ++-- telefuser/service/livekit/__init__.py | 7 + telefuser/service/livekit/app.py | 167 ++++++++++ telefuser/service/livekit/config.py | 93 ++++++ telefuser/service/livekit/data_protocol.py | 149 +++++++++ telefuser/service/livekit/main.py | 93 ++++++ telefuser/service/livekit/media_bridge.py | 86 +++++ telefuser/service/livekit/pipeline_adapter.py | 48 +++ telefuser/service/livekit/room_client.py | 186 +++++++++++ telefuser/service/livekit/runtime.py | 295 +++++++++++++++++ telefuser/service/livekit/scheduler.py | 176 ++++++++++ telefuser/service/livekit/schemas.py | 110 +++++++ telefuser/service/livekit/session_registry.py | 126 ++++++++ telefuser/service/livekit/token_service.py | 72 +++++ telefuser/service/livekit/worker.py | 270 ++++++++++++++++ telefuser/service/livekit/worker_pool.py | 84 +++++ tests/unit/service/livekit/__init__.py | 1 + tests/unit/service/livekit/test_app.py | 111 +++++++ tests/unit/service/livekit/test_cli.py | 45 +++ tests/unit/service/livekit/test_config.py | 31 ++ .../service/livekit/test_data_protocol.py | 147 +++++++++ .../unit/service/livekit/test_media_bridge.py | 98 ++++++ .../unit/service/livekit/test_room_client.py | 159 +++++++++ tests/unit/service/livekit/test_runtime.py | 135 ++++++++ tests/unit/service/livekit/test_scheduler.py | 54 ++++ .../service/livekit/test_session_registry.py | 38 +++ .../service/livekit/test_token_service.py | 105 ++++++ tests/unit/service/livekit/test_worker.py | 302 ++++++++++++++++++ 29 files changed, 3235 insertions(+), 21 deletions(-) create mode 100644 telefuser/service/livekit/__init__.py create mode 100644 telefuser/service/livekit/app.py create mode 100644 telefuser/service/livekit/config.py create mode 100644 telefuser/service/livekit/data_protocol.py create mode 100644 telefuser/service/livekit/main.py create mode 100644 telefuser/service/livekit/media_bridge.py create mode 100644 telefuser/service/livekit/pipeline_adapter.py create mode 100644 telefuser/service/livekit/room_client.py create mode 100644 telefuser/service/livekit/runtime.py create mode 100644 telefuser/service/livekit/scheduler.py create mode 100644 telefuser/service/livekit/schemas.py create mode 100644 telefuser/service/livekit/session_registry.py create mode 100644 telefuser/service/livekit/token_service.py create mode 100644 telefuser/service/livekit/worker.py create mode 100644 telefuser/service/livekit/worker_pool.py create mode 100644 tests/unit/service/livekit/__init__.py create mode 100644 tests/unit/service/livekit/test_app.py create mode 100644 tests/unit/service/livekit/test_cli.py create mode 100644 tests/unit/service/livekit/test_config.py create mode 100644 tests/unit/service/livekit/test_data_protocol.py create mode 100644 tests/unit/service/livekit/test_media_bridge.py create mode 100644 tests/unit/service/livekit/test_room_client.py create mode 100644 tests/unit/service/livekit/test_runtime.py create mode 100644 tests/unit/service/livekit/test_scheduler.py create mode 100644 tests/unit/service/livekit/test_session_registry.py create mode 100644 tests/unit/service/livekit/test_token_service.py create mode 100644 tests/unit/service/livekit/test_worker.py diff --git a/pyproject.toml b/pyproject.toml index a447275..83ac924 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,6 +38,8 @@ dependencies = [ "httpx", "imageio[ffmpeg]>=2.37.2", "loguru", + "livekit-api>=1.2.0,<2.0", + "livekit>=1.1.13,<2.0", "numpy", "opencv-python-headless>=4.0", "pillow", diff --git a/telefuser/entrypoints/cli/main.py b/telefuser/entrypoints/cli/main.py index fd1babf..c5d2ada 100644 --- a/telefuser/entrypoints/cli/main.py +++ b/telefuser/entrypoints/cli/main.py @@ -134,14 +134,23 @@ def serve( @main.command(name="stream-serve") @click.argument("pipe_path") -@click.option("--port", "-p", default=8088, type=int, help="Server port") -@click.option("--host", default="0.0.0.0", type=str, help="Server host") +@click.option("--host", default="0.0.0.0", type=str, help="HTTP API bind host") +@click.option("--port", "-p", default=8088, type=int, help="HTTP API bind port") +@click.option("--livekit-url", default=None, type=str, help="LiveKit server URL") +@click.option("--livekit-api-key", default=None, type=str, help="LiveKit API key") +@click.option("--livekit-api-secret", default=None, type=str, help="LiveKit API secret") +@click.option("--num-workers", default=1, type=int, help="Number of TeleFuser model workers") +@click.option("--worker-gpu-map", default=None, type=str, help="GPU groups, for example '0,1;2,3'") +@click.option("--queue-size", default=0, type=int, help="Maximum queued sessions; 0 rejects when busy") +@click.option("--session-timeout", default=1800, type=int, help="Maximum session lifetime in seconds") +@click.option("--token-ttl", default=3600, type=int, help="LiveKit join token TTL in seconds") +@click.option("--controller-timeout", default=60, type=int, help="Seconds to keep a session after controller leaves") +@click.option("--room-empty-timeout", default=30, type=int, help="Seconds to keep a session after room becomes empty") @click.option( - "--gpu-num", - "-g", - default=1, - type=click.IntRange(min=1), - help="Number of GPUs passed to a stream pipeline get_service(gpu_num=...) factory", + "--worker-mode", + type=click.Choice(["in-process", "process"], case_sensitive=False), + default="in-process", + help="Worker isolation mode", ) @click.option( "--security-level", @@ -157,22 +166,27 @@ def serve( ) def stream_serve( pipe_path: str, - port: int, host: str, - gpu_num: int, + port: int, + livekit_url: str | None, + livekit_api_key: str | None, + livekit_api_secret: str | None, + num_workers: int, + worker_gpu_map: str | None, + queue_size: int, + session_timeout: int, + token_ttl: int, + controller_timeout: int, + room_empty_timeout: int, + worker_mode: str, security_level: str, skip_validation: bool, ) -> None: - """Start the TeleFuser stream server (WebRTC / WebSocket). + """Start the LiveKit-backed TeleFuser stream server. \b PIPE_PATH is a Python file that defines get_service() returning - a ServerPushService (WebRTC) or BidirectionalService (WebSocket). - - \b - Examples: - telefuser stream-serve examples/stream_video_replay.py - telefuser stream-serve examples/stream_video_replay.py -p 8000 --host 0.0.0.0 + a ServerPushService or BidirectionalService stream service. """ if not skip_validation: level = SecurityLevel[security_level.upper()] @@ -184,13 +198,23 @@ def stream_serve( click.echo(f"\nTo bypass: telefuser stream-serve {pipe_path} --skip-validation", err=True) raise click.Abort() - from telefuser.service.main import run_stream_server + from telefuser.service.livekit.main import run_stream_server run_stream_server( - pipe_path, - port, - host, - gpu_num=gpu_num, + pipe_path=pipe_path, + host=host, + port=port, + livekit_url=livekit_url, + livekit_api_key=livekit_api_key, + livekit_api_secret=livekit_api_secret, + num_workers=num_workers, + worker_gpu_map=worker_gpu_map, + queue_size=queue_size, + session_timeout=session_timeout, + token_ttl=token_ttl, + controller_timeout=controller_timeout, + room_empty_timeout=room_empty_timeout, + worker_mode=worker_mode.lower(), skip_validation=skip_validation, security_level=security_level, ) diff --git a/telefuser/service/livekit/__init__.py b/telefuser/service/livekit/__init__.py new file mode 100644 index 0000000..f6b681a --- /dev/null +++ b/telefuser/service/livekit/__init__.py @@ -0,0 +1,7 @@ +"""LiveKit serving support for TeleFuser.""" + +from __future__ import annotations + +from .config import LiveKitServeConfig + +__all__ = ["LiveKitServeConfig"] diff --git a/telefuser/service/livekit/app.py b/telefuser/service/livekit/app.py new file mode 100644 index 0000000..f6b506b --- /dev/null +++ b/telefuser/service/livekit/app.py @@ -0,0 +1,167 @@ +"""FastAPI app for the LiveKit serving entrypoint.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + +from fastapi import FastAPI, HTTPException, Response, status +from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import JSONResponse + +from telefuser.metrics import get_service_metrics + +from .runtime import LiveKitServeRuntime, session_record_to_response +from .schemas import ( + LiveKitHealthResponse, + SessionCreateRequest, + SessionCreateResponse, + SessionDeleteResponse, + SessionStatusResponse, + SessionTokenRequest, + SessionTokenResponse, +) +from .token_service import LiveKitDependencyError + + +def create_livekit_app(runtime: LiveKitServeRuntime) -> FastAPI: + """Create the LiveKit-backed stream HTTP app.""" + + @asynccontextmanager + async def lifespan(_: FastAPI) -> AsyncIterator[None]: + try: + await runtime.start() + yield + finally: + await runtime.aclose() + + app = FastAPI( + title="TeleFuser Stream API", + description="LiveKit-backed real-time TeleFuser stream API.", + version="0.1.0", + docs_url="/docs", + redoc_url="/redoc", + openapi_url="/openapi.json", + lifespan=lifespan, + ) + app.add_middleware( + CORSMiddleware, + allow_origins=runtime.config.cors_allow_origins, + allow_methods=["*"], + allow_headers=["*"], + ) + + @app.post("/v1/stream/sessions", response_model=SessionCreateResponse) + async def create_session(request: SessionCreateRequest): + try: + result = runtime.create_session(request) + except LiveKitDependencyError as exc: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(exc)) from exc + + admission = result.admission + if admission.status == "rejected": + detail = admission.reason or "no_capacity" + raise HTTPException(status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail=detail) + + body = SessionCreateResponse( + session_id=result.record.session_id, + room=result.record.room_name, + livekit_url=runtime.config.livekit_url, + token=result.token, + worker_id=result.record.worker_id, + status=result.record.status, + expires_at=result.record.expires_at, + queue_position=admission.queue_position, + ) + if admission.status == "queued": + return JSONResponse(status_code=status.HTTP_202_ACCEPTED, content=body.model_dump()) + return body + + @app.post("/v1/stream/sessions/{session_id}/tokens", response_model=SessionTokenResponse) + async def create_token(session_id: str, request: SessionTokenRequest) -> SessionTokenResponse: + try: + record, token = runtime.create_viewer_token(session_id, request) + except KeyError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(exc)) from exc + except LiveKitDependencyError as exc: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(exc)) from exc + + return SessionTokenResponse( + session_id=record.session_id, + room=record.room_name, + livekit_url=runtime.config.livekit_url, + token=token, + role="viewer", + ) + + @app.get("/v1/stream/sessions/{session_id}", response_model=SessionStatusResponse) + async def get_session(session_id: str) -> SessionStatusResponse: + try: + return runtime.get_session_response(session_id) + except KeyError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + + @app.delete("/v1/stream/sessions/{session_id}", response_model=SessionDeleteResponse) + async def delete_session(session_id: str) -> SessionDeleteResponse: + try: + record = await runtime.delete_session(session_id) + except KeyError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + return SessionDeleteResponse(session_id=record.session_id, status=record.status) + + @app.get("/v1/stream/health", response_model=LiveKitHealthResponse) + async def livekit_health() -> LiveKitHealthResponse: + return runtime.health() + + @app.get("/v1/service/health") + async def service_health() -> dict: + health = runtime.health() + return { + "status": health.status, + "ready": runtime.is_ready and health.status != "unhealthy", + "service_type": "stream", + "transport": "livekit", + **health.model_dump(), + } + + @app.get("/v1/service/ready") + async def service_ready() -> JSONResponse: + health = runtime.health() + ready = runtime.is_ready and health.status != "unhealthy" + return JSONResponse( + status_code=status.HTTP_200_OK if ready else status.HTTP_503_SERVICE_UNAVAILABLE, + content={ + "status": "ready" if ready else "not_ready", + "ready": ready, + "service_type": "stream", + "transport": "livekit", + **health.model_dump(), + }, + ) + + @app.get("/v1/service/metadata") + async def service_metadata() -> dict: + return runtime.metadata() + + @app.get("/v1/service/metrics") + async def service_metrics() -> Response: + service_metrics_obj = get_service_metrics() + return Response( + content=service_metrics_obj.get_prometheus_format(), + media_type="text/plain; charset=utf-8", + ) + + @app.get("/v1/service/metrics/json") + async def service_metrics_json() -> dict: + service_metrics_obj = get_service_metrics() + health = runtime.health() + return { + "uptime_seconds": service_metrics_obj.service_uptime.value, + "service_type": "stream", + "transport": "livekit", + "livekit": health.model_dump(), + } + + return app diff --git a/telefuser/service/livekit/config.py b/telefuser/service/livekit/config.py new file mode 100644 index 0000000..58d072f --- /dev/null +++ b/telefuser/service/livekit/config.py @@ -0,0 +1,93 @@ +"""Configuration for the LiveKit-backed ``telefuser stream-serve`` command.""" + +from __future__ import annotations + +from typing import Literal + +from pydantic import Field, field_validator +from pydantic_settings import BaseSettings, SettingsConfigDict + + +class LiveKitServeConfig(BaseSettings): + """Runtime configuration for the LiveKit serving entrypoint.""" + + model_config = SettingsConfigDict( + env_prefix="TELEFUSER_LIVEKIT_", + case_sensitive=False, + extra="ignore", + ) + + host: str = Field(default="0.0.0.0", description="HTTP bind host") + port: int = Field(default=8088, ge=1, le=65535, description="HTTP bind port") + + livekit_url: str = Field(default="", description="LiveKit server URL") + livekit_api_key: str = Field(default="", description="LiveKit API key") + livekit_api_secret: str = Field(default="", description="LiveKit API secret") + + num_workers: int = Field(default=1, ge=1, le=64, description="Number of TeleFuser LiveKit workers") + worker_gpu_map: str | None = Field( + default=None, + description="Semicolon-separated worker GPU groups, for example '0,1;2,3'", + ) + worker_mode: Literal["in-process", "process"] = Field( + default="in-process", + description="Worker isolation mode", + ) + + queue_size: int = Field(default=0, ge=0, le=10000, description="Maximum queued sessions") + session_timeout: int = Field(default=1800, ge=1, description="Maximum session lifetime in seconds") + token_ttl: int = Field(default=3600, ge=1, description="LiveKit token TTL in seconds") + controller_timeout: int = Field( + default=60, + ge=0, + description="Seconds to keep a session after controller disconnect", + ) + room_empty_timeout: int = Field( + default=30, + ge=0, + description="Seconds to keep a session after the LiveKit room becomes empty", + ) + role_mode: Literal["single-controller"] = Field(default="single-controller") + + default_fps: int = Field(default=16, ge=1, le=120, description="Default output video FPS") + max_data_message_bytes: int = Field(default=12 * 1024, ge=1024, description="Maximum accepted data message size") + cors_allow_origins: list[str] = Field( + default_factory=lambda: ["*"], + description="CORS origins for the LiveKit serve API", + ) + + @field_validator("worker_gpu_map") + @classmethod + def validate_worker_gpu_map(cls: type[LiveKitServeConfig], value: str | None) -> str | None: + """Normalize empty GPU maps to ``None``.""" + if value is None: + return None + stripped = value.strip() + return stripped or None + + def require_livekit_credentials(self) -> None: + """Raise when the minimum LiveKit connection settings are missing.""" + missing = [] + if not self.livekit_url: + missing.append("livekit_url") + if not self.livekit_api_key: + missing.append("livekit_api_key") + if not self.livekit_api_secret: + missing.append("livekit_api_secret") + if missing: + joined = ", ".join(missing) + raise ValueError(f"Missing LiveKit configuration: {joined}") + + def worker_gpu_groups(self) -> list[list[str]]: + """Return one GPU-id group per configured worker.""" + if self.worker_gpu_map is None: + return [[] for _ in range(self.num_workers)] + + groups = [ + [gpu.strip() for gpu in group.split(",") if gpu.strip()] + for group in self.worker_gpu_map.split(";") + if group.strip() + ] + if len(groups) != self.num_workers: + raise ValueError(f"worker_gpu_map defines {len(groups)} worker groups, but num_workers={self.num_workers}") + return groups diff --git a/telefuser/service/livekit/data_protocol.py b/telefuser/service/livekit/data_protocol.py new file mode 100644 index 0000000..e4ba032 --- /dev/null +++ b/telefuser/service/livekit/data_protocol.py @@ -0,0 +1,149 @@ +"""LiveKit data topic and message validation.""" + +from __future__ import annotations + +import json +from typing import Any + +TF_CONTROL_TOPIC = "tf.control" +TF_STATUS_TOPIC = "tf.status" +TF_METRICS_TOPIC = "tf.metrics" +TF_ASSET_TOPIC = "tf.asset" + +KNOWN_CONTROL_TYPES = frozenset({"control_state", "control", "prompt", "reset", "stop"}) +KNOWN_CONTROLS = frozenset( + { + "ArrowUp", + "ArrowDown", + "ArrowLeft", + "ArrowRight", + "KeyW", + "KeyA", + "KeyS", + "KeyD", + "KeyI", + "KeyJ", + "KeyK", + "KeyL", + "w", + "a", + "s", + "d", + "i", + "j", + "k", + "l", + "up", + "down", + "left", + "right", + "forward", + "backward", + } +) +KNOWN_CONTROL_EVENTS = frozenset({"press", "release", "keyup", "end", "reset", "reset_pose"}) +MEDIA_KEYS = frozenset({"frames", "frames_b64", "audio_b64", "audio_sample_rate", "audio_channels"}) + + +class DataProtocolError(ValueError): + """Raised when a LiveKit data message violates the TeleFuser protocol.""" + + +def normalize_control_message( + message: bytes | str | dict[str, Any], + *, + topic: str, + session_id: str, + sender_identity: str, + controller_identity: str, + max_bytes: int = 12 * 1024, +) -> dict[str, Any]: + """Validate and normalize one client-to-worker control message.""" + if topic != TF_CONTROL_TOPIC: + raise DataProtocolError(f"Unsupported data topic: {topic}") + if sender_identity and sender_identity != controller_identity: + raise DataProtocolError("Only the controller may send control messages") + + decoded = _decode_json_message(message, max_bytes=max_bytes) + if "version" in decoded: + return _normalize_enveloped_message(decoded, session_id=session_id) + return _normalize_legacy_message(decoded) + + +def strip_media_fields(chunk: dict[str, Any]) -> dict[str, Any]: + """Return chunk metadata without video/audio payload fields.""" + top_level_media_keys = MEDIA_KEYS if isinstance(chunk.get("frames"), (list, tuple)) else MEDIA_KEYS - {"frames"} + metadata = {key: value for key, value in chunk.items() if key not in top_level_media_keys} + data = metadata.get("data") + if isinstance(data, dict): + nested_media_keys = MEDIA_KEYS if isinstance(data.get("frames"), (list, tuple)) else MEDIA_KEYS - {"frames"} + nested = {key: value for key, value in data.items() if key not in nested_media_keys} + if nested: + metadata["data"] = nested + else: + metadata.pop("data", None) + return metadata + + +def _decode_json_message(message: bytes | str | dict[str, Any], *, max_bytes: int) -> dict[str, Any]: + if isinstance(message, dict): + return dict(message) + if isinstance(message, bytes): + raw_size = len(message) + raw = message.decode("utf-8") + else: + raw = message + raw_size = len(raw.encode("utf-8")) + + if raw_size > max_bytes: + raise DataProtocolError(f"Data message exceeds {max_bytes} bytes") + + try: + decoded = json.loads(raw) + except json.JSONDecodeError as exc: + raise DataProtocolError("Data message must be JSON") from exc + if not isinstance(decoded, dict): + raise DataProtocolError("Data message must decode to a JSON object") + return decoded + + +def _normalize_enveloped_message(message: dict[str, Any], *, session_id: str) -> dict[str, Any]: + if message.get("version") != 1: + raise DataProtocolError("Unsupported data protocol version") + msg_session_id = message.get("session_id") + if msg_session_id is not None and msg_session_id != session_id: + raise DataProtocolError("Control message session_id does not match active session") + + msg_type = message.get("type") + if msg_type not in KNOWN_CONTROL_TYPES: + raise DataProtocolError(f"Unsupported control message type: {msg_type}") + + payload = message.get("payload") or {} + if not isinstance(payload, dict): + raise DataProtocolError("Control message payload must be an object") + normalized = {"type": msg_type} + normalized.update(payload) + return _normalize_legacy_message(normalized) + + +def _normalize_legacy_message(message: dict[str, Any]) -> dict[str, Any]: + msg_type = message.get("type") + if msg_type not in KNOWN_CONTROL_TYPES: + raise DataProtocolError(f"Unsupported control message type: {msg_type}") + + if msg_type == "control_state": + controls = message.get("controls") + if not isinstance(controls, list) or not all(isinstance(control, str) for control in controls): + raise DataProtocolError("control_state controls must be a list of strings") + if any(control not in KNOWN_CONTROLS for control in controls): + raise DataProtocolError("control_state contains an unsupported control") + if len(controls) != len(set(controls)): + raise DataProtocolError("control_state controls must not contain duplicates") + elif msg_type == "control": + control = message.get("control", message.get("key")) + event = str(message.get("event") or message.get("action") or "press").lower() + if control not in KNOWN_CONTROLS: + raise DataProtocolError(f"Unsupported control: {control}") + if event not in KNOWN_CONTROL_EVENTS: + raise DataProtocolError(f"Unsupported control event: {event}") + return dict(message) diff --git a/telefuser/service/livekit/main.py b/telefuser/service/livekit/main.py new file mode 100644 index 0000000..4470cad --- /dev/null +++ b/telefuser/service/livekit/main.py @@ -0,0 +1,93 @@ +"""Entrypoint helpers for ``telefuser stream-serve``.""" + +from __future__ import annotations + +import asyncio +import sys +from typing import Any + +import uvicorn + +from telefuser._logo import TELEFUSER_LOGO +from telefuser.utils.logging import logger + +from .app import create_livekit_app +from .config import LiveKitServeConfig +from .runtime import LiveKitServeRuntime + + +def run_stream_server( + *, + pipe_path: str, + host: str | None = None, + port: int | None = None, + livekit_url: str | None = None, + livekit_api_key: str | None = None, + livekit_api_secret: str | None = None, + num_workers: int | None = None, + worker_gpu_map: str | None = None, + queue_size: int | None = None, + session_timeout: int | None = None, + token_ttl: int | None = None, + controller_timeout: int | None = None, + room_empty_timeout: int | None = None, + worker_mode: str | None = None, + skip_validation: bool = False, + security_level: str | None = None, +) -> None: + """Run the LiveKit-backed streaming HTTP API.""" + config_kwargs: dict[str, Any] = _drop_none( + { + "host": host, + "port": port, + "livekit_url": livekit_url, + "livekit_api_key": livekit_api_key, + "livekit_api_secret": livekit_api_secret, + "num_workers": num_workers, + "worker_gpu_map": worker_gpu_map, + "queue_size": queue_size, + "session_timeout": session_timeout, + "token_ttl": token_ttl, + "controller_timeout": controller_timeout, + "room_empty_timeout": room_empty_timeout, + "worker_mode": worker_mode, + } + ) + config = LiveKitServeConfig(**config_kwargs) + config.require_livekit_credentials() + + runtime = LiveKitServeRuntime( + config=config, + pipeline_file=pipe_path, + skip_validation=skip_validation, + security_level=security_level, + ) + app = create_livekit_app(runtime) + + try: + print(TELEFUSER_LOGO) + logger.info( + "Starting TeleFuser LiveKit server on %s:%s, workers=%s, pipeline=%s, " + "skip_validation=%s, security_level=%s", + config.host, + config.port, + config.num_workers, + pipe_path, + skip_validation, + security_level, + ) + uvicorn.run(app, host=config.host, port=config.port, log_level="warning") + except KeyboardInterrupt: + logger.info("LiveKit server interrupted by user") + except Exception as exc: + logger.error(f"LiveKit server failed: {exc}") + sys.exit(1) + finally: + try: + asyncio.run(runtime.aclose()) + except Exception as exc: + logger.warning(f"Error during LiveKit server cleanup: {exc}") + + +def _drop_none(values: dict[str, Any]) -> dict[str, Any]: + return {key: value for key, value in values.items() if value is not None} diff --git a/telefuser/service/livekit/media_bridge.py b/telefuser/service/livekit/media_bridge.py new file mode 100644 index 0000000..e96879f --- /dev/null +++ b/telefuser/service/livekit/media_bridge.py @@ -0,0 +1,86 @@ +"""Utilities for adapting TeleFuser chunks to LiveKit media publishing.""" + +from __future__ import annotations + +import base64 +from dataclasses import dataclass + +import av +import cv2 +import numpy as np +from PIL import Image + +from .data_protocol import strip_media_fields + + +class MediaDecodeError(ValueError): + """Raised when chunk media payloads cannot be decoded.""" + + +@dataclass(frozen=True) +class AudioPayload: + """Decoded PCM16 audio carried by one stream chunk.""" + + pcm: bytes + sample_rate: int + channels: int + + +def frame_to_rgb(source: object) -> np.ndarray: + """Convert one native pipeline frame to contiguous RGB24 pixels.""" + if isinstance(source, av.VideoFrame): + frame = source.to_ndarray(format="rgb24") + elif isinstance(source, Image.Image): + frame = np.asarray(source.convert("RGB")) + elif isinstance(source, np.ndarray): + frame = source + else: + raise MediaDecodeError(f"Unsupported native frame type: {type(source).__name__}") + + if frame.ndim != 3 or frame.shape[2] != 3: + raise MediaDecodeError(f"Native frame must have shape HxWx3, got {frame.shape}") + if frame.shape[0] <= 0 or frame.shape[1] <= 0: + raise MediaDecodeError("Native frame dimensions must be positive") + if frame.dtype != np.uint8: + raise MediaDecodeError(f"Native frame dtype must be uint8, got {frame.dtype}") + return np.ascontiguousarray(frame) + + +def decode_jpeg_frames(frames_b64: list[str]) -> list[np.ndarray]: + """Decode base64 JPEG frames into RGB numpy arrays.""" + frames: list[np.ndarray] = [] + for frame_b64 in frames_b64: + raw = base64.b64decode(frame_b64) + np_arr = np.frombuffer(raw, dtype=np.uint8) + bgr = cv2.imdecode(np_arr, cv2.IMREAD_COLOR) + if bgr is None: + raise MediaDecodeError("Chunk contains an undecodable JPEG frame") + frames.append(cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)) + return frames + + +def split_chunk_media(chunk: dict) -> tuple[list[np.ndarray], AudioPayload | None, dict]: + """Split a TeleFuser stream chunk into decoded video, audio bytes, and metadata.""" + has_top_level_frames = "frames" in chunk or "frames_b64" in chunk + data = chunk if has_top_level_frames or not isinstance(chunk.get("data"), dict) else chunk["data"] + + raw_frames = data.get("frames") + if isinstance(raw_frames, (list, tuple)): + frames = [frame_to_rgb(frame) for frame in raw_frames] + else: + frames = decode_jpeg_frames(list(data.get("frames_b64") or [])) + + audio_b64 = data.get("audio_b64") + audio = None + if audio_b64: + pcm = base64.b64decode(audio_b64) + sample_rate = int(data.get("audio_sample_rate") or 48_000) + channels = int(data.get("audio_channels") or 1) + if sample_rate <= 0 or channels <= 0: + raise MediaDecodeError("Audio sample rate and channel count must be positive") + bytes_per_sample = channels * 2 + if len(pcm) % bytes_per_sample: + raise MediaDecodeError("PCM16 audio byte length must align with its channel count") + audio = AudioPayload(pcm=pcm, sample_rate=sample_rate, channels=channels) + metadata = strip_media_fields(chunk) + return frames, audio, metadata diff --git a/telefuser/service/livekit/pipeline_adapter.py b/telefuser/service/livekit/pipeline_adapter.py new file mode 100644 index 0000000..1a6d4eb --- /dev/null +++ b/telefuser/service/livekit/pipeline_adapter.py @@ -0,0 +1,48 @@ +"""Adapter from LiveKit workers to TeleFuser stream pipeline services.""" + +from __future__ import annotations + +from collections.abc import AsyncGenerator + +from telefuser.service.core.config import ServerConfig +from telefuser.service.core.stream_pipeline_service import StreamPipelineService +from telefuser.service.security.security_validator import SecurityLevel + + +class LiveKitPipelineAdapter: + """Thin wrapper around ``StreamPipelineService`` for LiveKit workers.""" + + def __init__(self, *, security_level: SecurityLevel | None = None, config: ServerConfig | None = None) -> None: + self.stream_service = StreamPipelineService(security_level=security_level, config=config) + + def start(self, pipeline_file: str, *, skip_validation: bool = False, gpu_num: int = 1) -> None: + """Load and start a stream pipeline.""" + if not self.stream_service.start_service(pipeline_file, skip_validation=skip_validation, gpu_num=gpu_num): + raise RuntimeError(f"Failed to start LiveKit stream pipeline: {pipeline_file}") + + @property + def stream_mode(self) -> str | None: + """Return the detected TeleFuser stream interaction mode.""" + return self.stream_service.stream_mode + + async def aclose(self) -> None: + """Stop the wrapped stream service.""" + await self.stream_service.aclose() + + def create_session(self, config: dict) -> str: + return self.stream_service.create_session(config) + + def push_chunk(self, session_id: str, chunk: dict) -> None: + self.stream_service.push_chunk(session_id, chunk) + + async def pull_chunks(self, session_id: str) -> AsyncGenerator[dict, None]: + async for chunk in self.stream_service.pull_chunks(session_id): + yield chunk + + async def stream_task(self, config: dict) -> AsyncGenerator[dict, None]: + """Yield chunks from a server-push service.""" + async for chunk in self.stream_service.stream_task(config): + yield chunk + + def close_session(self, session_id: str) -> None: + self.stream_service.close_session(session_id) diff --git a/telefuser/service/livekit/room_client.py b/telefuser/service/livekit/room_client.py new file mode 100644 index 0000000..df7554a --- /dev/null +++ b/telefuser/service/livekit/room_client.py @@ -0,0 +1,186 @@ +"""LiveKit room client abstraction and SDK-backed implementation.""" + +from __future__ import annotations + +import json +from collections.abc import Callable +from typing import Any, Protocol + +import numpy as np + +from .data_protocol import TF_METRICS_TOPIC, TF_STATUS_TOPIC +from .token_service import LiveKitDependencyError + +DataMessageHandler = Callable[[bytes | str | dict[str, Any], str, str], None] +_VIDEO_MAX_BITRATE = 3_000_000 + + +class RoomClient(Protocol): + """Minimal room operations required by a TeleFuser LiveKit worker.""" + + async def connect(self, url: str, token: str, on_data: DataMessageHandler) -> None: ... + async def publish_video_track(self, name: str, width: int, height: int, *, fps: float = 16.0) -> None: ... + async def publish_video_frame(self, frame_rgb: np.ndarray, *, fps: float = 16.0) -> None: ... + async def publish_audio_frame(self, pcm: bytes, *, sample_rate: int, channels: int) -> None: ... + async def publish_status(self, payload: dict[str, Any]) -> None: ... + async def publish_metrics(self, payload: dict[str, Any]) -> None: ... + async def disconnect(self) -> None: ... + + +class LiveKitRoomClient: + """SDK-backed LiveKit room client. + + Imports LiveKit lazily so configuration and API schemas remain importable + before the runtime worker starts. + """ + + def __init__(self) -> None: + self._rtc: Any | None = None + self._room: Any | None = None + self._video_source: Any | None = None + self._video_track: Any | None = None + self._video_track_sid: str | None = None + self._video_dimensions: tuple[int, int] | None = None + self._audio_source: Any | None = None + self._audio_track: Any | None = None + self._audio_track_sid: str | None = None + self._audio_format: tuple[int, int] | None = None + + async def connect(self, url: str, token: str, on_data: DataMessageHandler) -> None: + rtc = self._load_rtc() + self._rtc = rtc + room = rtc.Room() + self._room = room + + @room.on("data_received") + def _on_data_received(packet: Any) -> None: + participant = getattr(packet, "participant", None) + identity = getattr(participant, "identity", "") if participant is not None else "" + on_data(packet.data, packet.topic or "", identity) + + await room.connect(url, token) + + async def publish_video_track(self, name: str, width: int, height: int, *, fps: float = 16.0) -> None: + room = self._require_room() + rtc = self._require_rtc() + if self._video_source is not None: + return + + self._video_source = rtc.VideoSource(width, height) + self._video_track = rtc.LocalVideoTrack.create_video_track(name, self._video_source) + options = rtc.TrackPublishOptions( + source=rtc.TrackSource.SOURCE_CAMERA, + simulcast=False, + video_encoding=rtc.VideoEncoding( + max_framerate=fps, + max_bitrate=_VIDEO_MAX_BITRATE, + ), + video_codec=rtc.VideoCodec.VP8, + ) + publication = await room.local_participant.publish_track(self._video_track, options) + self._video_track_sid = getattr(publication, "sid", None) + self._video_dimensions = (width, height) + + async def publish_video_frame(self, frame_rgb: np.ndarray, *, fps: float = 16.0) -> None: + if self._video_source is None: + height, width = frame_rgb.shape[:2] + await self.publish_video_track("telefuser-output", width, height, fps=fps) + + rtc = self._require_rtc() + height, width = frame_rgb.shape[:2] + if self._video_dimensions != (width, height): + raise ValueError(f"LiveKit video dimensions changed from {self._video_dimensions} to {(width, height)}") + + contiguous = np.ascontiguousarray(frame_rgb) + frame = rtc.VideoFrame(width, height, rtc.VideoBufferType.RGB24, contiguous.tobytes()) + self._video_source.capture_frame(frame) + + async def publish_audio_frame(self, pcm: bytes, *, sample_rate: int, channels: int) -> None: + """Publish one PCM16 audio payload through a LiveKit audio source.""" + if sample_rate <= 0 or channels <= 0: + raise ValueError("LiveKit audio sample rate and channel count must be positive") + bytes_per_sample = channels * 2 + if not pcm or len(pcm) % bytes_per_sample: + raise ValueError("LiveKit PCM16 byte length must align with its channel count") + + room = self._require_room() + rtc = self._require_rtc() + audio_format = (sample_rate, channels) + if self._audio_source is None: + self._audio_source = rtc.AudioSource(sample_rate, channels) + self._audio_track = rtc.LocalAudioTrack.create_audio_track("telefuser-audio", self._audio_source) + options = rtc.TrackPublishOptions(source=rtc.TrackSource.SOURCE_MICROPHONE) + publication = await room.local_participant.publish_track(self._audio_track, options) + self._audio_track_sid = getattr(publication, "sid", None) + self._audio_format = audio_format + elif self._audio_format != audio_format: + raise ValueError(f"LiveKit audio format changed from {self._audio_format} to {audio_format}") + + samples_per_channel = len(pcm) // bytes_per_sample + frame = rtc.AudioFrame(pcm, sample_rate, channels, samples_per_channel) + await self._audio_source.capture_frame(frame) + + async def publish_status(self, payload: dict[str, Any]) -> None: + await self._publish_data(payload, topic=TF_STATUS_TOPIC, reliable=True) + + async def publish_metrics(self, payload: dict[str, Any]) -> None: + await self._publish_data(payload, topic=TF_METRICS_TOPIC, reliable=False) + + async def disconnect(self) -> None: + room = self._room + if room is None: + return + + try: + if self._video_track_sid: + try: + await room.local_participant.unpublish_track(self._video_track_sid) + except Exception: + pass + if self._audio_track_sid: + try: + await room.local_participant.unpublish_track(self._audio_track_sid) + except Exception: + pass + if self._video_source is not None: + await self._video_source.aclose() + if self._audio_source is not None: + await self._audio_source.aclose() + finally: + try: + await room.disconnect() + finally: + self._room = None + self._video_source = None + self._video_track = None + self._video_track_sid = None + self._video_dimensions = None + self._audio_source = None + self._audio_track = None + self._audio_track_sid = None + self._audio_format = None + + async def _publish_data(self, payload: dict[str, Any], *, topic: str, reliable: bool) -> None: + room = self._require_room() + await room.local_participant.publish_data(json.dumps(payload).encode("utf-8"), topic=topic, reliable=reliable) + + @staticmethod + def _load_rtc() -> Any: + try: + from livekit import rtc + except ModuleNotFoundError as exc: + raise LiveKitDependencyError( + "LiveKit RTC SDK is required for LiveKit worker connections. " + "Install the declared TeleFuser runtime dependencies." + ) from exc + return rtc + + def _require_room(self) -> Any: + if self._room is None: + raise RuntimeError("LiveKit room is not connected") + return self._room + + def _require_rtc(self) -> Any: + if self._rtc is None: + raise RuntimeError("LiveKit RTC SDK is not loaded") + return self._rtc diff --git a/telefuser/service/livekit/runtime.py b/telefuser/service/livekit/runtime.py new file mode 100644 index 0000000..f53e764 --- /dev/null +++ b/telefuser/service/livekit/runtime.py @@ -0,0 +1,295 @@ +"""Runtime coordinator for LiveKit-backed ``telefuser stream-serve``.""" + +from __future__ import annotations + +import threading +from dataclasses import dataclass + +from telefuser.service.security.security_validator import SecurityLevel + +from .config import LiveKitServeConfig +from .pipeline_adapter import LiveKitPipelineAdapter +from .scheduler import LiveKitScheduler, SchedulerAdmission +from .schemas import ( + LiveKitHealthResponse, + SessionCreateRequest, + SessionStatus, + SessionStatusResponse, + SessionTokenRequest, + new_session_id, +) +from .session_registry import TERMINAL_SESSION_STATUSES, SessionRecord, SessionRegistry +from .token_service import LiveKitTokenService +from .worker import LiveKitWorker +from .worker_pool import InProcessLiveKitWorkerPool, WorkerPool + + +@dataclass(frozen=True) +class CreateSessionResult: + """Internal result for a session creation request.""" + + record: SessionRecord + token: str + admission: SchedulerAdmission + + +class LiveKitServeRuntime: + """Coordinates LiveKit token/session APIs and TeleFuser worker capacity.""" + + def __init__( + self, + *, + config: LiveKitServeConfig, + pipeline_file: str, + token_service: LiveKitTokenService | None = None, + registry: SessionRegistry | None = None, + scheduler: LiveKitScheduler | None = None, + worker_pool: WorkerPool | None = None, + skip_validation: bool = False, + security_level: SecurityLevel | str | None = None, + ) -> None: + self.config = config + self.pipeline_file = pipeline_file + self.registry = registry or SessionRegistry() + self.scheduler = scheduler or LiveKitScheduler( + num_workers=config.num_workers, + gpu_groups=config.worker_gpu_groups(), + queue_size=config.queue_size, + ) + self.token_service = token_service or LiveKitTokenService( + api_key=config.livekit_api_key, + api_secret=config.livekit_api_secret, + token_ttl=config.token_ttl, + ) + self.skip_validation = skip_validation + self.security_level = security_level + self.worker_pool = worker_pool or self._create_worker_pool() + self._started = False + self._closing = False + self._closed = False + self._finished_sessions: set[str] = set() + self._lock = threading.RLock() + + @property + def is_ready(self) -> bool: + """Return whether workers are started and able to accept sessions.""" + return self._started and not self._closing and self.health().status != "unhealthy" + + async def start(self) -> None: + """Start runtime-owned workers and load their pipelines.""" + with self._lock: + if self._started: + return + if self._closed: + raise RuntimeError("LiveKit runtime is already closed") + if self.config.worker_mode != "in-process": + raise NotImplementedError("stream-serve currently supports only worker_mode='in-process'") + if self.config.num_workers != 1: + raise NotImplementedError("stream-serve currently supports exactly one in-process worker") + await self.worker_pool.start(skip_validation=self.skip_validation) + with self._lock: + self._started = True + + def create_session(self, request: SessionCreateRequest) -> CreateSessionResult: + """Create a session record, mint a controller token, and reserve capacity.""" + + session_id = new_session_id() + room_name = f"tf-world-{session_id}" + session_config = dict(request.config) + session_config["session_id"] = session_id + if request.prompt is not None: + session_config["prompt"] = request.prompt + if request.image_path is not None: + session_config["image_path"] = request.image_path + + record = self.registry.create( + session_id=session_id, + room_name=room_name, + controller_identity=request.identity, + config=session_config, + timeout_s=self.config.session_timeout, + ) + admission = self.scheduler.assign(session_id=session_id, room_name=room_name) + if admission.status == "rejected": + self.registry.delete(session_id) + return CreateSessionResult(record=record, token="", admission=admission) + if admission.status == "assigned" and admission.worker_id is not None: + record = self.registry.assign_worker(session_id, admission.worker_id) + else: + record = self.registry.update_status(session_id, "queued") + + try: + token = self.token_service.create_token( + identity=request.identity, + room_name=room_name, + role="controller", + ) + except Exception: + self.scheduler.release_session(session_id) + self.registry.delete(session_id) + raise + + if admission.status == "assigned": + try: + self.worker_pool.start_session(record) + except Exception as exc: + self.scheduler.release_session(session_id) + self.registry.fail(session_id, str(exc)) + raise + return CreateSessionResult(record=record, token=token, admission=admission) + + def create_viewer_token(self, session_id: str, request: SessionTokenRequest) -> tuple[SessionRecord, str]: + """Mint a subscribe-only viewer token for an existing session.""" + record = self.registry.require(session_id) + if record.status in TERMINAL_SESSION_STATUSES: + raise ValueError(f"Session {session_id} is not active") + token = self.token_service.create_token( + identity=request.identity, + room_name=record.room_name, + role="viewer", + ) + return record, token + + def get_session_response(self, session_id: str) -> SessionStatusResponse: + """Return public status for one session.""" + return session_record_to_response(self.registry.require(session_id)) + + async def delete_session(self, session_id: str) -> SessionRecord: + """Stop a session, close its room worker, and release capacity.""" + record = self.registry.require(session_id) + if record.status in TERMINAL_SESSION_STATUSES: + return record + self.registry.update_status(session_id, "draining") + await self.worker_pool.stop_session(session_id) + return self._finish_session(session_id) + + def on_worker_status(self, worker_id: str, status: str) -> None: + """Apply a worker lifecycle callback to scheduler state.""" + self.scheduler.update_worker_status(worker_id, status) + + def on_session_status(self, session_id: str, status: SessionStatus, error: str | None = None) -> None: + """Apply a worker-reported public session state.""" + if self.registry.require(session_id).status not in TERMINAL_SESSION_STATUSES: + self.registry.update_status(session_id, status, error=error) + + def on_pipeline_session(self, session_id: str, pipeline_session_id: str) -> None: + """Record the pipeline session created by a worker.""" + self.registry.set_pipeline_session(session_id, pipeline_session_id) + + def on_session_finished(self, worker_id: str, session_id: str, error: str | None = None) -> None: + """Release capacity after a worker session exits.""" + del worker_id + self._finish_session(session_id, error=error) + + def health(self) -> LiveKitHealthResponse: + """Return service health based on current scheduler state.""" + snapshot = self.scheduler.health_snapshot() + workers_total = snapshot["workers_total"] + workers_failed = snapshot["workers_failed"] + status = "healthy" + if workers_total and workers_failed == workers_total: + status = "unhealthy" + elif workers_failed: + status = "degraded" + running_statuses = {"joining_room", "starting_pipeline", "running", "draining"} + return LiveKitHealthResponse( + status=status, + livekit_connected=any(worker.status in running_statuses for worker in self.scheduler.workers()), + **snapshot, + ) + + def metadata(self) -> dict: + """Return runtime metadata for `/v1/service/metadata`.""" + health = self.health() + return { + "service_type": "stream", + "transport": "livekit", + "pipeline_file": self.pipeline_file, + "livekit_url": self.config.livekit_url, + "num_workers": self.config.num_workers, + "worker_mode": self.config.worker_mode, + "queue_size": self.config.queue_size, + **health.model_dump(), + } + + async def aclose(self) -> None: + """Stop runtime-owned background resources.""" + with self._lock: + if self._closed or self._closing: + return + self._closing = True + try: + await self.worker_pool.aclose() + for record in self.registry.list_records(): + if record.status not in TERMINAL_SESSION_STATUSES: + self._finish_session(record.session_id, error="runtime closed") + finally: + with self._lock: + self._started = False + self._closing = False + self._closed = True + + def _create_worker_pool(self) -> WorkerPool: + security_level = self.security_level + if isinstance(security_level, str): + security_level = SecurityLevel[security_level.upper()] + workers: dict[str, LiveKitWorker] = {} + for worker_state in self.scheduler.workers(): + workers[worker_state.worker_id] = LiveKitWorker( + worker_id=worker_state.worker_id, + config=self.config, + pipeline_file=self.pipeline_file, + token_service=self.token_service, + event_sink=self, + pipeline_adapter=LiveKitPipelineAdapter(security_level=security_level), + gpu_num=max(1, len(worker_state.gpu_ids)), + ) + return InProcessLiveKitWorkerPool(workers) + + def _finish_session(self, session_id: str, *, error: str | None = None) -> SessionRecord: + with self._lock: + current = self.registry.require(session_id) + if session_id in self._finished_sessions: + return current + self._finished_sessions.add(session_id) + + if current.status in TERMINAL_SESSION_STATUSES: + record = current + elif error is not None and error != "cancelled": + record = self.registry.fail(session_id, error) + else: + record = self.registry.close(session_id) + admission = self.scheduler.release_session(session_id) + if admission is not None and not self._closing: + self._start_queued_session(admission) + return record + + def _start_queued_session(self, admission: SchedulerAdmission) -> None: + if admission.worker_id is None: + return + worker_state = next(worker for worker in self.scheduler.workers() if worker.worker_id == admission.worker_id) + if worker_state.session_id is None: + return + session_id = worker_state.session_id + try: + record = self.registry.assign_worker(session_id, admission.worker_id) + self.worker_pool.start_session(record) + except Exception as exc: + self.registry.fail(session_id, str(exc)) + self._finish_session(session_id, error=str(exc)) + + +def session_record_to_response(record: SessionRecord) -> SessionStatusResponse: + """Convert an internal session record to public response schema.""" + return SessionStatusResponse( + session_id=record.session_id, + room=record.room_name, + status=record.status, + worker_id=record.worker_id, + pipeline_session_id=record.pipeline_session_id, + created_at=record.created_at, + updated_at=record.updated_at, + expires_at=record.expires_at, + participant_count=record.participant_count, + error=record.error, + ) diff --git a/telefuser/service/livekit/scheduler.py b/telefuser/service/livekit/scheduler.py new file mode 100644 index 0000000..775393b --- /dev/null +++ b/telefuser/service/livekit/scheduler.py @@ -0,0 +1,176 @@ +"""Worker admission control for LiveKit serving.""" + +from __future__ import annotations + +import threading +from collections import deque +from typing import Literal + +from pydantic import BaseModel, Field + +from .schemas import utc_timestamp + +WorkerStatus = Literal[ + "starting", + "idle", + "assigned", + "joining_room", + "starting_pipeline", + "running", + "draining", + "failed", + "stopped", +] +AdmissionStatus = Literal["assigned", "queued", "rejected"] + + +class WorkerState(BaseModel): + """Scheduler-visible state for one model worker.""" + + worker_id: str + status: WorkerStatus + gpu_ids: list[str] = Field(default_factory=list) + session_id: str | None = None + room_name: str | None = None + last_heartbeat_at: float + error: str | None = None + + +class SchedulerAdmission(BaseModel): + """Result of a scheduler admission attempt.""" + + status: AdmissionStatus + worker_id: str | None = None + queue_position: int | None = None + reason: str | None = None + + +class _QueuedSession(BaseModel): + session_id: str + room_name: str + + +class LiveKitScheduler: + """Simple FIFO scheduler with one active session per worker.""" + + def __init__(self, *, num_workers: int, gpu_groups: list[list[str]] | None = None, queue_size: int = 0) -> None: + if num_workers < 1: + raise ValueError("num_workers must be >= 1") + if queue_size < 0: + raise ValueError("queue_size must be >= 0") + + now = utc_timestamp() + groups = gpu_groups or [[] for _ in range(num_workers)] + if len(groups) != num_workers: + raise ValueError(f"Expected {num_workers} GPU groups, got {len(groups)}") + + self._workers: dict[str, WorkerState] = { + f"worker-{idx}": WorkerState( + worker_id=f"worker-{idx}", + status="idle", + gpu_ids=list(groups[idx]), + last_heartbeat_at=now, + ) + for idx in range(num_workers) + } + self._queue_size = queue_size + self._queue: deque[_QueuedSession] = deque() + self._lock = threading.RLock() + + def assign(self, *, session_id: str, room_name: str) -> SchedulerAdmission: + """Assign an idle worker or enqueue/reject the session.""" + with self._lock: + worker = self._first_idle_worker() + if worker is not None: + worker.status = "assigned" + worker.session_id = session_id + worker.room_name = room_name + worker.error = None + worker.last_heartbeat_at = utc_timestamp() + return SchedulerAdmission(status="assigned", worker_id=worker.worker_id) + + if self._queue_size == 0: + return SchedulerAdmission(status="rejected", reason="no_idle_worker") + if len(self._queue) >= self._queue_size: + return SchedulerAdmission(status="rejected", reason="queue_full") + + self._queue.append(_QueuedSession(session_id=session_id, room_name=room_name)) + return SchedulerAdmission(status="queued", queue_position=len(self._queue)) + + def release_session(self, session_id: str) -> SchedulerAdmission | None: + """Release the worker owning ``session_id`` and assign the next queued session if present.""" + with self._lock: + worker = self._worker_for_session(session_id) + if worker is None: + self._queue = deque(item for item in self._queue if item.session_id != session_id) + return None + + worker.session_id = None + worker.room_name = None + worker.error = None + worker.status = "idle" + worker.last_heartbeat_at = utc_timestamp() + + if not self._queue: + return None + + queued = self._queue.popleft() + worker.status = "assigned" + worker.session_id = queued.session_id + worker.room_name = queued.room_name + worker.last_heartbeat_at = utc_timestamp() + return SchedulerAdmission(status="assigned", worker_id=worker.worker_id) + + def update_worker_status(self, worker_id: str, status: WorkerStatus) -> WorkerState: + """Update a worker lifecycle status.""" + with self._lock: + worker = self._workers[worker_id] + worker.status = status + worker.last_heartbeat_at = utc_timestamp() + return worker.model_copy(deep=True) + + def heartbeat(self, worker_id: str) -> WorkerState: + """Record a worker heartbeat.""" + with self._lock: + worker = self._workers[worker_id] + worker.last_heartbeat_at = utc_timestamp() + return worker.model_copy(deep=True) + + def fail_worker(self, worker_id: str, error: str) -> WorkerState: + """Mark a worker failed.""" + with self._lock: + worker = self._workers[worker_id] + worker.status = "failed" + worker.error = error + worker.last_heartbeat_at = utc_timestamp() + return worker.model_copy(deep=True) + + def workers(self) -> list[WorkerState]: + """Return copies of all workers.""" + with self._lock: + return [worker.model_copy(deep=True) for worker in self._workers.values()] + + def health_snapshot(self) -> dict[str, int]: + """Return scheduler capacity counts.""" + with self._lock: + workers = list(self._workers.values()) + busy_statuses = {"assigned", "joining_room", "starting_pipeline", "running", "draining"} + return { + "workers_total": len(workers), + "workers_idle": sum(1 for worker in workers if worker.status == "idle"), + "workers_busy": sum(1 for worker in workers if worker.status in busy_statuses), + "workers_failed": sum(1 for worker in workers if worker.status == "failed"), + "queued_sessions": len(self._queue), + } + + def _first_idle_worker(self) -> WorkerState | None: + for worker in self._workers.values(): + if worker.status == "idle": + return worker + return None + + def _worker_for_session(self, session_id: str) -> WorkerState | None: + for worker in self._workers.values(): + if worker.session_id == session_id: + return worker + return None diff --git a/telefuser/service/livekit/schemas.py b/telefuser/service/livekit/schemas.py new file mode 100644 index 0000000..56c3f14 --- /dev/null +++ b/telefuser/service/livekit/schemas.py @@ -0,0 +1,110 @@ +"""Pydantic schemas for the LiveKit serving API.""" + +from __future__ import annotations + +import time +import uuid +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field + +LiveKitClientRole = Literal["controller", "viewer", "admin"] +LiveKitTokenRole = Literal["controller", "viewer", "admin", "worker"] +SessionStatus = Literal[ + "pending", + "queued", + "assigned", + "joining_room", + "starting_pipeline", + "running", + "draining", + "closed", + "failed", + "expired", +] + + +def new_session_id() -> str: + """Generate a public TeleFuser LiveKit session id.""" + return str(uuid.uuid4()) + + +def utc_timestamp() -> float: + """Return current Unix timestamp in seconds.""" + return time.time() + + +class SessionCreateRequest(BaseModel): + """Request body for creating a LiveKit-backed stream session.""" + + model_config = ConfigDict(extra="allow") + + identity: str = Field(min_length=1) + role: Literal["controller"] = "controller" + prompt: str | None = None + image_path: str | None = None + config: dict = Field(default_factory=dict) + + +class SessionTokenRequest(BaseModel): + """Request body for minting an additional room token.""" + + identity: str = Field(min_length=1) + role: Literal["viewer"] = "viewer" + + +class SessionCreateResponse(BaseModel): + """Response body for creating a LiveKit-backed stream session.""" + + session_id: str + room: str + livekit_url: str + token: str + worker_id: str | None + status: SessionStatus + expires_at: float | None = None + queue_position: int | None = None + + +class SessionTokenResponse(BaseModel): + """Response body for minting an additional room token.""" + + session_id: str + room: str + livekit_url: str + token: str + role: Literal["viewer"] + + +class SessionStatusResponse(BaseModel): + """Public session status.""" + + session_id: str + room: str + status: SessionStatus + worker_id: str | None = None + pipeline_session_id: str | None = None + created_at: float + updated_at: float + expires_at: float | None = None + participant_count: int = 0 + error: str | None = None + + +class SessionDeleteResponse(BaseModel): + """Response body for deleting a session.""" + + session_id: str + status: SessionStatus + + +class LiveKitHealthResponse(BaseModel): + """Liveness and scheduler health for the LiveKit service.""" + + status: Literal["healthy", "degraded", "unhealthy"] + livekit_connected: bool + workers_total: int + workers_idle: int + workers_busy: int + workers_failed: int + queued_sessions: int diff --git a/telefuser/service/livekit/session_registry.py b/telefuser/service/livekit/session_registry.py new file mode 100644 index 0000000..020f9c9 --- /dev/null +++ b/telefuser/service/livekit/session_registry.py @@ -0,0 +1,126 @@ +"""In-memory session registry for LiveKit serving.""" + +from __future__ import annotations + +import threading + +from pydantic import BaseModel, Field + +from .schemas import SessionStatus, new_session_id, utc_timestamp + +TERMINAL_SESSION_STATUSES: frozenset[str] = frozenset({"closed", "failed", "expired"}) + + +class SessionRecord(BaseModel): + """Authoritative API-side state for one LiveKit-backed TeleFuser session.""" + + session_id: str + room_name: str + controller_identity: str + status: SessionStatus + worker_id: str | None = None + pipeline_session_id: str | None = None + config: dict = Field(default_factory=dict) + error: str | None = None + created_at: float + updated_at: float + expires_at: float | None = None + participant_count: int = 0 + + +class SessionRegistry: + """Thread-safe in-memory registry for session records.""" + + def __init__(self) -> None: + self._records: dict[str, SessionRecord] = {} + self._lock = threading.RLock() + + def create( + self, + *, + controller_identity: str, + config: dict, + session_id: str | None = None, + room_name: str | None = None, + timeout_s: int | None = None, + ) -> SessionRecord: + """Create and store a new pending session.""" + now = utc_timestamp() + sid = session_id or new_session_id() + record = SessionRecord( + session_id=sid, + room_name=room_name or f"tf-world-{sid}", + controller_identity=controller_identity, + status="pending", + config=dict(config), + created_at=now, + updated_at=now, + expires_at=now + timeout_s if timeout_s else None, + ) + with self._lock: + if sid in self._records: + raise ValueError(f"Session {sid} already exists") + self._records[sid] = record + return record.model_copy(deep=True) + + def get(self, session_id: str) -> SessionRecord | None: + """Return a session record copy if present.""" + with self._lock: + record = self._records.get(session_id) + return record.model_copy(deep=True) if record is not None else None + + def require(self, session_id: str) -> SessionRecord: + """Return a session record or raise ``KeyError``.""" + record = self.get(session_id) + if record is None: + raise KeyError(f"Session {session_id} not found") + return record + + def assign_worker(self, session_id: str, worker_id: str) -> SessionRecord: + """Mark a session as assigned to a worker.""" + return self.update(session_id, status="assigned", worker_id=worker_id) + + def set_pipeline_session(self, session_id: str, pipeline_session_id: str) -> SessionRecord: + """Record the pipeline-owned session id returned by ``create_session``.""" + return self.update(session_id, pipeline_session_id=pipeline_session_id) + + def update_status(self, session_id: str, status: SessionStatus, *, error: str | None = None) -> SessionRecord: + """Update only the public session status and optional error.""" + return self.update(session_id, status=status, error=error) + + def set_participant_count(self, session_id: str, participant_count: int) -> SessionRecord: + """Update the current room participant count.""" + return self.update(session_id, participant_count=participant_count) + + def close(self, session_id: str) -> SessionRecord: + """Mark a session closed.""" + return self.update_status(session_id, "closed") + + def fail(self, session_id: str, error: str) -> SessionRecord: + """Mark a session failed.""" + return self.update_status(session_id, "failed", error=error) + + def expire(self, session_id: str) -> SessionRecord: + """Mark a session expired.""" + return self.update_status(session_id, "expired") + + def update(self, session_id: str, **updates: object) -> SessionRecord: + """Apply field updates and return a copy of the new record.""" + with self._lock: + record = self._records.get(session_id) + if record is None: + raise KeyError(f"Session {session_id} not found") + for key, value in updates.items(): + setattr(record, key, value) + record.updated_at = utc_timestamp() + return record.model_copy(deep=True) + + def delete(self, session_id: str) -> bool: + """Delete a session record.""" + with self._lock: + return self._records.pop(session_id, None) is not None + + def list_records(self) -> list[SessionRecord]: + """Return copies of all known sessions.""" + with self._lock: + return [record.model_copy(deep=True) for record in self._records.values()] diff --git a/telefuser/service/livekit/token_service.py b/telefuser/service/livekit/token_service.py new file mode 100644 index 0000000..49f33a5 --- /dev/null +++ b/telefuser/service/livekit/token_service.py @@ -0,0 +1,72 @@ +"""LiveKit token generation.""" + +from __future__ import annotations + +import datetime +from typing import Any + +from .schemas import LiveKitTokenRole + + +class LiveKitDependencyError(RuntimeError): + """Raised when LiveKit SDK dependencies are unavailable.""" + + +class LiveKitTokenService: + """Generate LiveKit JWTs with TeleFuser role-specific grants.""" + + def __init__(self, *, api_key: str, api_secret: str, token_ttl: int) -> None: + self.api_key = api_key + self.api_secret = api_secret + self.token_ttl = token_ttl + + def create_token( + self, + *, + identity: str, + room_name: str, + role: LiveKitTokenRole, + name: str | None = None, + metadata: str | None = None, + ) -> str: + """Create a LiveKit join token for a scoped TeleFuser role.""" + api = self._load_livekit_api() + grants = self._create_video_grants(api, room_name=room_name, role=role) + + token = api.AccessToken(self.api_key, self.api_secret) + token = token.with_identity(identity) + token = token.with_name(name or identity) + token = token.with_grants(grants) + token = token.with_ttl(datetime.timedelta(seconds=self.token_ttl)) + if metadata is not None: + token = token.with_metadata(metadata) + return token.to_jwt() + + @staticmethod + def _load_livekit_api() -> Any: + try: + from livekit import api + except ModuleNotFoundError as exc: + raise LiveKitDependencyError( + "LiveKit Python SDK is required for `telefuser stream-serve`. " + "Install the declared TeleFuser runtime dependencies before starting this service." + ) from exc + return api + + @staticmethod + def _create_video_grants(api: Any, *, room_name: str, role: LiveKitTokenRole) -> Any: + kwargs: dict[str, Any] = { + "room_join": True, + "room": room_name, + } + if role == "controller": + kwargs.update(can_publish=False, can_publish_data=True, can_subscribe=True) + elif role == "viewer": + kwargs.update(can_publish=False, can_publish_data=False, can_subscribe=True) + elif role == "worker": + kwargs.update(can_publish=True, can_publish_data=True, can_subscribe=False) + elif role == "admin": + kwargs.update(can_publish=True, can_publish_data=True, can_subscribe=True, room_admin=True) + else: + raise ValueError(f"Unsupported LiveKit role: {role}") + return api.VideoGrants(**kwargs) diff --git a/telefuser/service/livekit/worker.py b/telefuser/service/livekit/worker.py new file mode 100644 index 0000000..2379100 --- /dev/null +++ b/telefuser/service/livekit/worker.py @@ -0,0 +1,270 @@ +"""LiveKit worker lifecycle for TeleFuser stream pipelines.""" + +from __future__ import annotations + +import asyncio +import contextlib +import time +from collections.abc import AsyncGenerator +from typing import Any, Protocol + +from telefuser.service.api.stream_schema import StreamChunkMessage, StreamDoneMessage, serialisable_chunk +from telefuser.service.core.stream_pipeline_service import STREAM_MODE_BIDIRECTIONAL, STREAM_MODE_SERVER_PUSH +from telefuser.utils.logging import logger + +from .config import LiveKitServeConfig +from .data_protocol import TF_CONTROL_TOPIC, normalize_control_message +from .media_bridge import split_chunk_media +from .pipeline_adapter import LiveKitPipelineAdapter +from .room_client import LiveKitRoomClient, RoomClient +from .schemas import SessionStatus +from .session_registry import SessionRecord +from .token_service import LiveKitTokenService + +_ROOM_DISCONNECT_TIMEOUT_SECONDS = 5.0 + + +class WorkerEventSink(Protocol): + """Callbacks used by a worker to report lifecycle changes.""" + + def on_worker_status(self, worker_id: str, status: str) -> None: ... + def on_session_status(self, session_id: str, status: SessionStatus, error: str | None = None) -> None: ... + def on_pipeline_session(self, session_id: str, pipeline_session_id: str) -> None: ... + def on_session_finished(self, worker_id: str, session_id: str, error: str | None = None) -> None: ... + + +class NullWorkerEventSink: + """No-op worker event sink for tests and isolated worker use.""" + + def on_worker_status(self, worker_id: str, status: str) -> None: + return None + + def on_session_status(self, session_id: str, status: SessionStatus, error: str | None = None) -> None: + return None + + def on_pipeline_session(self, session_id: str, pipeline_session_id: str) -> None: + return None + + def on_session_finished(self, worker_id: str, session_id: str, error: str | None = None) -> None: + return None + + +class LiveKitWorker: + """Owns one active LiveKit room and one active TeleFuser pipeline session.""" + + def __init__( + self, + *, + worker_id: str, + config: LiveKitServeConfig, + pipeline_file: str, + token_service: LiveKitTokenService, + event_sink: WorkerEventSink | None = None, + pipeline_adapter: LiveKitPipelineAdapter | None = None, + room_client: RoomClient | None = None, + gpu_num: int = 1, + ) -> None: + self.worker_id = worker_id + self.config = config + self.pipeline_file = pipeline_file + self.token_service = token_service + self.event_sink = event_sink or NullWorkerEventSink() + self.pipeline_adapter = pipeline_adapter or LiveKitPipelineAdapter() + self.room_client = room_client or LiveKitRoomClient() + self.gpu_num = gpu_num + self._active_session_id: str | None = None + self._pipeline_session_id: str | None = None + self._stop_event = asyncio.Event() + + async def start(self, *, skip_validation: bool = False) -> None: + """Load the stream pipeline owned by this worker.""" + self.pipeline_adapter.start(self.pipeline_file, skip_validation=skip_validation, gpu_num=self.gpu_num) + self.event_sink.on_worker_status(self.worker_id, "idle") + + async def stop(self) -> None: + """Stop the worker and close any active session.""" + self._stop_event.set() + await self._close_active_session() + await self.pipeline_adapter.aclose() + self.event_sink.on_worker_status(self.worker_id, "stopped") + + async def run_session(self, record: SessionRecord) -> None: + """Join a LiveKit room, create a pipeline session, and publish output chunks.""" + if self._active_session_id is not None: + raise RuntimeError(f"Worker {self.worker_id} is already running a session") + + self._active_session_id = record.session_id + self._stop_event.clear() + error: str | None = None + try: + self.event_sink.on_worker_status(self.worker_id, "joining_room") + self.event_sink.on_session_status(record.session_id, "joining_room") + worker_token = self.token_service.create_token( + identity=f"telefuser-{self.worker_id}", + room_name=record.room_name, + role="worker", + ) + await self.room_client.connect( + self.config.livekit_url, + worker_token, + lambda message, topic, identity: self._on_data_message(record, message, topic, identity), + ) + + self.event_sink.on_worker_status(self.worker_id, "starting_pipeline") + self.event_sink.on_session_status(record.session_id, "starting_pipeline") + if self.pipeline_adapter.stream_mode == STREAM_MODE_BIDIRECTIONAL: + self._pipeline_session_id = self.pipeline_adapter.create_session(record.config) + self.event_sink.on_pipeline_session(record.session_id, self._pipeline_session_id) + chunks = self.pipeline_adapter.pull_chunks(self._pipeline_session_id) + elif self.pipeline_adapter.stream_mode == STREAM_MODE_SERVER_PUSH: + chunks = self.pipeline_adapter.stream_task(record.config) + else: + raise RuntimeError(f"Unsupported stream mode: {self.pipeline_adapter.stream_mode}") + + self.event_sink.on_worker_status(self.worker_id, "running") + self.event_sink.on_session_status(record.session_id, "running") + await self.room_client.publish_status( + StreamChunkMessage( + session_id=record.session_id, + data={ + "type": "status", + "stage": "worker_running", + "worker_id": self.worker_id, + }, + ).model_dump(mode="json") + ) + await self._publish_pipeline_chunks(record.session_id, chunks) + except asyncio.CancelledError: + error = "cancelled" + raise + except Exception as exc: + error = str(exc) + logger.exception(f"LiveKit worker failed: worker={self.worker_id} session={record.session_id}") + self.event_sink.on_session_status(record.session_id, "failed", error=error) + with contextlib.suppress(Exception): + await self.room_client.publish_status( + StreamChunkMessage( + session_id=record.session_id, + error=error, + data={ + "type": "error", + "error": error, + }, + ).model_dump(mode="json") + ) + finally: + await self._close_active_session() + self.event_sink.on_session_finished(self.worker_id, record.session_id, error) + + async def stop_session(self, session_id: str) -> None: + """Request the active session to stop.""" + if self._active_session_id != session_id: + return + self._stop_event.set() + if self._pipeline_session_id is not None: + with contextlib.suppress(Exception): + self.pipeline_adapter.push_chunk(self._pipeline_session_id, {"type": "stop"}) + + def _on_data_message( + self, + record: SessionRecord, + message: bytes | str | dict[str, Any], + topic: str, + sender_identity: str, + ) -> None: + if self._pipeline_session_id is None: + return + try: + chunk = normalize_control_message( + message, + topic=topic, + session_id=record.session_id, + sender_identity=sender_identity, + controller_identity=record.controller_identity, + max_bytes=self.config.max_data_message_bytes, + ) + except Exception as exc: + logger.warning(f"LiveKit control message rejected: session={record.session_id} error={exc}") + return + self.pipeline_adapter.push_chunk(self._pipeline_session_id, chunk) + if chunk.get("type") == "stop": + self._stop_event.set() + + async def _publish_pipeline_chunks( + self, + session_id: str, + chunks: AsyncGenerator[dict, None], + ) -> None: + chunk_count = 0 + next_frame_at: float | None = None + async for chunk in chunks: + if self._stop_event.is_set(): + break + + frames, audio, metadata = split_chunk_media(chunk) + chunk_data = chunk.get("data") if isinstance(chunk.get("data"), dict) else chunk + fps_value = chunk_data.get("fps", chunk.get("fps", self.config.default_fps)) + try: + fps = float(fps_value) + except (TypeError, ValueError): + fps = float(self.config.default_fps) + if fps <= 0: + fps = float(self.config.default_fps) + frame_interval = 1.0 / fps + + for frame in frames: + if self._stop_event.is_set(): + break + now = time.monotonic() + if next_frame_at is None or now - next_frame_at > frame_interval: + next_frame_at = now + delay = next_frame_at - now + if delay > 0: + await asyncio.sleep(delay) + await self.room_client.publish_video_frame(frame, fps=fps) + next_frame_at += frame_interval + + if audio is not None: + await self.room_client.publish_audio_frame( + audio.pcm, + sample_rate=audio.sample_rate, + channels=audio.channels, + ) + + if self._stop_event.is_set(): + break + if frames: + chunk_count += 1 + if metadata: + index = chunk.get("index") + if index is None: + index = chunk_data.get("index") + message = StreamChunkMessage( + session_id=session_id, + index=index if isinstance(index, int) else None, + data=serialisable_chunk(metadata), + ) + await self.room_client.publish_status(message.model_dump(mode="json")) + + done = StreamDoneMessage(session_id=session_id, total_chunks=chunk_count).model_dump(mode="json") + await self.room_client.publish_status(done) + + async def _close_active_session(self) -> None: + pipeline_session_id = self._pipeline_session_id + self._pipeline_session_id = None + if pipeline_session_id is not None: + with contextlib.suppress(Exception): + await asyncio.to_thread(self.pipeline_adapter.close_session, pipeline_session_id) + try: + await asyncio.wait_for( + self.room_client.disconnect(), + timeout=_ROOM_DISCONNECT_TIMEOUT_SECONDS, + ) + except TimeoutError: + logger.warning( + f"LiveKit room disconnect timed out after {_ROOM_DISCONNECT_TIMEOUT_SECONDS:g}s: " + f"worker={self.worker_id}" + ) + except Exception as exc: + logger.warning(f"LiveKit room disconnect failed: worker={self.worker_id} error={exc}") + self._active_session_id = None diff --git a/telefuser/service/livekit/worker_pool.py b/telefuser/service/livekit/worker_pool.py new file mode 100644 index 0000000..9722f58 --- /dev/null +++ b/telefuser/service/livekit/worker_pool.py @@ -0,0 +1,84 @@ +"""Worker pool implementations for LiveKit serving.""" + +from __future__ import annotations + +import asyncio +from typing import Protocol + +from telefuser.utils.logging import logger + +from .session_registry import SessionRecord +from .worker import LiveKitWorker + + +class WorkerPool(Protocol): + """Worker-pool operations used by the API runtime.""" + + async def start(self, *, skip_validation: bool = False) -> None: ... + + def start_session(self, record: SessionRecord) -> None: ... + async def stop_session(self, session_id: str) -> None: ... + async def aclose(self) -> None: ... + + +class InProcessLiveKitWorkerPool: + """Run LiveKit workers as asyncio tasks in the API server process.""" + + def __init__(self, workers: dict[str, LiveKitWorker]) -> None: + self._workers = workers + self._started = False + self._tasks: dict[str, asyncio.Task] = {} + + async def start(self, *, skip_validation: bool = False) -> None: + """Load all worker-owned pipelines.""" + if self._started: + return + for worker in self._workers.values(): + await worker.start(skip_validation=skip_validation) + + self._started = True + + def start_session(self, record: SessionRecord) -> None: + """Start a worker task for an assigned session.""" + if not self._started: + raise RuntimeError("LiveKit worker pool is not started") + if record.worker_id is None: + raise RuntimeError(f"Session {record.session_id} has no assigned worker") + if record.session_id in self._tasks: + raise RuntimeError(f"Session {record.session_id} is already running") + + worker = self._workers[record.worker_id] + task = asyncio.create_task(worker.run_session(record), name=f"livekit-worker-{record.worker_id}") + self._tasks[record.session_id] = task + task.add_done_callback(lambda done: self._on_task_done(record.session_id, done)) + + async def stop_session(self, session_id: str) -> None: + """Request an active session to stop and wait for cleanup.""" + task = self._tasks.get(session_id) + if task is None: + return + for worker in self._workers.values(): + await worker.stop_session(session_id) + if not task.done(): + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + async def aclose(self) -> None: + """Stop every active session and worker.""" + session_ids = list(self._tasks.keys()) + self._started = False + for session_id in session_ids: + await self.stop_session(session_id) + for worker in self._workers.values(): + await worker.stop() + + def _on_task_done(self, session_id: str, task: asyncio.Task) -> None: + self._tasks.pop(session_id, None) + if task.cancelled(): + return + exc = task.exception() + if exc is not None: + logger.warning(f"LiveKit worker task failed: session={session_id} error={exc}") diff --git a/tests/unit/service/livekit/__init__.py b/tests/unit/service/livekit/__init__.py new file mode 100644 index 0000000..8995e72 --- /dev/null +++ b/tests/unit/service/livekit/__init__.py @@ -0,0 +1 @@ +"""Unit tests for TeleFuser LiveKit serving support.""" diff --git a/tests/unit/service/livekit/test_app.py b/tests/unit/service/livekit/test_app.py new file mode 100644 index 0000000..687313a --- /dev/null +++ b/tests/unit/service/livekit/test_app.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +from telefuser.service.livekit.app import create_livekit_app +from telefuser.service.livekit.config import LiveKitServeConfig +from telefuser.service.livekit.runtime import LiveKitServeRuntime +from tests.unit.openai._asgi_test_client import ASGITestClient + + +class FakeTokenService: + def create_token(self, *, identity: str, room_name: str, role: str, **kwargs: object) -> str: + return f"{role}:{identity}:{room_name}" + + +class FakeWorkerPool: + async def start(self, *, skip_validation: bool = False) -> None: + return None + + def start_session(self, record) -> None: + return None + + async def stop_session(self, session_id: str) -> None: + return None + + async def aclose(self) -> None: + return None + + +def _make_runtime(*, num_workers: int = 1, queue_size: int = 0) -> LiveKitServeRuntime: + config = LiveKitServeConfig( + livekit_url="wss://livekit.example", + livekit_api_key="key", + livekit_api_secret="secret", + num_workers=num_workers, + queue_size=queue_size, + ) + return LiveKitServeRuntime( + config=config, pipeline_file="pipeline.py", token_service=FakeTokenService(), worker_pool=FakeWorkerPool() + ) + + +def test_livekit_session_lifecycle_routes() -> None: + runtime = _make_runtime() + app = create_livekit_app(runtime) + + with ASGITestClient(app) as client: + create_resp = client.post( + "/v1/stream/sessions", + json={"identity": "controller-1", "config": {"fps": 16}}, + ) + assert create_resp.status_code == 200 + created = create_resp.json() + assert created["status"] == "assigned" + assert created["worker_id"] == "worker-0" + assert created["token"].startswith("controller:controller-1:") + + session_id = created["session_id"] + viewer_resp = client.post( + f"/v1/stream/sessions/{session_id}/tokens", + json={"identity": "viewer-1"}, + ) + assert viewer_resp.status_code == 200 + assert viewer_resp.json()["token"].startswith("viewer:viewer-1:") + + status_resp = client.get(f"/v1/stream/sessions/{session_id}") + assert status_resp.status_code == 200 + assert status_resp.json()["status"] == "assigned" + + delete_resp = client.delete(f"/v1/stream/sessions/{session_id}") + assert delete_resp.status_code == 200 + assert delete_resp.json() == {"session_id": session_id, "status": "closed"} + + +def test_livekit_session_queue_returns_accepted() -> None: + runtime = _make_runtime(queue_size=1) + app = create_livekit_app(runtime) + + with ASGITestClient(app) as client: + first_resp = client.post("/v1/stream/sessions", json={"identity": "controller-1"}) + second_resp = client.post("/v1/stream/sessions", json={"identity": "controller-2"}) + + assert first_resp.status_code == 200 + assert second_resp.status_code == 202 + assert second_resp.json()["status"] == "queued" + assert second_resp.json()["queue_position"] == 1 + + +def test_livekit_session_rejects_when_busy() -> None: + runtime = _make_runtime(queue_size=0) + app = create_livekit_app(runtime) + + with ASGITestClient(app) as client: + first_resp = client.post("/v1/stream/sessions", json={"identity": "controller-1"}) + second_resp = client.post("/v1/stream/sessions", json={"identity": "controller-2"}) + + assert first_resp.status_code == 200 + assert second_resp.status_code == 429 + + +def test_livekit_health_and_service_metadata_routes() -> None: + runtime = _make_runtime() + app = create_livekit_app(runtime) + + with ASGITestClient(app) as client: + health = client.get("/v1/stream/health") + metadata = client.get("/v1/service/metadata") + + assert health.status_code == 200 + assert health.json()["workers_total"] == 1 + assert metadata.status_code == 200 + assert metadata.json()["service_type"] == "stream" + assert metadata.json()["transport"] == "livekit" diff --git a/tests/unit/service/livekit/test_cli.py b/tests/unit/service/livekit/test_cli.py new file mode 100644 index 0000000..20c5066 --- /dev/null +++ b/tests/unit/service/livekit/test_cli.py @@ -0,0 +1,45 @@ +from __future__ import annotations + +from click.testing import CliRunner + +from telefuser.entrypoints.cli.main import main + + +def test_cli_stream_serve_forwards_livekit_options(monkeypatch) -> None: + captured = {} + + def fake_run_stream_server(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr("telefuser.service.livekit.main.run_stream_server", fake_run_stream_server) + + result = CliRunner().invoke( + main, + [ + "stream-serve", + "pipeline.py", + "--skip-validation", + "--livekit-url", + "wss://livekit.example", + "--livekit-api-key", + "key", + "--livekit-api-secret", + "secret", + "--num-workers", + "2", + "--worker-gpu-map", + "0;1", + "--queue-size", + "3", + ], + ) + + assert result.exit_code == 0 + assert captured["pipe_path"] == "pipeline.py" + assert captured["livekit_url"] == "wss://livekit.example" + assert captured["livekit_api_key"] == "key" + assert captured["livekit_api_secret"] == "secret" + assert captured["num_workers"] == 2 + assert captured["worker_gpu_map"] == "0;1" + assert captured["queue_size"] == 3 + assert captured["skip_validation"] is True diff --git a/tests/unit/service/livekit/test_config.py b/tests/unit/service/livekit/test_config.py new file mode 100644 index 0000000..e0f0627 --- /dev/null +++ b/tests/unit/service/livekit/test_config.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import pytest + +from telefuser.service.livekit.config import LiveKitServeConfig + + +def test_worker_gpu_groups_default_to_empty_groups() -> None: + config = LiveKitServeConfig(num_workers=2) + + assert config.worker_gpu_groups() == [[], []] + + +def test_worker_gpu_groups_parse_semicolon_map() -> None: + config = LiveKitServeConfig(num_workers=2, worker_gpu_map="0,1;2,3") + + assert config.worker_gpu_groups() == [["0", "1"], ["2", "3"]] + + +def test_worker_gpu_groups_reject_wrong_group_count() -> None: + config = LiveKitServeConfig(num_workers=2, worker_gpu_map="0,1") + + with pytest.raises(ValueError, match="worker groups"): + config.worker_gpu_groups() + + +def test_require_livekit_credentials_reports_missing_fields() -> None: + config = LiveKitServeConfig(livekit_url="wss://example.livekit.cloud") + + with pytest.raises(ValueError, match="livekit_api_key, livekit_api_secret"): + config.require_livekit_credentials() diff --git a/tests/unit/service/livekit/test_data_protocol.py b/tests/unit/service/livekit/test_data_protocol.py new file mode 100644 index 0000000..466163e --- /dev/null +++ b/tests/unit/service/livekit/test_data_protocol.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +import json + +import pytest + +from telefuser.service.livekit.data_protocol import ( + TF_CONTROL_TOPIC, + DataProtocolError, + normalize_control_message, + strip_media_fields, +) + + +def test_normalize_enveloped_control_message() -> None: + payload = { + "version": 1, + "type": "control", + "session_id": "session-1", + "payload": {"event": "press", "key": "ArrowUp"}, + } + + chunk = normalize_control_message( + json.dumps(payload), + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="controller", + controller_identity="controller", + ) + + assert chunk == {"type": "control", "event": "press", "key": "ArrowUp"} + + +def test_normalize_legacy_stop_message() -> None: + chunk = normalize_control_message( + {"type": "stop"}, + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="controller", + controller_identity="controller", + ) + + assert chunk == {"type": "stop"} + + +def test_rejects_viewer_control_message() -> None: + with pytest.raises(DataProtocolError, match="controller"): + normalize_control_message( + {"type": "stop"}, + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="viewer", + controller_identity="controller", + ) + + +def test_rejects_session_mismatch() -> None: + with pytest.raises(DataProtocolError, match="session_id"): + normalize_control_message( + {"version": 1, "type": "stop", "session_id": "other"}, + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="controller", + controller_identity="controller", + ) + + +def test_strip_media_fields_handles_nested_data() -> None: + metadata = strip_media_fields( + { + "type": "chunk", + "frames_b64": ["frame"], + "audio_b64": "audio", + "data": {"frames_b64": ["nested"], "stage": "ok"}, + "index": 1, + } + ) + + assert metadata == {"type": "chunk", "data": {"stage": "ok"}, "index": 1} + + +def test_strip_media_fields_preserves_numeric_frame_counts() -> None: + metadata = strip_media_fields( + { + "type": "status", + "stage": "chunk_decoded", + "frames": 13, + "data": {"stage": "nested", "frames": 7}, + } + ) + + assert metadata == { + "type": "status", + "stage": "chunk_decoded", + "frames": 13, + "data": {"stage": "nested", "frames": 7}, + } + + +@pytest.mark.parametrize( + "message", + [ + {"type": "control_state", "controls": ["a", "d", "i", "j", "k", "l", "s", "w"]}, + {"type": "control", "control": "up", "event": "reset"}, + {"type": "control", "control": "up", "event": "reset_pose"}, + ], +) +def test_normalize_livekit_demo_controls(message: dict) -> None: + chunk = normalize_control_message( + message, + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="controller", + controller_identity="controller", + ) + + assert chunk == message + + +@pytest.mark.parametrize( + "payload", + [ + {"controls": ["w", "w"]}, + {"controls": ["unsupported"]}, + ], +) +def test_rejects_invalid_enveloped_control_state(payload: dict) -> None: + with pytest.raises(DataProtocolError): + normalize_control_message( + {"version": 1, "type": "control_state", "payload": payload}, + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="controller", + controller_identity="controller", + ) + + +def test_accepts_control_when_livekit_omits_sender_participant() -> None: + chunk = normalize_control_message( + {"type": "control_state", "controls": ["w"]}, + topic=TF_CONTROL_TOPIC, + session_id="session-1", + sender_identity="", + controller_identity="controller", + ) + + assert chunk == {"type": "control_state", "controls": ["w"]} diff --git a/tests/unit/service/livekit/test_media_bridge.py b/tests/unit/service/livekit/test_media_bridge.py new file mode 100644 index 0000000..244fec8 --- /dev/null +++ b/tests/unit/service/livekit/test_media_bridge.py @@ -0,0 +1,98 @@ +from __future__ import annotations + +import base64 + +import cv2 +import numpy as np +import pytest +from PIL import Image + +from telefuser.service.livekit.media_bridge import MediaDecodeError, frame_to_rgb, split_chunk_media + + +def test_split_chunk_media_accepts_native_pil_frames() -> None: + image = Image.new("RGB", (4, 3), color=(10, 20, 30)) + + frames, audio, metadata = split_chunk_media( + { + "type": "chunk", + "index": 2, + "frames": [image], + "stream_progress": {"completed_chunks": 3}, + } + ) + + assert audio is None + assert len(frames) == 1 + assert frames[0].shape == (3, 4, 3) + assert frames[0][0, 0].tolist() == [10, 20, 30] + assert metadata == { + "type": "chunk", + "index": 2, + "stream_progress": {"completed_chunks": 3}, + } + + +def test_split_chunk_media_keeps_base64_jpeg_compatibility() -> None: + bgr = np.zeros((3, 4, 3), dtype=np.uint8) + bgr[:, :] = (30, 20, 10) + ok, encoded = cv2.imencode(".jpg", bgr) + assert ok + + frames, _, metadata = split_chunk_media( + { + "data": { + "frames_b64": [base64.b64encode(encoded.tobytes()).decode()], + "fps": 16, + }, + "index": 1, + } + ) + + assert frames[0].shape == (3, 4, 3) + assert metadata == {"data": {"fps": 16}, "index": 1} + + +def test_split_chunk_media_treats_numeric_frames_as_status_metadata() -> None: + frames, audio, metadata = split_chunk_media( + { + "type": "status", + "stage": "chunk_decoded", + "index": 0, + "frames": 13, + } + ) + + assert frames == [] + assert audio is None + assert metadata == { + "type": "status", + "stage": "chunk_decoded", + "index": 0, + "frames": 13, + } + + +def test_split_chunk_media_decodes_pcm16_audio_format() -> None: + pcm = np.zeros(960 * 2, dtype=np.int16).tobytes() + + frames, audio, metadata = split_chunk_media( + { + "type": "chunk", + "audio_b64": base64.b64encode(pcm).decode(), + "audio_sample_rate": 48_000, + "audio_channels": 2, + } + ) + + assert frames == [] + assert audio is not None + assert audio.pcm == pcm + assert audio.sample_rate == 48_000 + assert audio.channels == 2 + assert metadata == {"type": "chunk"} + + +def test_frame_to_rgb_rejects_wrong_pixel_shape() -> None: + with pytest.raises(MediaDecodeError, match="HxWx3"): + frame_to_rgb(np.zeros((3, 4), dtype=np.uint8)) diff --git a/tests/unit/service/livekit/test_room_client.py b/tests/unit/service/livekit/test_room_client.py new file mode 100644 index 0000000..2dd98a7 --- /dev/null +++ b/tests/unit/service/livekit/test_room_client.py @@ -0,0 +1,159 @@ +from __future__ import annotations + +import asyncio +import sys +import types + +import numpy as np +import pytest + +from telefuser.service.livekit.room_client import LiveKitRoomClient + + +def test_livekit_room_client_uses_sdk_room_publish_and_video_source(monkeypatch) -> None: + captured: dict[str, object] = {} + + class FakeVideoFrame: + def __init__(self, width: int, height: int, buffer_type: str, data: bytes) -> None: + captured["frame"] = { + "width": width, + "height": height, + "buffer_type": buffer_type, + "data": data, + } + + class FakeVideoSource: + def __init__(self, width: int, height: int) -> None: + captured["source"] = (width, height) + self.frames = [] + + def capture_frame(self, frame) -> None: + self.frames.append(frame) + captured["captured_frame"] = frame + + async def aclose(self) -> None: + captured["source_closed"] = True + + class FakeLocalVideoTrack: + @staticmethod + def create_video_track(name: str, source: FakeVideoSource): + captured["track"] = (name, source) + return object() + + class FakeAudioFrame: + def __init__( + self, + data: bytes, + sample_rate: int, + channels: int, + samples_per_channel: int, + ) -> None: + captured["audio_frame"] = (data, sample_rate, channels, samples_per_channel) + + class FakeAudioSource: + def __init__(self, sample_rate: int, channels: int) -> None: + captured["audio_source"] = (sample_rate, channels) + + async def capture_frame(self, frame) -> None: + captured["captured_audio_frame"] = frame + + async def aclose(self) -> None: + captured["audio_source_closed"] = True + + class FakeLocalAudioTrack: + @staticmethod + def create_audio_track(name: str, source: FakeAudioSource): + captured["audio_track"] = (name, source) + return object() + + class FakePublication: + sid = "track-sid" + + class FakeLocalParticipant: + async def publish_track(self, track, options): + captured["published_track"] = (track, options) + return FakePublication() + + async def publish_data(self, data: bytes, *, topic: str, reliable: bool) -> None: + captured["published_data"] = (data, topic, reliable) + + async def unpublish_track(self, sid: str) -> None: + captured["unpublished_track"] = sid + + class FakeRoom: + def __init__(self) -> None: + self.local_participant = FakeLocalParticipant() + self.handlers = {} + + def on(self, event: str): + def _decorator(fn): + self.handlers[event] = fn + return fn + + return _decorator + + async def connect(self, url: str, token: str) -> None: + captured["connect"] = (url, token) + + async def disconnect(self) -> None: + captured["disconnect"] = True + + fake_room = FakeRoom() + fake_rtc = types.SimpleNamespace( + Room=lambda: fake_room, + TrackPublishOptions=lambda **kwargs: types.SimpleNamespace(**kwargs), + TrackSource=types.SimpleNamespace(SOURCE_CAMERA="camera"), + VideoEncoding=lambda **kwargs: types.SimpleNamespace(**kwargs), + VideoCodec=types.SimpleNamespace(VP8="VP8"), + VideoSource=FakeVideoSource, + LocalVideoTrack=FakeLocalVideoTrack, + VideoFrame=FakeVideoFrame, + VideoBufferType=types.SimpleNamespace(RGB24="RGB24"), + AudioSource=FakeAudioSource, + LocalAudioTrack=FakeLocalAudioTrack, + AudioFrame=FakeAudioFrame, + ) + fake_rtc.TrackSource.SOURCE_MICROPHONE = "microphone" + monkeypatch.setitem(sys.modules, "livekit", types.SimpleNamespace(rtc=fake_rtc)) + + async def _run() -> None: + messages = [] + client = LiveKitRoomClient() + await client.connect( + "wss://livekit.example", "token", lambda data, topic, identity: messages.append((data, topic, identity)) + ) + + participant = types.SimpleNamespace(identity="controller") + packet = types.SimpleNamespace(data=b"{}", topic="tf.control", participant=participant) + fake_room.handlers["data_received"](packet) + + frame = np.zeros((2, 3, 3), dtype=np.uint8) + await client.publish_video_frame(frame, fps=16) + _video_track, video_options = captured["published_track"] + with pytest.raises(ValueError, match="dimensions changed"): + await client.publish_video_frame(np.zeros((4, 3, 3), dtype=np.uint8), fps=16) + pcm = np.zeros(960, dtype=np.int16).tobytes() + await client.publish_audio_frame(pcm, sample_rate=48_000, channels=1) + with pytest.raises(ValueError, match="audio format changed"): + await client.publish_audio_frame(pcm, sample_rate=24_000, channels=1) + await client.publish_status({"type": "status"}) + await client.disconnect() + + assert messages == [(b"{}", "tf.control", "controller")] + assert captured["connect"] == ("wss://livekit.example", "token") + assert captured["source"] == (3, 2) + assert video_options.simulcast is False + assert video_options.video_codec == "VP8" + assert video_options.video_encoding.max_framerate == 16 + assert video_options.video_encoding.max_bitrate == 3_000_000 + assert captured["frame"]["buffer_type"] == "RGB24" + assert captured["audio_source"] == (48_000, 1) + assert captured["audio_frame"] == (pcm, 48_000, 1, 960) + assert "captured_audio_frame" in captured + assert captured["published_data"] == (b'{"type": "status"}', "tf.status", True) + assert captured["unpublished_track"] == "track-sid" + assert captured["source_closed"] is True + assert captured["audio_source_closed"] is True + assert captured["disconnect"] is True + + asyncio.run(_run()) diff --git a/tests/unit/service/livekit/test_runtime.py b/tests/unit/service/livekit/test_runtime.py new file mode 100644 index 0000000..809179a --- /dev/null +++ b/tests/unit/service/livekit/test_runtime.py @@ -0,0 +1,135 @@ +from __future__ import annotations + +import asyncio + +from telefuser.service.livekit.config import LiveKitServeConfig +from telefuser.service.livekit.runtime import LiveKitServeRuntime +from telefuser.service.livekit.schemas import SessionCreateRequest + + +class FakeTokenService: + def create_token(self, *, identity: str, room_name: str, role: str, **kwargs: object) -> str: + return f"{role}:{identity}:{room_name}" + + +class FakeWorkerPool: + def __init__(self) -> None: + self.started: list[str] = [] + self.stopped: list[str] = [] + self.closed = False + self.start_options: list[bool] = [] + self.close_calls = 0 + + async def start(self, *, skip_validation: bool = False) -> None: + self.start_options.append(skip_validation) + + def start_session(self, record) -> None: + self.started.append(record.session_id) + + async def stop_session(self, session_id: str) -> None: + self.stopped.append(session_id) + + async def aclose(self) -> None: + self.closed = True + self.close_calls += 1 + + +def test_runtime_starts_queued_session_when_worker_is_released() -> None: + async def _run() -> None: + config = LiveKitServeConfig( + livekit_url="wss://livekit.example", + livekit_api_key="key", + livekit_api_secret="secret", + queue_size=1, + ) + worker_pool = FakeWorkerPool() + runtime = LiveKitServeRuntime( + config=config, + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + worker_pool=worker_pool, + ) + + first = runtime.create_session(SessionCreateRequest(identity="controller-1")) + second = runtime.create_session(SessionCreateRequest(identity="controller-2")) + await runtime.delete_session(first.record.session_id) + + assert first.record.status == "assigned" + assert second.record.status == "queued" + assert worker_pool.stopped == [first.record.session_id] + assert worker_pool.started == [first.record.session_id, second.record.session_id] + assert runtime.registry.require(first.record.session_id).status == "closed" + second_record = runtime.registry.require(second.record.session_id) + assert second_record.status == "assigned" + assert second_record.worker_id == "worker-0" + + asyncio.run(_run()) + + +def test_runtime_worker_callbacks_release_capacity() -> None: + config = LiveKitServeConfig(livekit_url="wss://livekit.example", livekit_api_key="key", livekit_api_secret="secret") + worker_pool = FakeWorkerPool() + runtime = LiveKitServeRuntime( + config=config, + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + worker_pool=worker_pool, + ) + result = runtime.create_session(SessionCreateRequest(identity="controller-1")) + + runtime.on_pipeline_session(result.record.session_id, "pipeline-1") + runtime.on_session_status(result.record.session_id, "running") + runtime.on_session_finished("worker-0", result.record.session_id) + + record = runtime.registry.require(result.record.session_id) + assert record.status == "closed" + assert record.pipeline_session_id == "pipeline-1" + assert runtime.scheduler.health_snapshot()["workers_idle"] == 1 + + +def test_runtime_start_and_close_are_idempotent() -> None: + async def _run() -> None: + config = LiveKitServeConfig( + livekit_url="wss://livekit.example", + livekit_api_key="key", + livekit_api_secret="secret", + ) + worker_pool = FakeWorkerPool() + runtime = LiveKitServeRuntime( + config=config, + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + worker_pool=worker_pool, + skip_validation=True, + ) + + await runtime.start() + await runtime.start() + assert runtime.is_ready is True + assert worker_pool.start_options == [True] + + await runtime.aclose() + await runtime.aclose() + assert runtime.is_ready is False + assert worker_pool.close_calls == 1 + + asyncio.run(_run()) + + +def test_runtime_releases_capacity_after_worker_reports_failure() -> None: + config = LiveKitServeConfig(livekit_url="wss://livekit.example", livekit_api_key="key", livekit_api_secret="secret") + runtime = LiveKitServeRuntime( + config=config, + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + worker_pool=FakeWorkerPool(), + ) + result = runtime.create_session(SessionCreateRequest(identity="controller-1")) + + runtime.on_session_status(result.record.session_id, "failed", error="room connect failed") + runtime.on_session_finished("worker-0", result.record.session_id, error="room connect failed") + + record = runtime.registry.require(result.record.session_id) + assert record.status == "failed" + assert record.error == "room connect failed" + assert runtime.scheduler.health_snapshot()["workers_idle"] == 1 diff --git a/tests/unit/service/livekit/test_scheduler.py b/tests/unit/service/livekit/test_scheduler.py new file mode 100644 index 0000000..d9b70b3 --- /dev/null +++ b/tests/unit/service/livekit/test_scheduler.py @@ -0,0 +1,54 @@ +from __future__ import annotations + +from telefuser.service.livekit.scheduler import LiveKitScheduler + + +def test_scheduler_assigns_first_idle_worker() -> None: + scheduler = LiveKitScheduler(num_workers=1, gpu_groups=[["0"]]) + + admission = scheduler.assign(session_id="session-1", room_name="room-1") + + assert admission.status == "assigned" + assert admission.worker_id == "worker-0" + worker = scheduler.workers()[0] + assert worker.session_id == "session-1" + assert worker.gpu_ids == ["0"] + + +def test_scheduler_rejects_when_busy_and_queue_disabled() -> None: + scheduler = LiveKitScheduler(num_workers=1, queue_size=0) + scheduler.assign(session_id="session-1", room_name="room-1") + + admission = scheduler.assign(session_id="session-2", room_name="room-2") + + assert admission.status == "rejected" + assert admission.reason == "no_idle_worker" + + +def test_scheduler_queues_and_assigns_on_release() -> None: + scheduler = LiveKitScheduler(num_workers=1, queue_size=1) + scheduler.assign(session_id="session-1", room_name="room-1") + + queued = scheduler.assign(session_id="session-2", room_name="room-2") + next_admission = scheduler.release_session("session-1") + + assert queued.status == "queued" + assert queued.queue_position == 1 + assert next_admission is not None + assert next_admission.status == "assigned" + worker = scheduler.workers()[0] + assert worker.session_id == "session-2" + + +def test_scheduler_health_counts_failed_workers() -> None: + scheduler = LiveKitScheduler(num_workers=2) + scheduler.assign(session_id="session-1", room_name="room-1") + scheduler.fail_worker("worker-1", "boom") + + assert scheduler.health_snapshot() == { + "workers_total": 2, + "workers_idle": 0, + "workers_busy": 1, + "workers_failed": 1, + "queued_sessions": 0, + } diff --git a/tests/unit/service/livekit/test_session_registry.py b/tests/unit/service/livekit/test_session_registry.py new file mode 100644 index 0000000..66b17fb --- /dev/null +++ b/tests/unit/service/livekit/test_session_registry.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from telefuser.service.livekit.session_registry import SessionRegistry + + +def test_session_registry_lifecycle() -> None: + registry = SessionRegistry() + + record = registry.create( + session_id="session-1", + room_name="room-1", + controller_identity="user-1", + config={"fps": 16}, + timeout_s=60, + ) + + assert record.status == "pending" + assert record.expires_at is not None + + assigned = registry.assign_worker("session-1", "worker-0") + assert assigned.status == "assigned" + assert assigned.worker_id == "worker-0" + + pipeline_record = registry.set_pipeline_session("session-1", "pipeline-1") + assert pipeline_record.pipeline_session_id == "pipeline-1" + + closed = registry.close("session-1") + assert closed.status == "closed" + + +def test_session_registry_returns_copies() -> None: + registry = SessionRegistry() + record = registry.create(controller_identity="user-1", config={}, session_id="session-1") + record.status = "failed" + + stored = registry.require("session-1") + + assert stored.status == "pending" diff --git a/tests/unit/service/livekit/test_token_service.py b/tests/unit/service/livekit/test_token_service.py new file mode 100644 index 0000000..e39bee9 --- /dev/null +++ b/tests/unit/service/livekit/test_token_service.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +import sys +import types + +from telefuser.service.livekit.token_service import LiveKitTokenService + + +def test_token_service_builds_viewer_grants(monkeypatch) -> None: + captured: dict[str, object] = {} + + class FakeVideoGrants: + def __init__(self, **kwargs) -> None: + captured["grants"] = kwargs + + class FakeAccessToken: + def __init__(self, api_key, api_secret) -> None: + captured["api_key"] = api_key + captured["api_secret"] = api_secret + + def with_identity(self, identity): + captured["identity"] = identity + return self + + def with_name(self, name): + captured["name"] = name + return self + + def with_grants(self, grants): + captured["grant_obj"] = grants + return self + + def with_ttl(self, ttl): + captured["ttl"] = ttl.total_seconds() + return self + + def to_jwt(self): + return "jwt-token" + + fake_api = types.SimpleNamespace(AccessToken=FakeAccessToken, VideoGrants=FakeVideoGrants) + monkeypatch.setitem(sys.modules, "livekit", types.SimpleNamespace(api=fake_api)) + + token = LiveKitTokenService(api_key="key", api_secret="secret", token_ttl=123).create_token( + identity="viewer-1", + room_name="room-1", + role="viewer", + ) + + assert token == "jwt-token" + assert captured["api_key"] == "key" + assert captured["api_secret"] == "secret" + assert captured["identity"] == "viewer-1" + assert captured["ttl"] == 123 + assert captured["grants"] == { + "room_join": True, + "room": "room-1", + "can_publish": False, + "can_publish_data": False, + "can_subscribe": True, + } + + +def test_token_service_builds_worker_grants(monkeypatch) -> None: + captured: dict[str, object] = {} + + class FakeVideoGrants: + def __init__(self, **kwargs) -> None: + captured["grants"] = kwargs + + class FakeAccessToken: + def __init__(self, api_key, api_secret) -> None: + return None + + def with_identity(self, identity): + return self + + def with_name(self, name): + return self + + def with_grants(self, grants): + return self + + def with_ttl(self, ttl): + return self + + def to_jwt(self): + return "worker-token" + + fake_api = types.SimpleNamespace(AccessToken=FakeAccessToken, VideoGrants=FakeVideoGrants) + monkeypatch.setitem(sys.modules, "livekit", types.SimpleNamespace(api=fake_api)) + + token = LiveKitTokenService(api_key="key", api_secret="secret", token_ttl=60).create_token( + identity="worker-0", + room_name="room-1", + role="worker", + ) + + assert token == "worker-token" + assert captured["grants"] == { + "room_join": True, + "room": "room-1", + "can_publish": True, + "can_publish_data": True, + "can_subscribe": False, + } diff --git a/tests/unit/service/livekit/test_worker.py b/tests/unit/service/livekit/test_worker.py new file mode 100644 index 0000000..acaea6d --- /dev/null +++ b/tests/unit/service/livekit/test_worker.py @@ -0,0 +1,302 @@ +from __future__ import annotations + +import asyncio +import base64 +import json + +import cv2 +import numpy as np +from PIL import Image + +from telefuser.service.core.stream_pipeline_service import STREAM_MODE_BIDIRECTIONAL, STREAM_MODE_SERVER_PUSH +from telefuser.service.livekit import worker as worker_module +from telefuser.service.livekit.config import LiveKitServeConfig +from telefuser.service.livekit.session_registry import SessionRecord +from telefuser.service.livekit.worker import LiveKitWorker + + +class FakeTokenService: + def create_token(self, *, identity: str, room_name: str, role: str, **kwargs: object) -> str: + return f"{role}:{identity}:{room_name}" + + +class FakePipelineAdapter: + def __init__(self, stream_mode: str = STREAM_MODE_BIDIRECTIONAL) -> None: + self.stream_mode = stream_mode + self.started: list[dict[str, object]] = [] + self.created_config: dict | None = None + self.served_config: dict | None = None + self.pushed: list[tuple[str, dict]] = [] + self.closed: list[str] = [] + self.closed_service = False + self.created = asyncio.Event() + self.output_queue: asyncio.Queue[dict | None] = asyncio.Queue() + + def start(self, pipeline_file: str, *, skip_validation: bool = False, gpu_num: int = 1) -> None: + self.started.append({"pipeline_file": pipeline_file, "skip_validation": skip_validation, "gpu_num": gpu_num}) + + async def aclose(self) -> None: + self.closed_service = True + + def create_session(self, config: dict) -> str: + self.created_config = config + self.created.set() + return "pipeline-session-1" + + def push_chunk(self, session_id: str, chunk: dict) -> None: + self.pushed.append((session_id, chunk)) + + async def pull_chunks(self, session_id: str): + while True: + item = await self.output_queue.get() + if item is None: + break + yield item + + async def stream_task(self, config: dict): + self.served_config = config + while True: + item = await self.output_queue.get() + if item is None: + break + yield item + + def close_session(self, session_id: str) -> None: + self.closed.append(session_id) + + +class FakeRoomClient: + def __init__(self) -> None: + self.connected = asyncio.Event() + self.connect_args: tuple[str, str] | None = None + self.on_data = None + self.video_frames: list[np.ndarray] = [] + self.video_frame_fps: list[float] = [] + self.audio_frames: list[tuple[bytes, int, int]] = [] + self.statuses: list[dict] = [] + self.disconnected = False + self.disconnect_gate: asyncio.Event | None = None + + async def connect(self, url: str, token: str, on_data) -> None: + self.connect_args = (url, token) + self.on_data = on_data + self.connected.set() + + async def publish_video_track(self, name: str, width: int, height: int, *, fps: float = 16.0) -> None: + return None + + async def publish_video_frame(self, frame_rgb: np.ndarray, *, fps: float = 16.0) -> None: + self.video_frames.append(frame_rgb) + self.video_frame_fps.append(fps) + + async def publish_audio_frame(self, pcm: bytes, *, sample_rate: int, channels: int) -> None: + self.audio_frames.append((pcm, sample_rate, channels)) + + async def publish_status(self, payload: dict) -> None: + self.statuses.append(payload) + + async def publish_metrics(self, payload: dict) -> None: + return None + + async def disconnect(self) -> None: + if self.disconnect_gate is not None: + await self.disconnect_gate.wait() + self.disconnected = True + + def emit_control(self, payload: dict, *, identity: str = "controller") -> None: + assert self.on_data is not None + self.on_data(json.dumps(payload), "tf.control", identity) + + +class FakeSink: + def __init__(self) -> None: + self.worker_statuses: list[tuple[str, str]] = [] + self.session_statuses: list[tuple[str, str, str | None]] = [] + self.pipeline_sessions: list[tuple[str, str]] = [] + self.finished: list[tuple[str, str, str | None]] = [] + + def on_worker_status(self, worker_id: str, status: str) -> None: + self.worker_statuses.append((worker_id, status)) + + def on_session_status(self, session_id: str, status: str, error: str | None = None) -> None: + self.session_statuses.append((session_id, status, error)) + + def on_pipeline_session(self, session_id: str, pipeline_session_id: str) -> None: + self.pipeline_sessions.append((session_id, pipeline_session_id)) + + def on_session_finished(self, worker_id: str, session_id: str, error: str | None = None) -> None: + self.finished.append((worker_id, session_id, error)) + + +def _jpeg_chunk() -> dict: + frame = np.zeros((8, 8, 3), dtype=np.uint8) + ok, encoded = cv2.imencode(".jpg", frame) + assert ok + return { + "type": "chunk", + "index": 0, + "fps": 16, + "frames_b64": [base64.b64encode(encoded.tobytes()).decode("ascii")], + } + + +def _native_chunk() -> dict: + return { + "type": "chunk", + "index": 1, + "fps": 16, + "frames": [Image.new("RGB", (8, 8), color=(1, 2, 3))], + "stream_progress": {"completed_chunks": 2}, + } + + +def _audio_chunk() -> dict: + pcm = np.zeros(960, dtype=np.int16).tobytes() + return { + "type": "chunk", + "index": 2, + "audio_b64": base64.b64encode(pcm).decode("ascii"), + "audio_sample_rate": 48_000, + "audio_channels": 1, + } + + +async def _wait_for(predicate, *, timeout: float = 1.0) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while not predicate(): + if asyncio.get_running_loop().time() > deadline: + raise AssertionError("timed out waiting for condition") + await asyncio.sleep(0.01) + + +def test_livekit_worker_runs_pipeline_and_forwards_control() -> None: + async def _run() -> None: + config = LiveKitServeConfig( + livekit_url="wss://livekit.example", livekit_api_key="key", livekit_api_secret="secret" + ) + adapter = FakePipelineAdapter() + room = FakeRoomClient() + sink = FakeSink() + worker = LiveKitWorker( + worker_id="worker-0", + config=config, + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + event_sink=sink, + pipeline_adapter=adapter, + room_client=room, + ) + record = SessionRecord( + session_id="session-1", + room_name="room-1", + controller_identity="controller", + status="assigned", + worker_id="worker-0", + config={"session_id": "session-1", "fps": 16}, + created_at=0, + updated_at=0, + ) + + await worker.start(skip_validation=True) + task = asyncio.create_task(worker.run_session(record)) + await room.connected.wait() + await adapter.created.wait() + room.emit_control({"type": "control", "event": "press", "key": "ArrowUp"}) + await _wait_for(lambda: len(adapter.pushed) == 1) + await adapter.output_queue.put({"type": "status", "stage": "chunk_decoded", "frames": 13}) + await adapter.output_queue.put(_jpeg_chunk()) + await adapter.output_queue.put(_native_chunk()) + await adapter.output_queue.put(_audio_chunk()) + await adapter.output_queue.put(None) + await task + + assert adapter.started == [{"pipeline_file": "pipeline.py", "skip_validation": True, "gpu_num": 1}] + assert adapter.created_config == {"session_id": "session-1", "fps": 16} + assert adapter.pushed == [("pipeline-session-1", {"type": "control", "event": "press", "key": "ArrowUp"})] + assert adapter.closed == ["pipeline-session-1"] + assert room.connect_args == ("wss://livekit.example", "worker:telefuser-worker-0:room-1") + assert len(room.video_frames) == 2 + assert room.video_frame_fps == [16.0, 16.0] + assert room.audio_frames == [(np.zeros(960, dtype=np.int16).tobytes(), 48_000, 1)] + assert any(status.get("data", {}).get("frames") == 13 for status in room.statuses) + assert any(status.get("data", {}).get("stream_progress") == {"completed_chunks": 2} for status in room.statuses) + assert room.statuses[-1]["type"] == "done" + assert room.disconnected is True + assert sink.pipeline_sessions == [("session-1", "pipeline-session-1")] + assert room.statuses[-1]["total_chunks"] == 2 + assert sink.finished == [("worker-0", "session-1", None)] + + asyncio.run(_run()) + + +def test_livekit_worker_runs_server_push_pipeline() -> None: + async def _run() -> None: + adapter = FakePipelineAdapter(stream_mode=STREAM_MODE_SERVER_PUSH) + room = FakeRoomClient() + sink = FakeSink() + worker = LiveKitWorker( + worker_id="worker-0", + config=LiveKitServeConfig( + livekit_url="wss://livekit.example", + livekit_api_key="key", + livekit_api_secret="secret", + ), + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + event_sink=sink, + pipeline_adapter=adapter, + room_client=room, + ) + record = SessionRecord( + session_id="session-1", + room_name="room-1", + controller_identity="controller", + status="assigned", + worker_id="worker-0", + config={"session_id": "session-1", "prompt": "sunset", "fps": 16}, + created_at=0, + updated_at=0, + ) + + await worker.start(skip_validation=True) + task = asyncio.create_task(worker.run_session(record)) + await room.connected.wait() + await adapter.output_queue.put(_native_chunk()) + await adapter.output_queue.put(None) + await task + + assert adapter.served_config == record.config + assert adapter.created_config is None + assert adapter.closed == [] + assert sink.pipeline_sessions == [] + assert len(room.video_frames) == 1 + assert room.statuses[-1]["type"] == "done" + + asyncio.run(_run()) + + +def test_livekit_worker_bounds_room_disconnect(monkeypatch) -> None: + async def _run() -> None: + room = FakeRoomClient() + room.disconnect_gate = asyncio.Event() + worker = LiveKitWorker( + worker_id="worker-0", + config=LiveKitServeConfig( + livekit_url="wss://livekit.example", + livekit_api_key="key", + livekit_api_secret="secret", + ), + pipeline_file="pipeline.py", + token_service=FakeTokenService(), + pipeline_adapter=FakePipelineAdapter(), + room_client=room, + ) + worker._active_session_id = "session-1" + monkeypatch.setattr(worker_module, "_ROOM_DISCONNECT_TIMEOUT_SECONDS", 0.01) + + await asyncio.wait_for(worker._close_active_session(), timeout=1) + + assert worker._active_session_id is None + assert room.disconnected is False + + asyncio.run(_run()) From 9583911aedb4280b203d5dedcf3a95d4ea86235d Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 05:44:31 +0000 Subject: [PATCH 02/11] fix(attention): stop loading SageAttention from tf-kernel Use the standalone sageattention package as the only SageAttention backend source, avoiding the tf-kernel path that triggered CUDA misaligned-address failures. Verification: PYTHONDONTWRITEBYTECODE=1 .venv/bin/python -m pytest tests/unit/ops/test_attention_backends.py tests/unit/service/livekit -q; ruff check and ruff format --check on the changed attention and LiveKit files; git diff --check. --- telefuser/ops/attention/backends.py | 16 +++++++--------- tests/unit/ops/test_attention_backends.py | 10 +++++----- 2 files changed, 12 insertions(+), 14 deletions(-) diff --git a/telefuser/ops/attention/backends.py b/telefuser/ops/attention/backends.py index 1d63f4f..8558ae1 100644 --- a/telefuser/ops/attention/backends.py +++ b/telefuser/ops/attention/backends.py @@ -88,15 +88,13 @@ def _try_import_sage_attn() -> None: """Import Sage Attention.""" global SAGE_ATTN_AVAILABLE, sageattention - for module_name in ["tf_kernel.sageattn2", "sageattention"]: - try: - if importlib.util.find_spec(module_name) is not None: - sageattention = importlib.import_module(module_name) - SAGE_ATTN_AVAILABLE = True - logger.debug(f"Sage Attention loaded from {module_name}") - return - except (ModuleNotFoundError, ImportError): - continue + try: + if importlib.util.find_spec("sageattention") is not None: + sageattention = importlib.import_module("sageattention") + SAGE_ATTN_AVAILABLE = True + logger.debug("Sage Attention available") + except (ModuleNotFoundError, ImportError): + pass def _try_import_sparge_attn() -> None: diff --git a/tests/unit/ops/test_attention_backends.py b/tests/unit/ops/test_attention_backends.py index 56155b5..7434621 100644 --- a/tests/unit/ops/test_attention_backends.py +++ b/tests/unit/ops/test_attention_backends.py @@ -4,15 +4,15 @@ from telefuser.ops.attention import backends -def test_sage_attention_prefers_tf_kernel() -> None: +def test_sage_attention_uses_standalone_package() -> None: imported_modules: list[str] = [] - tf_kernel_module = ModuleType("tf_kernel.sageattn2") + sageattention_module = ModuleType("sageattention") previous_available = backends.SAGE_ATTN_AVAILABLE previous_backend = backends.sageattention def import_module(name: str) -> ModuleType: imported_modules.append(name) - return tf_kernel_module + return sageattention_module try: backends.SAGE_ATTN_AVAILABLE = False @@ -23,9 +23,9 @@ def import_module(name: str) -> ModuleType: ): backends._try_import_sage_attn() - assert imported_modules == ["tf_kernel.sageattn2"] + assert imported_modules == ["sageattention"] assert backends.SAGE_ATTN_AVAILABLE is True - assert backends.sageattention is tf_kernel_module + assert backends.sageattention is sageattention_module finally: backends.SAGE_ATTN_AVAILABLE = previous_available backends.sageattention = previous_backend From d1fcdda471c321b3ca816e1fd6bb5171f89ad825 Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 10:45:42 +0000 Subject: [PATCH 03/11] feat(streaming): add LiveKit browser control demo Add the shared camera-control UI and LiveKit browser adapter with TURN relay configuration, reliable control messages, video playback, and streaming telemetry. Cover the rendered page and LingBot example defaults. Verification: - .venv/bin/python -m pytest tests/unit/service/livekit/test_demo.py tests/unit/pipelines/lingbot_world_fast/test_stream_example.py -q - 7 passed - git diff --cached --check --- examples/stream_server/_control_demo_ui.py | 604 ++++++++++++++++++ .../livekit_bidirectional_demo.py | 350 ++++++++++ .../lingbot_world_fast/test_stream_example.py | 37 +- tests/unit/service/livekit/test_demo.py | 39 ++ 4 files changed, 1015 insertions(+), 15 deletions(-) create mode 100644 examples/stream_server/_control_demo_ui.py create mode 100644 examples/stream_server/livekit_bidirectional_demo.py create mode 100644 tests/unit/service/livekit/test_demo.py diff --git a/examples/stream_server/_control_demo_ui.py b/examples/stream_server/_control_demo_ui.py new file mode 100644 index 0000000..fe6e553 --- /dev/null +++ b/examples/stream_server/_control_demo_ui.py @@ -0,0 +1,604 @@ +"""Shared browser layout and camera-control UI fragments for stream demos.""" + +from __future__ import annotations + +import json +from pathlib import Path + +DEFAULT_SERVER_URL = "http://localhost:8088" +DEFAULT_PORT = 8091 +_PROJECT_ROOT = Path(__file__).resolve().parents[2] +DEFAULT_IMAGE_PATH = str(_PROJECT_ROOT / "examples" / "data" / "lingbot_world_fast" / "image.jpg") +DEFAULT_PROMPT = ( + "A serene lakeside scene with a lone tree standing in calm water, surrounded by distant snow-capped " + "mountains under a bright blue sky with drifting white clouds. Gentle ripples reflect the tree and sky." +) + +HTML_TEMPLATE = r""" + + + +LingBot-World-Fast LiveKit Demo + + + +
+

LingBot-World-Fast LiveKit Demo

+
+
+
+

Server Output

+
Ready.
+
+ +
+
Server limit--
+
Target video--
+
Generated video--
+
Frames / chunks--
+
Output cadence--
+
Pipeline residence--
+
Applied control latency--
+
Queue / dropped video--
+
+
+ + +
+
+ + + +""" diff --git a/examples/stream_server/livekit_bidirectional_demo.py b/examples/stream_server/livekit_bidirectional_demo.py new file mode 100644 index 0000000..80d785a --- /dev/null +++ b/examples/stream_server/livekit_bidirectional_demo.py @@ -0,0 +1,350 @@ +"""LingBot-World-Fast LiveKit control demo. + +The page reuses the shared control UI asset to keep the prompt, image, controls, +and telemetry behavior consistent across interactive examples. + +Usage: + # 1. Start a LiveKit server and export its URL/key/secret. + # 2. Start TeleFuser: + telefuser stream-serve examples/lingbot/lingbot_world_fast_image_to_video_h100.py --skip-validation + # 3. Start this browser client: + python examples/stream_server/livekit_bidirectional_demo.py --server-url http://localhost:8088 +""" + +from __future__ import annotations + +import argparse +import functools +import http.server +import json +import runpy +import threading +import urllib.error +import urllib.request +import webbrowser +from pathlib import Path + +DEFAULT_SERVER_URL = "http://localhost:8088" +DEFAULT_PORT = 8092 +LIVEKIT_CLIENT_URL = "https://cdn.jsdelivr.net/npm/livekit-client@2.21.0/dist/livekit-client.umd.min.js" +_PROJECT_ROOT = Path(__file__).resolve().parents[2] +_CONTROL_DEMO_UI_PATH = _PROJECT_ROOT / "examples" / "stream_server" / "_control_demo_ui.py" + + +def _shared_demo_parts() -> tuple[str, str, str, str, str]: + namespace = runpy.run_path(str(_CONTROL_DEMO_UI_PATH)) + template = str(namespace["HTML_TEMPLATE"]) + script = template.split("", 1)[0] + shell = template.split("\n\n\n' + + +def main() -> None: + parser = argparse.ArgumentParser(description="LingBot-World-Fast LiveKit control demo") + parser.add_argument("--server-url", default=DEFAULT_SERVER_URL, help="LiveKit API server base URL") + parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Local HTTP server port") + parser.add_argument( + "--proxy-backend", + action=argparse.BooleanOptionalAction, + default=True, + help="Proxy /v1/stream/* via this demo server", + ) + parser.add_argument("--no-open", action="store_true", help="Do not open the browser automatically") + args = parser.parse_args() + + server_url_for_browser = "" if args.proxy_backend else args.server_url + html = _render_html(server_url_for_browser) + + class Handler(http.server.BaseHTTPRequestHandler): + _opener = urllib.request.build_opener(urllib.request.ProxyHandler({})) + + def _is_backend_request(self) -> bool: + return bool(args.proxy_backend) and self.path.startswith("/v1/stream/") + + def _proxy(self) -> None: + url = f"{args.server_url.rstrip('/')}{self.path}" + content_length = int(self.headers.get("Content-Length") or "0") + body = self.rfile.read(content_length) if content_length else None + headers = {"Content-Type": self.headers["Content-Type"]} if self.headers.get("Content-Type") else {} + request = urllib.request.Request(url, data=body, headers=headers, method=self.command) + try: + with self._opener.open(request, timeout=30) as response: + response_body = response.read() + response_status = getattr(response, "status", 200) + response_type = response.headers.get("Content-Type", "application/octet-stream") + except urllib.error.HTTPError as exc: + response_body = exc.read() + response_status = exc.code + response_type = exc.headers.get("Content-Type", "application/json") + except Exception as exc: + response_body = json.dumps({"detail": f"Demo proxy error: {exc}"}).encode() + response_status = 502 + response_type = "application/json" + self.send_response(response_status) + self.send_header("Content-Type", response_type) + self.send_header("Content-Length", str(len(response_body))) + self.end_headers() + self.wfile.write(response_body) + + def do_GET(self) -> None: + if self._is_backend_request(): + self._proxy() + return + if self.path == "/default-image": + body = Path(DEFAULT_IMAGE_PATH).read_bytes() + content_type = "image/jpeg" + else: + body = html.encode() + content_type = "text/html; charset=utf-8" + self.send_response(200) + self.send_header("Content-Type", content_type) + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_POST(self) -> None: + self._proxy() if self._is_backend_request() else self.send_error(404) + + def do_DELETE(self) -> None: + self._proxy() if self._is_backend_request() else self.send_error(404) + + def log_message(self, format: str, *_args: object) -> None: + return None + + server = http.server.ThreadingHTTPServer(("0.0.0.0", args.port), Handler) + url = f"http://localhost:{args.port}" + print(f"Serving LingBot-World-Fast LiveKit demo at {url}") + print(f"LiveKit API server: {args.server_url}") + print(f"LiveKit JS client: {LIVEKIT_CLIENT_URL}") + if args.proxy_backend: + print("Proxy: enabled for the TeleFuser API") + print("VS Code forwarding required: TCP 8092 (page), 7880 (LiveKit), and 3478 (TURN)") + print("Press Ctrl+C to stop.\n") + if not args.no_open: + threading.Timer(0.5, functools.partial(webbrowser.open, url)).start() + try: + server.serve_forever() + except KeyboardInterrupt: + print("\nStopped.") + finally: + server.server_close() + + +if __name__ == "__main__": + main() diff --git a/tests/unit/pipelines/lingbot_world_fast/test_stream_example.py b/tests/unit/pipelines/lingbot_world_fast/test_stream_example.py index 8f45eac..241da03 100644 --- a/tests/unit/pipelines/lingbot_world_fast/test_stream_example.py +++ b/tests/unit/pipelines/lingbot_world_fast/test_stream_example.py @@ -4,28 +4,35 @@ import torch from examples.lingbot import lingbot_world_fast_image_to_video_h100 as offline_example -from examples.stream_server import webrtc_bidirectional_demo as webrtc_demo +from examples.stream_server import livekit_bidirectional_demo as stream_demo from telefuser.core.config import AttnImplType -def test_webrtc_demo_uses_stream_service_defaults() -> None: - source = inspect.getsource(webrtc_demo) +def test_livekit_demo_uses_stream_service_defaults() -> None: + source = inspect.getsource(stream_demo) + html = stream_demo._render_html("") - assert webrtc_demo.DEFAULT_IMAGE_PATH == offline_example.DEFAULT_IMAGE_PATH - assert webrtc_demo.DEFAULT_PROMPT == offline_example.DEFAULT_PROMPT - assert source.count("HTML_TEMPLATE =") == 1 + assert stream_demo.DEFAULT_IMAGE_PATH == offline_example.DEFAULT_IMAGE_PATH + assert stream_demo.DEFAULT_PROMPT == offline_example.DEFAULT_PROMPT assert "DEFAULT_OPTIONS" not in source assert "--intrinsics-path" not in source assert "--image-path" not in source - assert 'type="file"' in webrtc_demo.HTML_TEMPLATE - assert 'id="image-preview" src="/default-image"' in webrtc_demo.HTML_TEMPLATE - assert '$("image-preview").src = imagePreviewObjectUrl' in webrtc_demo.HTML_TEMPLATE - assert "requestBody.image = image" in webrtc_demo.HTML_TEMPLATE - assert "requestBody.image_path = DEFAULT_IMAGE_PATH" in webrtc_demo.HTML_TEMPLATE - assert 'type: "control_state"' in webrtc_demo.HTML_TEMPLATE - assert 'window.addEventListener("blur", () => releaseAllControls(true))' in webrtc_demo.HTML_TEMPLATE - assert 'document.addEventListener("visibilitychange"' in webrtc_demo.HTML_TEMPLATE - assert 'id="reset-pose"' in webrtc_demo.HTML_TEMPLATE + assert 'type="file"' in html + assert 'id="image-preview" src="/default-image"' in html + assert '$("image-preview").src = imagePreviewObjectUrl' in html + assert "requestBody.config = image" not in html + assert "config: image ? { image } : {}" in html + assert "requestBody.image_path = DEFAULT_IMAGE_PATH" in html + assert 'type: "control_state"' in html + assert 'window.addEventListener("blur", () => releaseAllControls(true))' in html + assert 'document.addEventListener("visibilitychange"' in html + assert 'id="reset-pose"' in html + assert 'Output cadence--' in html + assert 'Pipeline residence--' in html + assert 'Applied control latency--' in html + assert "metrics.output_cadence_seconds" in html + assert "metrics.pipeline_residence_seconds ?? metrics.chunk_elapsed_seconds" in html + assert "metrics.applied_control_latency_seconds ?? metrics.control_to_chunk_seconds" in html def test_unified_example_get_pipeline_maps_ppl_config_to_internal_workers() -> None: diff --git a/tests/unit/service/livekit/test_demo.py b/tests/unit/service/livekit/test_demo.py new file mode 100644 index 0000000..aa35ac4 --- /dev/null +++ b/tests/unit/service/livekit/test_demo.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +import runpy +from pathlib import Path + + +def test_stream_demo_preserves_controls_and_uses_livekit_transport() -> None: + project_root = Path(__file__).resolve().parents[4] + namespace = runpy.run_path(str(project_root / "examples" / "stream_server" / "livekit_bidirectional_demo.py")) + + html = namespace["_render_html"]("") + + required_fragments = [ + "LingBot-World-Fast LiveKit Demo", + "livekit-client@2.21.0", + "RoomEvent.TrackSubscribed", + "RoomEvent.DataReceived", + "adaptiveStream: false", + "dynacast: false", + 'SERVER_URL + "/v1/stream/sessions"', + 'urls: ["turn:127.0.0.1:3478?transport=tcp"]', + 'iceTransportPolicy: "relay"', + "{ rtcConfig: TURN_RTC_CONFIG }", + "topic: CONTROL_TOPIC", + 'type: "control_state"', + 'event: "reset"', + 'event: "reset_pose"', + 'type: "stop"', + "telemetry-progress", + "telemetry-cadence", + "output_cadence_seconds: data.output_cadence_seconds", + "pipeline_residence_seconds: data.pipeline_residence_seconds", + "applied_control_latency_seconds: data.applied_control_latency_seconds", + ] + for fragment in required_fragments: + assert fragment in html + + assert "__UTILITIES__" not in html + assert "__CONTROLS__" not in html From 28bd24bc0dded447f3714d7d55a92cade9025b32 Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 10:46:25 +0000 Subject: [PATCH 04/11] fix(lingbot): stabilize realtime camera sessions Keep actor-stage work asynchronous while reporting queue, pipeline, and applied-control latency separately. Preserve v2 camera intrinsics, rebalance translation controls, and make session stop and cache release behavior deterministic. Verification: - .venv/bin/python -m pytest tests/unit/pipelines/lingbot_world_fast/test_service_action_loop.py tests/unit/pipelines/lingbot_world_fast/test_service_metrics.py tests/unit/pipelines/lingbot_world_v2/test_service.py -q - 52 passed - git diff --cached --check --- .../lingbot_world_fast_image_to_video_h100.py | 6 +- .../lingbot_world_v2_image_to_video_h100.py | 11 ++- .../pipelines/lingbot_world_fast/service.py | 85 +++++++++++++----- .../pipelines/lingbot_world_fast/session.py | 4 +- .../test_service_action_loop.py | 88 ++++++++++++++++--- .../test_service_metrics.py | 2 +- .../lingbot_world_v2/test_service.py | 7 ++ 7 files changed, 164 insertions(+), 39 deletions(-) diff --git a/examples/lingbot/lingbot_world_fast_image_to_video_h100.py b/examples/lingbot/lingbot_world_fast_image_to_video_h100.py index ceb4de2..e875f37 100644 --- a/examples/lingbot/lingbot_world_fast_image_to_video_h100.py +++ b/examples/lingbot/lingbot_world_fast_image_to_video_h100.py @@ -5,9 +5,11 @@ Four GPUs with Ulysses sequence parallelism: python examples/lingbot/lingbot_world_fast_image_to_video_h100.py --gpu_num 4 -WebRTC streaming service: +LiveKit streaming service: telefuser stream-serve examples/lingbot/lingbot_world_fast_image_to_video_h100.py \ - --gpu-num 4 -p 8088 --skip-validation + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey --livekit-api-secret secret \ + --worker-gpu-map 0,1,2,3 -p 8088 --skip-validation """ diff --git a/examples/lingbot/lingbot_world_v2_image_to_video_h100.py b/examples/lingbot/lingbot_world_v2_image_to_video_h100.py index 8eb21a1..3d219d5 100644 --- a/examples/lingbot/lingbot_world_v2_image_to_video_h100.py +++ b/examples/lingbot/lingbot_world_v2_image_to_video_h100.py @@ -6,9 +6,11 @@ Four GPUs with Ulysses sequence parallelism: python examples/lingbot/lingbot_world_v2_image_to_video_h100.py --gpu_num 4 Multi-GPU runs configure the VAE worker and DiT SP group independently in PPL_CONFIG. -WebRTC streaming service: +LiveKit streaming service: telefuser stream-serve examples/lingbot/lingbot_world_v2_image_to_video_h100.py \ - --gpu-num 4 -p 8088 --skip-validation + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey --livekit-api-secret secret \ + --worker-gpu-map 0,1,2,3 -p 8088 --skip-validation """ @@ -83,7 +85,7 @@ max_attention_size=None, control_move_step=0.10, control_lateral_step=0.10, - control_translation_scale=3.0, + control_translation_scale=1.0, control_yaw_step_degrees=0.5, control_pitch_step_degrees=0.5, vae_torch_dtype=torch.float32, @@ -200,6 +202,9 @@ def get_service(gpu_num: int = PPL_CONFIG["parallelism"]) -> LingBotWorldFastSer "frame_policy": PPL_CONFIG["frame_policy"], "sample_shift": PPL_CONFIG["sample_shift"], "max_attention_size": PPL_CONFIG["max_attention_size"], + "intrinsics_path": DEFAULT_INTRINSICS_PATH, + "intrinsics_width": 832, + "intrinsics_height": 480, "control_move_step": PPL_CONFIG["control_move_step"], "control_lateral_step": PPL_CONFIG["control_lateral_step"], "control_translation_scale": PPL_CONFIG["control_translation_scale"], diff --git a/telefuser/pipelines/lingbot_world_fast/service.py b/telefuser/pipelines/lingbot_world_fast/service.py index 8bb7dac..73c9f96 100644 --- a/telefuser/pipelines/lingbot_world_fast/service.py +++ b/telefuser/pipelines/lingbot_world_fast/service.py @@ -190,9 +190,10 @@ def create_session(self, config: dict) -> str: session_id = config.get("session_id") or str(uuid.uuid4()) image = self._load_image(config) - intrinsics = config.get("intrinsics") - if intrinsics is None and config.get("intrinsics_path"): - intrinsics = np.load(Path(config["intrinsics_path"])) + intrinsics = config.get("intrinsics", defaults.get("intrinsics")) + intrinsics_path = config.get("intrinsics_path", defaults.get("intrinsics_path")) + if intrinsics is None and intrinsics_path: + intrinsics = np.load(Path(intrinsics_path)) fps_value = config.get("fps", defaults.get("fps", self.default_fps)) if fps_value is None: @@ -672,7 +673,10 @@ def _positive_float(value: object, name: str) -> float: return result @staticmethod - def _queue_direction_snapshot(state: LingBotWorldFastSessionState) -> None: + def _queue_direction_snapshot( + state: LingBotWorldFastSessionState, + received_at_monotonic: float, + ) -> None: """Store only the newest short-press snapshot for the next chunk.""" if not state.pressed_controls: return @@ -683,6 +687,7 @@ def _queue_direction_snapshot(state: LingBotWorldFastSessionState) -> None: state.pending_direction_command = LingBotWorldFastDirectionCommand( revision=state.next_control_revision, controls=frozenset(state.pressed_controls), + received_at_monotonic=received_at_monotonic, ) @staticmethod @@ -696,7 +701,12 @@ def _controls_from_snapshot(chunk: dict) -> set[str] | None: return None return controls - def _update_direction_controls(self, state: LingBotWorldFastSessionState, chunk: dict) -> bool: + def _update_direction_controls( + self, + state: LingBotWorldFastSessionState, + chunk: dict, + received_at_monotonic: float, + ) -> bool: if chunk.get("type") == "control_state": controls = self._controls_from_snapshot(chunk) if controls is None: @@ -705,8 +715,8 @@ def _update_direction_controls(self, state: LingBotWorldFastSessionState, chunk: with state.control_lock: previous_controls = set(state.pressed_controls) state.pressed_controls = controls - if controls - previous_controls: - self._queue_direction_snapshot(state) + if controls and controls != previous_controls: + self._queue_direction_snapshot(state, received_at_monotonic) active_controls = sorted(state.pressed_controls) pending_count = int(state.pending_direction_command is not None) revision = state.next_control_revision @@ -732,6 +742,7 @@ def _update_direction_controls(self, state: LingBotWorldFastSessionState, chunk: logger.warning(f"Ignoring LingBot direction control with unsupported event={event!r}") return False with state.control_lock: + previous_controls = set(state.pressed_controls) if event in {"release", "keyup", "end"}: state.pressed_controls.discard(direction) elif event == "reset": @@ -744,9 +755,13 @@ def _update_direction_controls(self, state: LingBotWorldFastSessionState, chunk: state.control_pitch = 0.0 state.control_initialized = False else: - if direction not in state.pressed_controls: - state.pressed_controls.add(direction) - self._queue_direction_snapshot(state) + state.pressed_controls.add(direction) + if ( + event not in {"reset", "reset_pose"} + and state.pressed_controls + and state.pressed_controls != previous_controls + ): + self._queue_direction_snapshot(state, received_at_monotonic) controls = sorted(state.pressed_controls) pending_count = int(state.pending_direction_command is not None) revision = state.next_control_revision @@ -794,12 +809,14 @@ def _next_realtime_control( chunk_index: int, emit_status: Callable[..., None], block: bool, - ) -> tuple[object, list[str] | None] | None: + ) -> tuple[object, list[str] | None, float | None] | None: """Select one queued tap or continue a direction that remains held.""" while state.active: with state.control_lock: explicit_control = state.latest_explicit_control + explicit_control_received_at = state.latest_explicit_control_received_at_monotonic state.latest_explicit_control = None + state.latest_explicit_control_received_at_monotonic = None held_controls = frozenset(state.pressed_controls) direction_command = None if explicit_control is None and state.pending_direction_command is not None: @@ -807,14 +824,16 @@ def _next_realtime_control( state.pending_direction_command = None if explicit_control is not None: - return control_builder.defer(explicit_control), None + return control_builder.defer(explicit_control), None, explicit_control_received_at if direction_command is not None: controls = set(direction_command.controls) revision: int | None = direction_command.revision + control_received_at: float | None = direction_command.received_at_monotonic elif held_controls and block: controls = set(held_controls) revision = None + control_received_at = None else: if not block: return None @@ -844,7 +863,7 @@ def _next_realtime_control( lateral_step=state.config.control_lateral_step, pitch_step_degrees=state.config.control_pitch_step_degrees, ) - return control_builder.defer(directional_chunk), applied_controls + return control_builder.defer(directional_chunk), applied_controls, control_received_at return None def _benchmark_devices(self) -> tuple[str | torch.device, ...]: @@ -924,6 +943,7 @@ def _run_actor_worker_loop( submitted = 0 controls_by_chunk: dict[int, list[str] | None] = {} + control_received_at_by_chunk: dict[int, float | None] = {} measurements_by_chunk: dict[int, RuntimeMeasurement] = {} def raise_scheduler_error() -> None: @@ -931,9 +951,9 @@ def raise_scheduler_error() -> None: if error is not None: raise RuntimeError("LingBot streaming scheduler failed") from error - def submit_chunk(item: tuple[object, list[str] | None]) -> None: + def submit_chunk(item: tuple[object, list[str] | None, float | None]) -> None: nonlocal submitted - deferred_control, applied_controls = item + deferred_control, applied_controls, control_received_at = item control = self.pipeline._resolve_control(deferred_control) self.pipeline._validate_control(runtime, control) chunk_measurement = self._start_benchmark_measurement(state) @@ -941,6 +961,7 @@ def submit_chunk(item: tuple[object, list[str] | None]) -> None: self._finish_benchmark_measurement(chunk_measurement) raise RuntimeError("LingBot streaming ingress became unavailable after capacity check") controls_by_chunk[submitted] = applied_controls + control_received_at_by_chunk[submitted] = control_received_at if chunk_measurement is not None: measurements_by_chunk[submitted] = chunk_measurement emit_status( @@ -965,6 +986,7 @@ def submit_chunk(item: tuple[object, list[str] | None]) -> None: if not frames: raise RuntimeError(f"LingBot actor emitted no frames for chunk {result_index}") applied_controls = controls_by_chunk.pop(result_index, None) + control_received_at = control_received_at_by_chunk.pop(result_index, None) chunk_facts = self._finish_benchmark_measurement(measurements_by_chunk.pop(result_index, None)) if state.config.show_control_hud: frames = self._overlay_control_hud(frames, applied_controls) @@ -984,9 +1006,14 @@ def submit_chunk(item: tuple[object, list[str] | None]) -> None: started_at = state.chunk_started_at_monotonic.pop(result_index, now) if state.first_chunk_sent_at_monotonic is None: state.first_chunk_sent_at_monotonic = now - control_to_chunk_seconds = ( - now - state.last_control_at_monotonic if state.last_control_at_monotonic is not None else None + output_cadence_seconds = ( + now - state.last_chunk_sent_at_monotonic + if state.last_chunk_sent_at_monotonic is not None + else None ) + state.last_chunk_sent_at_monotonic = now + pipeline_residence_seconds = now - started_at + applied_control_latency_seconds = now - control_received_at if control_received_at is not None else None runtime.current_chunk_index += 1 runtime.emitted_frames += len(frames) emit_status( @@ -994,9 +1021,21 @@ def submit_chunk(item: tuple[object, list[str] | None]) -> None: index=result_index, controls=applied_controls or [], frames=len(frames), - chunk_elapsed_seconds=round(now - started_at, 6), + output_cadence_seconds=( + round(output_cadence_seconds, 6) if output_cadence_seconds is not None else None + ), + pipeline_residence_seconds=round(pipeline_residence_seconds, 6), + applied_control_latency_seconds=( + round(applied_control_latency_seconds, 6) + if applied_control_latency_seconds is not None + else None + ), + # Backward-compatible aliases for clients that have not adopted the explicit metric names. + chunk_elapsed_seconds=round(pipeline_residence_seconds, 6), control_to_chunk_seconds=( - round(control_to_chunk_seconds, 6) if control_to_chunk_seconds is not None else None + round(applied_control_latency_seconds, 6) + if applied_control_latency_seconds is not None + else None ), runtime_metrics=self._runtime_metrics(state), stream_progress=self._stream_progress(state, runtime), @@ -1140,10 +1179,11 @@ def push_chunk(self, session_id: str, chunk: dict) -> None: state = self._sessions.get(session_id) if state is None or not state.active: return - with state.metrics_lock: - state.last_control_at_monotonic = time.monotonic() + received_at_monotonic = time.monotonic() is_direction_action = chunk.get("type") in {"control", "control_state"} and self._update_direction_controls( - state, chunk + state, + chunk, + received_at_monotonic, ) if is_direction_action: self._wake_control_worker(state, {"type": "direction_control"}) @@ -1153,6 +1193,7 @@ def push_chunk(self, session_id: str, chunk: dict) -> None: with state.metrics_lock: state.overwritten_explicit_controls += 1 state.latest_explicit_control = chunk + state.latest_explicit_control_received_at_monotonic = received_at_monotonic self._wake_control_worker(state) async def pull_chunks(self, session_id: str) -> AsyncGenerator[dict, None]: diff --git a/telefuser/pipelines/lingbot_world_fast/session.py b/telefuser/pipelines/lingbot_world_fast/session.py index 6a3604e..da38e82 100644 --- a/telefuser/pipelines/lingbot_world_fast/session.py +++ b/telefuser/pipelines/lingbot_world_fast/session.py @@ -113,6 +113,7 @@ class LingBotWorldFastDirectionCommand: revision: int controls: frozenset[str] + received_at_monotonic: float @dataclass @@ -131,8 +132,8 @@ class LingBotWorldFastSessionState: active: bool = True created_at_monotonic: float = field(default_factory=time.monotonic) worker_started_at_monotonic: float | None = None - last_control_at_monotonic: float | None = None first_chunk_sent_at_monotonic: float | None = None + last_chunk_sent_at_monotonic: float | None = None chunk_started_at_monotonic: dict[int, float] = field(default_factory=dict) output_queue_high_watermark: int = 0 dropped_video_payloads: int = 0 @@ -144,6 +145,7 @@ class LingBotWorldFastSessionState: last_applied_control_revision: int | None = None overwritten_direction_commands: int = 0 latest_explicit_control: dict | None = None + latest_explicit_control_received_at_monotonic: float | None = None overwritten_explicit_controls: int = 0 dropped_control_signals: int = 0 scheduler_metrics: dict[str, object] | None = None diff --git a/tests/unit/pipelines/lingbot_world_fast/test_service_action_loop.py b/tests/unit/pipelines/lingbot_world_fast/test_service_action_loop.py index b1e9659..7b31714 100644 --- a/tests/unit/pipelines/lingbot_world_fast/test_service_action_loop.py +++ b/tests/unit/pipelines/lingbot_world_fast/test_service_action_loop.py @@ -77,11 +77,13 @@ def test_actor_worker_submits_control_and_emits_ordered_chunk() -> None: state.control_context = SimpleNamespace() control_builder = MagicMock() emit_status = MagicMock() - first_control = (object(), ["w"]) + first_control = (object(), ["w"], 8.0) + state.last_chunk_sent_at_monotonic = 9.0 with ( patch.object(service, "_next_realtime_control", return_value=first_control), patch.object(service, "_put_output") as put_output, + patch("telefuser.pipelines.lingbot_world_fast.service.time.monotonic", side_effect=[10.0, 12.0, 13.0]), ): service._run_actor_worker_loop(state, state.control_context, control_builder, emit_status) @@ -94,6 +96,12 @@ def test_actor_worker_submits_control_and_emits_ordered_chunk() -> None: assert runtime.current_chunk_index == 1 assert runtime.emitted_frames == 1 assert state.streaming_session is streaming_session + chunk_sent = next(call for call in emit_status.call_args_list if call.args[0] == "chunk_sent") + assert chunk_sent.kwargs["output_cadence_seconds"] == 3.0 + assert chunk_sent.kwargs["pipeline_residence_seconds"] == 2.0 + assert chunk_sent.kwargs["applied_control_latency_seconds"] == 4.0 + assert chunk_sent.kwargs["chunk_elapsed_seconds"] == 2.0 + assert chunk_sent.kwargs["control_to_chunk_seconds"] == 4.0 def test_chunk_hud_and_metadata_use_the_control_snapshot_submitted_to_the_model() -> None: @@ -127,7 +135,7 @@ def test_chunk_hud_and_metadata_use_the_control_snapshot_submitted_to_the_model( state = LingBotWorldFastSessionState(config=runtime.config, control_context=SimpleNamespace()) emit_status = MagicMock() with ( - patch.object(service, "_next_realtime_control", return_value=(object(), ["w", "j"])), + patch.object(service, "_next_realtime_control", return_value=(object(), ["w", "j"], None)), patch.object(service, "_overlay_control_hud", return_value=frames) as overlay, patch.object(service, "_put_output") as put_output, ): @@ -150,7 +158,7 @@ def initialize(*_args: object, **_kwargs: object) -> LingBotWorldFastGenerationS return runtime pipeline._create_initialized_session.side_effect = initialize - with patch.object(service, "_next_realtime_control", return_value=(object(), None)): + with patch.object(service, "_next_realtime_control", return_value=(object(), None, None)): service._run_actor_worker_loop(state, MagicMock(), MagicMock(), MagicMock()) pipeline._get_streaming_runtime.assert_not_called() @@ -208,7 +216,7 @@ def stop_wait(*_args: object, **_kwargs: object) -> bool: streaming_runtime.wait_until_idle.side_effect = stop_wait pipeline._get_streaming_runtime.return_value = streaming_runtime - with patch.object(service, "_next_realtime_control", return_value=(object(), ["w"])): + with patch.object(service, "_next_realtime_control", return_value=(object(), ["w"], None)): service._run_actor_worker_loop(state, state.control_context, MagicMock(), MagicMock()) assert streaming_runtime.try_submit_chunk.call_args_list == [ @@ -222,14 +230,16 @@ def test_direction_action_updates_state_and_wakes_worker() -> None: state = _state() service._sessions["session-a"] = state - service.push_chunk( - "session-a", - {"type": "control", "direction": "up", "event": "press"}, - ) + with patch("telefuser.pipelines.lingbot_world_fast.service.time.monotonic", return_value=123.0): + service.push_chunk( + "session-a", + {"type": "control", "direction": "up", "event": "press"}, + ) assert state.pressed_controls == {"w"} assert state.pending_direction_command is not None assert (state.pending_direction_command.revision, state.pending_direction_command.controls) == (1, frozenset({"w"})) + assert state.pending_direction_command.received_at_monotonic == 123.0 assert state.pending_inputs.get_nowait() == {"type": "direction_control"} @@ -357,6 +367,40 @@ def test_control_state_supports_combined_translation_and_rotation() -> None: assert [call.args[0]["controls"] for call in control_builder.defer.call_args_list] == [["j", "w"]] +def test_control_state_release_replaces_pending_combined_snapshot() -> None: + service = LingBotWorldFastService(MagicMock()) + state = _state() + state.control_context = SimpleNamespace(control_type="cam", chunk_size=3) + service._sessions["session-a"] = state + control_builder = MagicMock() + + service.push_chunk("session-a", {"type": "control_state", "controls": ["j"]}) + service._next_realtime_control(state, state.control_context, control_builder, 0, MagicMock(), block=True) + service.push_chunk("session-a", {"type": "control_state", "controls": ["j", "w"]}) + service.push_chunk("session-a", {"type": "control_state", "controls": ["w"]}) + service._next_realtime_control(state, state.control_context, control_builder, 1, MagicMock(), block=True) + + assert state.pressed_controls == {"w"} + assert [call.args[0]["controls"] for call in control_builder.defer.call_args_list] == [["j"], ["w"]] + + +def test_direction_release_replaces_pending_combined_snapshot() -> None: + service = LingBotWorldFastService(MagicMock()) + state = _state() + state.control_context = SimpleNamespace(control_type="cam", chunk_size=3) + service._sessions["session-a"] = state + control_builder = MagicMock() + + service.push_chunk("session-a", {"type": "control", "direction": "left", "event": "press"}) + service._next_realtime_control(state, state.control_context, control_builder, 0, MagicMock(), block=True) + service.push_chunk("session-a", {"type": "control", "direction": "forward", "event": "press"}) + service.push_chunk("session-a", {"type": "control", "direction": "left", "event": "release"}) + service._next_realtime_control(state, state.control_context, control_builder, 1, MagicMock(), block=True) + + assert state.pressed_controls == {"w"} + assert [call.args[0]["controls"] for call in control_builder.defer.call_args_list] == [["a"], ["w"]] + + def test_held_direction_continues_only_after_a_chunk_is_completed() -> None: service = LingBotWorldFastService(MagicMock()) state = _state() @@ -394,10 +438,12 @@ def test_explicit_controls_keep_only_the_latest_pending_value() -> None: first = {"control_tensor": "first"} second = {"control_tensor": "second"} - service.push_chunk("session-a", first) - service.push_chunk("session-a", second) + with patch("telefuser.pipelines.lingbot_world_fast.service.time.monotonic", side_effect=[10.0, 20.0]): + service.push_chunk("session-a", first) + service.push_chunk("session-a", second) assert state.latest_explicit_control == second + assert state.latest_explicit_control_received_at_monotonic == 20.0 assert state.overwritten_explicit_controls == 1 assert state.pending_inputs.qsize() == 1 @@ -611,6 +657,28 @@ def test_create_session_initializes_fixed_intrinsics_from_intrinsics_path() -> N assert service._sessions[session_id].control_context is pipeline.control_context.return_value +def test_create_session_inherits_fixed_intrinsics_from_service_defaults() -> None: + pipeline = MagicMock() + service = LingBotWorldFastService( + pipeline, + default_session_config={ + "intrinsics_path": "/controls/default-intrinsics.npy", + "intrinsics_width": 832, + "intrinsics_height": 480, + }, + ) + intrinsics = np.asarray([[415.0, 416.0, 415.5, 239.5]]) + + with patch("telefuser.pipelines.lingbot_world_fast.service.np.load", return_value=intrinsics) as load: + service.create_session({"image": Image.new("RGB", (16, 9))}) + + load.assert_called_once_with(Path("/controls/default-intrinsics.npy")) + session_config = pipeline.control_context.call_args.args[0] + assert session_config.intrinsics is intrinsics + assert session_config.intrinsics_width == 832 + assert session_config.intrinsics_height == 480 + + def test_pull_chunks_drains_terminal_messages_after_session_becomes_inactive() -> None: service = LingBotWorldFastService(MagicMock()) state = _state() diff --git a/tests/unit/pipelines/lingbot_world_fast/test_service_metrics.py b/tests/unit/pipelines/lingbot_world_fast/test_service_metrics.py index e2c3110..ef83466 100644 --- a/tests/unit/pipelines/lingbot_world_fast/test_service_metrics.py +++ b/tests/unit/pipelines/lingbot_world_fast/test_service_metrics.py @@ -62,7 +62,7 @@ def emit_status(stage: str, **data: object) -> None: statuses.append({"stage": stage, **data}) with ( - patch.object(service, "_next_realtime_control", return_value=(object(), None)), + patch.object(service, "_next_realtime_control", return_value=(object(), None, None)), patch.object(service, "_put_output"), patch( "telefuser.pipelines.lingbot_world_fast.service.start_runtime_measurement", diff --git a/tests/unit/pipelines/lingbot_world_v2/test_service.py b/tests/unit/pipelines/lingbot_world_v2/test_service.py index b033828..84a52fa 100644 --- a/tests/unit/pipelines/lingbot_world_v2/test_service.py +++ b/tests/unit/pipelines/lingbot_world_v2/test_service.py @@ -1,5 +1,6 @@ from unittest.mock import MagicMock, patch +import numpy as np import torch from PIL import Image from click.testing import CliRunner @@ -116,8 +117,14 @@ def test_v2_unified_example_service_constructs_v2_session_from_ppl_config() -> N assert session_config.chunk_size == 4 assert session_config.frame_policy == "truncate" assert session_config.sample_shift == 10.0 + assert session_config.control_translation_scale == 1.0 + np.testing.assert_array_equal(session_config.intrinsics, np.load(offline_example.DEFAULT_INTRINSICS_PATH)) + assert session_config.intrinsics_width == 832 + assert session_config.intrinsics_height == 480 assert service.default_fps == 16 assert service.max_generation_seconds == 120.0 + assert service.default_session_config["control_translation_scale"] == 1.0 + assert service.default_session_config["intrinsics_path"] == offline_example.DEFAULT_INTRINSICS_PATH assert session_id in service._sessions From 6d4fa032cdb8f5b76d090e382986b1c4a7233cf7 Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 10:47:20 +0000 Subject: [PATCH 05/11] refactor(streaming): remove the direct aiortc backend Make LiveKit the only stream transport, remove direct SDP and WebRTC API routes, delete the legacy aiortc session implementation and browser clients, and retain server-push support through the unified stream service contract. Verification: - .venv/bin/python -m pytest tests/unit/service -q - 182 passed - git diff --cached --check --- .../stream_server/stream_arrow_overlay.py | 5 +- examples/stream_server/stream_video_replay.py | 13 +- .../webrtc_arrow_overlay_demo.py | 257 ----- .../webrtc_bidirectional_demo.py | 957 ------------------ examples/stream_server/webrtc_client_demo.py | 221 ---- pyproject.toml | 6 - scripts/run_ci_tests.sh | 2 +- telefuser/service/api/api_server.py | 47 +- telefuser/service/api/routers/__init__.py | 8 +- telefuser/service/api/routers/service.py | 26 - telefuser/service/api/routers/stream.py | 95 -- telefuser/service/api/routers/webrtc.py | 186 ---- telefuser/service/api/stream_schema.py | 28 +- telefuser/service/core/config.py | 40 - telefuser/service/core/container.py | 28 - .../service/core/stream_pipeline_service.py | 6 +- telefuser/service/main.py | 34 - telefuser/service/webrtc/__init__.py | 32 - telefuser/service/webrtc/chunk_router.py | 143 --- telefuser/service/webrtc/session_manager.py | 469 --------- telefuser/service/webrtc/track.py | 343 ------- tests/integration/conftest.py | 163 --- tests/integration/test_stream_api.py | 188 ---- tests/integration/test_webrtc_api.py | 180 ---- tests/unit/service/test_service_routes.py | 476 --------- .../service/test_webrtc_session_manager.py | 208 ---- webui/stream_app.py | 191 ---- 27 files changed, 23 insertions(+), 4329 deletions(-) delete mode 100644 examples/stream_server/webrtc_arrow_overlay_demo.py delete mode 100644 examples/stream_server/webrtc_bidirectional_demo.py delete mode 100644 examples/stream_server/webrtc_client_demo.py delete mode 100644 telefuser/service/api/routers/stream.py delete mode 100644 telefuser/service/api/routers/webrtc.py delete mode 100644 telefuser/service/webrtc/__init__.py delete mode 100644 telefuser/service/webrtc/chunk_router.py delete mode 100644 telefuser/service/webrtc/session_manager.py delete mode 100644 telefuser/service/webrtc/track.py delete mode 100644 tests/integration/conftest.py delete mode 100644 tests/integration/test_stream_api.py delete mode 100644 tests/integration/test_webrtc_api.py delete mode 100644 tests/unit/service/test_webrtc_session_manager.py delete mode 100644 webui/stream_app.py diff --git a/examples/stream_server/stream_arrow_overlay.py b/examples/stream_server/stream_arrow_overlay.py index fe67c92..84e23bb 100644 --- a/examples/stream_server/stream_arrow_overlay.py +++ b/examples/stream_server/stream_arrow_overlay.py @@ -5,7 +5,10 @@ overlay via ``pull_chunks()``. Usage: - telefuser stream-serve examples/stream_server/stream_arrow_overlay.py -p 8088 --skip-validation + telefuser stream-serve examples/stream_server/stream_arrow_overlay.py \ + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey --livekit-api-secret secret \ + -p 8088 --skip-validation """ from __future__ import annotations diff --git a/examples/stream_server/stream_video_replay.py b/examples/stream_server/stream_video_replay.py index 6145975..7cdeda3 100644 --- a/examples/stream_server/stream_video_replay.py +++ b/examples/stream_server/stream_video_replay.py @@ -5,7 +5,12 @@ base64 when the source contains an audio track. Usage: - telefuser stream-serve examples/stream_video_replay.py -p 8088 --skip-validation + telefuser stream-serve examples/stream_server/stream_video_replay.py \ + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey \ + --livekit-api-secret secret \ + -p 8088 \ + --skip-validation """ from __future__ import annotations @@ -23,11 +28,7 @@ VIDEO_PATH = str(Path(__file__).parent / "data" / "liveact_1.mp4") FRAMES_PER_CHUNK = 32 OUTPUT_FPS = 24 - -try: - from telefuser.service.webrtc.track import AUDIO_SAMPLE_RATE -except ImportError: - AUDIO_SAMPLE_RATE = 48_000 +AUDIO_SAMPLE_RATE = 48_000 class VideoReplayService: diff --git a/examples/stream_server/webrtc_arrow_overlay_demo.py b/examples/stream_server/webrtc_arrow_overlay_demo.py deleted file mode 100644 index 44b8510..0000000 --- a/examples/stream_server/webrtc_arrow_overlay_demo.py +++ /dev/null @@ -1,257 +0,0 @@ -"""WebRTC arrow overlay client demo. - -Connects to a bidirectional ``ArrowOverlayService`` and sends keyboard -arrow-key events over a DataChannel. The server overlays a D-pad HUD -on the video and streams it back via a WebRTC media track. - -Usage: - # 1. Start the server: - telefuser stream-serve examples/stream_server/stream_arrow_overlay.py -p 8088 --skip-validation - - # 2. Start this client: - python examples/stream_server/webrtc_arrow_overlay_demo.py --server-url http://localhost:8088 - - # 3. Click Connect, then press arrow keys. -""" - -from __future__ import annotations - -import argparse -import functools -import http.server -import threading -import webbrowser - -DEFAULT_SERVER_URL = "http://localhost:8088" -DEFAULT_PORT = 8092 - -HTML_TEMPLATE = """ - - - -TeleFuser Arrow Overlay Demo - - - -

Arrow Overlay Demo

- - -
- - -
-
Ready. Click Connect then press arrow keys.
-

Use keyboard arrow keys. The server overlays a D-pad HUD on the video.

- -
-
-
-
-
-
-
-
-
-
-
- -

DataChannel Log

-
- - - -""" - - -def main() -> None: - parser = argparse.ArgumentParser(description="TeleFuser arrow overlay WebRTC client demo") - parser.add_argument("--server-url", default=DEFAULT_SERVER_URL, help="Stream server base URL") - parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Local HTTP server port") - parser.add_argument("--no-open", action="store_true", help="Don't open browser automatically") - args = parser.parse_args() - - html = HTML_TEMPLATE.format(server_url=args.server_url) - - class Handler(http.server.BaseHTTPRequestHandler): - def do_GET(self) -> None: - self.send_response(200) - self.send_header("Content-Type", "text/html; charset=utf-8") - self.end_headers() - self.wfile.write(html.encode("utf-8")) - - def log_message(self, format: str, *_args: object) -> None: - pass - - server = http.server.HTTPServer(("0.0.0.0", args.port), Handler) - url = f"http://localhost:{args.port}" - print(f"Serving arrow overlay demo at {url}") - print(f"Stream server: {args.server_url}") - print("Press Ctrl+C to stop.\n") - - if not args.no_open: - threading.Timer(0.5, functools.partial(webbrowser.open, url)).start() - - try: - server.serve_forever() - except KeyboardInterrupt: - print("\nStopped.") - server.server_close() - - -if __name__ == "__main__": - main() diff --git a/examples/stream_server/webrtc_bidirectional_demo.py b/examples/stream_server/webrtc_bidirectional_demo.py deleted file mode 100644 index bfdd712..0000000 --- a/examples/stream_server/webrtc_bidirectional_demo.py +++ /dev/null @@ -1,957 +0,0 @@ -"""LingBot-World-Fast WebRTC control demo. - -Demonstrates the LingBot-World-Fast bidirectional WebRTC protocol: - -* Client creates a DataChannel ("telefuser") for JSON control messages. -* Client sends prompt and direction controls. -* Server sends generated video via media tracks and metadata via DataChannel. - -Usage: - # 1. Start the LingBot stream server: - telefuser stream-serve examples/lingbot/lingbot_world_fast_image_to_video_h100.py -p 8088 --skip-validation - - # 2. Start this client (opens browser): - python examples/stream_server/webrtc_bidirectional_demo.py --server-url http://localhost:8088 - - # 3. Select an image, enter a prompt, click Connect, then use arrow keys or the D-pad. -""" - -from __future__ import annotations - -import argparse -import functools -import http.server -import json -import os -import threading -import urllib.error -import urllib.request -import webbrowser -from pathlib import Path - -DEFAULT_SERVER_URL = "http://localhost:8088" -DEFAULT_PORT = 8091 -_PROJECT_ROOT = Path(__file__).resolve().parents[2] -DEFAULT_IMAGE_PATH = str(_PROJECT_ROOT / "examples" / "data" / "lingbot_world_fast" / "image.jpg") -DEFAULT_PROMPT = ( - "A serene lakeside scene with a lone tree standing in calm water, surrounded by distant snow-capped " - "mountains under a bright blue sky with drifting white clouds. Gentle ripples reflect the tree and sky." -) - -HTML_TEMPLATE = r""" - - - -LingBot-World-Fast WebRTC Demo - - - -
-

LingBot-World-Fast WebRTC Demo

-
-
-
-

Server Output

-
Ready.
-
- -
-
Server limit--
-
Target video--
-
Generated video--
-
Frames / chunks--
-
Chunk / control latency--
-
Queue / dropped video--
-
-
- - -
-
- - - -""" - - -def main() -> None: - parser = argparse.ArgumentParser(description="LingBot-World-Fast WebRTC control demo") - parser.add_argument("--server-url", default=DEFAULT_SERVER_URL, help="Stream server base URL") - parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Local HTTP server port") - parser.add_argument( - "--ice-gather-timeout-ms", - type=int, - default=10000, - help="How long the browser waits for ICE candidates before sending the SDP offer", - ) - parser.add_argument( - "--proxy-backend", - action=argparse.BooleanOptionalAction, - default=True, - help="Proxy /v1/stream/webrtc/* via this demo server to --server-url (recommended for VS Code Remote port forwarding)", - ) - parser.add_argument( - "--turn-url", - default=os.environ.get("TELEFUSER_TURN_SERVER", ""), - help="TURN server URL for browser WebRTC ICE, e.g. turn:localhost:3478?transport=tcp", - ) - parser.add_argument( - "--turn-username", - default=os.environ.get("TELEFUSER_TURN_USERNAME", ""), - help="TURN username; defaults to TELEFUSER_TURN_USERNAME", - ) - parser.add_argument( - "--turn-credential", - default=os.environ.get("TELEFUSER_TURN_CREDENTIAL", ""), - help="TURN credential; defaults to TELEFUSER_TURN_CREDENTIAL", - ) - parser.add_argument( - "--force-turn-relay", - action="store_true", - help="Force WebRTC to use relay candidates only. Useful when testing through SSH port forwarding.", - ) - parser.add_argument("--no-open", action="store_true", help="Don't open browser automatically") - args = parser.parse_args() - - rtc_config: dict[str, object] = {} - if args.turn_url: - turn_server: dict[str, str] = {"urls": args.turn_url} - if args.turn_username: - turn_server["username"] = args.turn_username - if args.turn_credential: - turn_server["credential"] = args.turn_credential - rtc_config["iceServers"] = [turn_server] - if args.force_turn_relay: - rtc_config["iceTransportPolicy"] = "relay" - - # When proxying, the browser should call the demo origin (no separate port forward needed for --server-url). - server_url_for_browser = "" if args.proxy_backend else args.server_url - - html = ( - HTML_TEMPLATE.replace("__SERVER_URL__", json.dumps(server_url_for_browser)) - .replace("__RTC_CONFIG__", json.dumps(rtc_config)) - .replace("__DEFAULT_IMAGE_PATH__", json.dumps(DEFAULT_IMAGE_PATH)) - .replace("__PROMPT__", json.dumps(DEFAULT_PROMPT)) - .replace("__ICE_GATHER_TIMEOUT_MS__", str(args.ice_gather_timeout_ms)) - ) - - class Handler(http.server.BaseHTTPRequestHandler): - _opener = urllib.request.build_opener(urllib.request.ProxyHandler({})) - - def _proxy_backend(self) -> bool: - return bool(args.proxy_backend) and self.path.startswith("/v1/stream/webrtc/") - - def _proxy(self) -> None: - backend = args.server_url.rstrip("/") - url = f"{backend}{self.path}" - content_len = int(self.headers.get("Content-Length") or "0") - body = self.rfile.read(content_len) if content_len > 0 else None - - headers: dict[str, str] = {} - content_type = self.headers.get("Content-Type") - if content_type: - headers["Content-Type"] = content_type - - req = urllib.request.Request(url, data=body, headers=headers, method=self.command) - try: - with self._opener.open(req, timeout=30) as resp: - resp_body = resp.read() - status = getattr(resp, "status", 200) - resp_headers = resp.headers - except urllib.error.HTTPError as exc: - status = exc.code - resp_headers = exc.headers - resp_body = exc.read() - except Exception as exc: - resp_body = json.dumps({"detail": f"Demo proxy error: {exc}"}).encode("utf-8") - self.send_response(502) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(resp_body))) - self.end_headers() - self.wfile.write(resp_body) - return - - self.send_response(status) - if resp_headers.get("Content-Type"): - self.send_header("Content-Type", resp_headers.get("Content-Type")) - else: - self.send_header("Content-Type", "application/octet-stream") - self.send_header("Content-Length", str(len(resp_body))) - self.end_headers() - self.wfile.write(resp_body) - - def do_GET(self) -> None: - if self.path == "/default-image": - body = Path(DEFAULT_IMAGE_PATH).read_bytes() - self.send_response(200) - self.send_header("Content-Type", "image/jpeg") - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) - return - body = html.encode("utf-8") - self.send_response(200) - self.send_header("Content-Type", "text/html; charset=utf-8") - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) - - def do_POST(self) -> None: - if self._proxy_backend(): - self._proxy() - return - self.send_error(404) - - def do_DELETE(self) -> None: - if self._proxy_backend(): - self._proxy() - return - self.send_error(404) - - def log_message(self, format: str, *_args: object) -> None: - pass - - server = http.server.ThreadingHTTPServer(("0.0.0.0", args.port), Handler) - url = f"http://localhost:{args.port}" - print(f"Serving LingBot-World-Fast WebRTC demo at {url}") - print(f"Stream server: {args.server_url}") - print(f"ICE gather timeout: {args.ice_gather_timeout_ms} ms") - if args.proxy_backend: - print("Proxy: enabled (browser will call this demo origin; no separate port forward needed for --server-url)") - if args.turn_url and "localhost" in args.turn_url: - print( - "TURN uses localhost. This is valid only when local port 3478 is forwarded to the remote TURN server; " - "otherwise use the remote public IP/DNS or disable --force-turn-relay." - ) - if rtc_config: - print(f"WebRTC config: {json.dumps(rtc_config)}") - print("Press Ctrl+C to stop.\n") - - if not args.no_open: - threading.Timer(0.5, functools.partial(webbrowser.open, url)).start() - - try: - server.serve_forever() - except KeyboardInterrupt: - print("\nStopped.") - server.server_close() - - -if __name__ == "__main__": - main() diff --git a/examples/stream_server/webrtc_client_demo.py b/examples/stream_server/webrtc_client_demo.py deleted file mode 100644 index b35d6e7..0000000 --- a/examples/stream_server/webrtc_client_demo.py +++ /dev/null @@ -1,221 +0,0 @@ -"""WebRTC client demo: serves a minimal HTML page that streams video via WebRTC. - -Usage: - # 1. Start the stream server: - telefuser stream-serve examples/stream_video_replay.py -p 8088 --skip-validation - - # 2. Start this client (opens browser): - python examples/webrtc_client_demo.py --server-url http://localhost:8088 - - # 3. Enter a prompt and click Connect — video plays in real-time via WebRTC. -""" - -from __future__ import annotations - -import argparse -import functools -import http.server -import json -import os -import threading -import webbrowser - -DEFAULT_SERVER_URL = "http://localhost:8088" -DEFAULT_PORT = 8090 - -HTML_TEMPLATE = """ - - - -TeleFuser WebRTC Demo - - - -

TeleFuser WebRTC Demo

- -
- - - - -
-
Ready.
- - - -""" - - -def main() -> None: - parser = argparse.ArgumentParser(description="TeleFuser WebRTC client demo") - parser.add_argument("--server-url", default=DEFAULT_SERVER_URL, help="Stream server base URL") - parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Local HTTP server port") - parser.add_argument( - "--turn-url", - default=os.environ.get("TELEFUSER_TURN_SERVER", ""), - help="TURN server URL for browser WebRTC ICE, e.g. turn:localhost:3478?transport=tcp", - ) - parser.add_argument( - "--turn-username", - default=os.environ.get("TELEFUSER_TURN_USERNAME", ""), - help="TURN username; defaults to TELEFUSER_TURN_USERNAME", - ) - parser.add_argument( - "--turn-credential", - default=os.environ.get("TELEFUSER_TURN_CREDENTIAL", ""), - help="TURN credential; defaults to TELEFUSER_TURN_CREDENTIAL", - ) - parser.add_argument( - "--force-turn-relay", - action="store_true", - help="Force WebRTC to use relay candidates only. Useful when testing through SSH port forwarding.", - ) - parser.add_argument("--no-open", action="store_true", help="Don't open browser automatically") - args = parser.parse_args() - - rtc_config: dict[str, object] = {} - if args.turn_url: - turn_server: dict[str, str] = {"urls": args.turn_url} - if args.turn_username: - turn_server["username"] = args.turn_username - if args.turn_credential: - turn_server["credential"] = args.turn_credential - rtc_config["iceServers"] = [turn_server] - if args.force_turn_relay: - rtc_config["iceTransportPolicy"] = "relay" - - html = HTML_TEMPLATE.format(server_url=args.server_url, rtc_config=json.dumps(rtc_config)) - - class Handler(http.server.BaseHTTPRequestHandler): - def do_GET(self) -> None: - self.send_response(200) - self.send_header("Content-Type", "text/html; charset=utf-8") - self.end_headers() - self.wfile.write(html.encode("utf-8")) - - def log_message(self, format: str, *_args: object) -> None: - pass - - server = http.server.HTTPServer(("0.0.0.0", args.port), Handler) - url = f"http://localhost:{args.port}" - print(f"Serving WebRTC demo at {url}") - print(f"Stream server: {args.server_url}") - if rtc_config: - print(f"WebRTC config: {json.dumps(rtc_config)}") - print("Press Ctrl+C to stop.\n") - - if not args.no_open: - threading.Timer(0.5, functools.partial(webbrowser.open, url)).start() - - try: - server.serve_forever() - except KeyboardInterrupt: - print("\nStopped.") - server.server_close() - - -if __name__ == "__main__": - main() diff --git a/pyproject.toml b/pyproject.toml index 83ac924..bbd1678 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,7 +29,6 @@ classifiers = [ dependencies = [ # Default install follows vLLM-style inference readiness: service runtime # plus the core PyTorch/model/distributed stack needed by standard pipelines. - "aiortc>=1.9.0", "click", "diffusers>=0.36.0", "einops", @@ -70,11 +69,6 @@ distributed = [ "ray", ] -livekit = [ - "livekit-api>=1.0.0", - "livekit>=1.0.0", -] - dev = [ "gradio==5.50", "ray", diff --git a/scripts/run_ci_tests.sh b/scripts/run_ci_tests.sh index 6f03dd4..9c883f4 100755 --- a/scripts/run_ci_tests.sh +++ b/scripts/run_ci_tests.sh @@ -40,7 +40,7 @@ fi # Install dependencies if needed print_section "Installing dependencies" -pip install -e ".[dev,webrtc]" -q +pip install -e ".[dev]" -q pip install torch --index-url https://download.pytorch.org/whl/cpu -q check_result "Dependencies installation" diff --git a/telefuser/service/api/api_server.py b/telefuser/service/api/api_server.py index 89c2841..f139850 100644 --- a/telefuser/service/api/api_server.py +++ b/telefuser/service/api/api_server.py @@ -10,7 +10,6 @@ import httpx from fastapi import FastAPI, HTTPException -from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import StreamingResponse from telefuser.utils.logging import logger @@ -25,7 +24,6 @@ if TYPE_CHECKING: from ..core.pipeline_service import PipelineService - from ..core.stream_pipeline_service import StreamPipelineService class ApiServer: @@ -46,7 +44,6 @@ def __init__( enable_logging: bool = False, enable_openai_api: bool = True, config: ServerConfig | None = None, - route_profile: str = "all", ) -> None: self.server_config = config or server_config self.app = app or FastAPI( @@ -60,8 +57,6 @@ def __init__( ) self.file_service: FileService | None = None self.inference_service: PipelineService | None = None - self.stream_service: StreamPipelineService | None = None - self._webrtc_routes: object | None = None self.media_service: MediaGenerationService | None = None self.task_app_service = TaskApplicationService(self) self.cache_service: Any | None = None @@ -70,7 +65,6 @@ def __init__( self.configured_max_concurrent_tasks = configured_max_concurrent_tasks or max_concurrent_tasks self._task_manager = task_manager self.enable_openai_api = enable_openai_api - self.route_profile = route_profile self.task_processor: AsyncTaskProcessor | None = None self._task_processor_lock = asyncio.Lock() @@ -99,25 +93,12 @@ def _setup_routes(self) -> None: service_router = routers.service.setup_routes(self) self.app.include_router(service_router) - if self.route_profile in {"all", "request_response"}: - tasks_router = routers.tasks.setup_routes(self) - files_router = routers.files.setup_routes(self) - self.app.include_router(tasks_router) - self.app.include_router(files_router) - - if self.route_profile in {"all", "stream"}: - stream_router = routers.setup_stream_routes(self) - self.app.include_router(stream_router) - - if routers.setup_webrtc_routes is not None: - try: - webrtc_router = routers.setup_webrtc_routes(self) - self.app.include_router(webrtc_router) - logger.info("WebRTC routes enabled at /v1/stream/webrtc") - except Exception as e: - logger.info(f"WebRTC routes not available: {e}") - - if self.enable_openai_api and self.route_profile in {"all", "request_response"}: + tasks_router = routers.tasks.setup_routes(self) + files_router = routers.files.setup_routes(self) + self.app.include_router(tasks_router) + self.app.include_router(files_router) + + if self.enable_openai_api: try: from .openai import image_routes, video_routes @@ -363,16 +344,6 @@ def initialize_services( max_concurrent=self.max_concurrent_tasks, ) - def initialize_stream_service(self, stream_service: StreamPipelineService) -> None: - """Initialize stream pipeline service for stream-mode endpoints.""" - self.stream_service = stream_service - self.app.add_middleware( - CORSMiddleware, - allow_origins=["*"], - allow_methods=["*"], - allow_headers=["*"], - ) - async def cleanup(self) -> None: """Cleanup resources and stop processing workers.""" if self._artifact_cleanup_task is not None: @@ -386,12 +357,6 @@ async def cleanup(self) -> None: if self.task_processor is not None: await self.task_processor.stop() - if self._webrtc_routes is not None: - await self._webrtc_routes.cleanup() - - if self.stream_service is not None: - await self.stream_service.aclose() - if self.file_service: try: self.run_artifact_cleanup() diff --git a/telefuser/service/api/routers/__init__.py b/telefuser/service/api/routers/__init__.py index a52e318..13b08db 100644 --- a/telefuser/service/api/routers/__init__.py +++ b/telefuser/service/api/routers/__init__.py @@ -9,12 +9,6 @@ from .files import router as files_router from .service import router as service_router -from .stream import setup_routes as setup_stream_routes from .tasks import router as tasks_router -try: - from .webrtc import setup_routes as setup_webrtc_routes -except ImportError: - setup_webrtc_routes = None # type: ignore[assignment] - -__all__ = ["tasks_router", "files_router", "service_router", "setup_stream_routes", "setup_webrtc_routes"] +__all__ = ["tasks_router", "files_router", "service_router"] diff --git a/telefuser/service/api/routers/service.py b/telefuser/service/api/routers/service.py index 3231960..8dbd0be 100644 --- a/telefuser/service/api/routers/service.py +++ b/telefuser/service/api/routers/service.py @@ -36,10 +36,6 @@ async def get_status(self) -> dict: status["pool"] = pool_status else: status["execution_mode"] = "serial_single_pipeline" - webrtc_stats = self._webrtc_session_stats() - status.update(webrtc_stats) - if webrtc_stats.get("webrtc_active_sessions", 0) > 0 and status.get("service_status") == "idle": - status["service_status"] = "active" return status def _pipeline_pool_status(self) -> list[dict] | None: @@ -54,13 +50,6 @@ def _pipeline_pool_status(self) -> list[dict] | None: return None return pool_status - def _webrtc_session_stats(self) -> dict: - """Return WebRTC session stats if available.""" - routes = self.api._webrtc_routes - if routes is None: - return {} - return routes._session_manager.session_stats() - async def get_metadata(self) -> dict: """Get service metadata.""" if self.api.inference_service is not None: @@ -70,12 +59,6 @@ async def get_metadata(self) -> dict: metadata["max_queue_size"] = self.api.max_queue_size return metadata - if self.api.stream_service is not None: - metadata = self.api.stream_service.server_metadata() - metadata["max_queue_size"] = self.api.max_queue_size - metadata.update(self._webrtc_session_stats()) - return metadata - raise HTTPException(status_code=503, detail="No service is initialized") async def health_check(self) -> dict: @@ -92,11 +75,6 @@ async def health_check(self) -> dict: if self.api.inference_service: status["pipeline_ready"] = self.api.inference_service.is_running - if self.api.stream_service: - status["stream_ready"] = self.api.stream_service.is_running - status["stream_mode"] = self.api.stream_service.stream_mode - status.update(self._webrtc_session_stats()) - return status def _is_ready(self) -> bool: @@ -107,9 +85,6 @@ def _is_ready(self) -> bool: return any(replica.get("status") != "dead" for replica in pool_status) return bool(getattr(self.api.inference_service, "is_running", False)) - if self.api.stream_service is not None: - return bool(getattr(self.api.stream_service, "is_running", False)) - return False async def readiness_check(self) -> JSONResponse: @@ -174,7 +149,6 @@ async def get_metrics_json() -> dict: "metrics_count": len(registry.list_metrics()), "registered_stages": registry.list_stages(), } - result["webrtc"] = routes._webrtc_session_stats() return result return new_router diff --git a/telefuser/service/api/routers/stream.py b/telefuser/service/api/routers/stream.py deleted file mode 100644 index fb60d83..0000000 --- a/telefuser/service/api/routers/stream.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Stream routes: session status and lifecycle endpoints. - -Bidirectional streaming uses WebRTC (DataChannel + media tracks). - -Endpoints: - DELETE /v1/stream/sessions/{session_id} – close session - GET /v1/stream/sessions/{session_id}/status -""" - -from __future__ import annotations - -import asyncio -from typing import TYPE_CHECKING - -from fastapi import APIRouter, HTTPException - -from telefuser.utils.logging import logger - -if TYPE_CHECKING: - from ..api_server import ApiServer - - -# --------------------------------------------------------------------------- -# Route handlers -# --------------------------------------------------------------------------- - - -class StreamRoutes: - def __init__(self, api_server: ApiServer) -> None: - self.api = api_server - - def _require_service(self): - svc = self.api.stream_service - if svc is None or not svc.is_running: - raise HTTPException(status_code=503, detail="Stream service is not running") - return svc - - async def close_session(self, session_id: str) -> dict: - svc = self._require_service() - pipeline_closed = False - try: - await asyncio.to_thread(svc.close_session, session_id) - pipeline_closed = True - except Exception as exc: - logger.warning(f"Failed to close pipeline stream session {session_id}: {exc}") - webrtc_closed = False - webrtc_routes = self.api._webrtc_routes - if webrtc_routes is not None: - webrtc_closed = await webrtc_routes._session_manager.close_session( - session_id, - reason="stream_session_delete", - notify_pipeline=False, - ) - if not pipeline_closed and not webrtc_closed: - raise HTTPException(status_code=404, detail=f"Session {session_id} not found") - return {"session_id": session_id, "status": "closed"} - - async def session_status(self, session_id: str) -> dict: - task = self.api.task_manager.get_task_status(session_id) - if task: - return task - - svc = self.api.stream_service - if svc is not None and svc.is_running and svc.has_session(session_id): - return {"session_id": session_id, "status": "active", "stream_mode": svc.stream_mode} - - webrtc_routes = self.api._webrtc_routes - if webrtc_routes is not None and webrtc_routes._session_manager.has_session(session_id): - return {"session_id": session_id, "status": "active", "stream_mode": svc.stream_mode if svc else "unknown"} - - return {"session_id": session_id, "status": "unknown"} - - -# --------------------------------------------------------------------------- -# Router factory -# --------------------------------------------------------------------------- - - -def create_router(api_server: ApiServer) -> APIRouter: - router = APIRouter(prefix="/v1/stream", tags=["stream"]) - routes = StreamRoutes(api_server) - - @router.delete("/sessions/{session_id}", summary="Close session") - async def close_session(session_id: str): - return await routes.close_session(session_id) - - @router.get("/sessions/{session_id}/status", summary="Session status") - async def session_status(session_id: str): - return await routes.session_status(session_id) - - return router - - -def setup_routes(api_server: ApiServer) -> APIRouter: - return create_router(api_server) diff --git a/telefuser/service/api/routers/webrtc.py b/telefuser/service/api/routers/webrtc.py deleted file mode 100644 index 5e2f615..0000000 --- a/telefuser/service/api/routers/webrtc.py +++ /dev/null @@ -1,186 +0,0 @@ -"""WebRTC signaling routes for SDP offer/answer exchange. - -Supports both stream modes: - -* **server_push** – output-only tracks, no DataChannel. -* **bidirectional** – client-created DataChannel + optional media tracks. - -Endpoints: - POST /v1/stream/webrtc/offer – SDP offer → answer - DELETE /v1/stream/webrtc/{session_id} – close session -""" - -from __future__ import annotations - -import asyncio -from typing import TYPE_CHECKING - -from fastapi import APIRouter, HTTPException - -from telefuser.utils.logging import logger - -from ...core.stream_pipeline_service import STREAM_MODE_BIDIRECTIONAL, STREAM_MODE_SERVER_PUSH -from ..stream_schema import WebRTCOfferRequest, WebRTCOfferResponse - -if TYPE_CHECKING: - from ..api_server import ApiServer - - -class WebRTCRoutes: - def __init__(self, api_server: ApiServer) -> None: - self.api = api_server - - from aiortc import RTCConfiguration, RTCIceServer - - from ...webrtc.session_manager import WebRTCSessionManager - - config = api_server.server_config - ice_servers: list[RTCIceServer] = [] - for url in config.stun_servers: - ice_servers.append(RTCIceServer(urls=url)) - if config.turn_server: - ice_servers.append( - RTCIceServer( - urls=config.turn_server, - username=config.turn_username or "", - credential=config.turn_credential or "", - ) - ) - - configuration = RTCConfiguration(iceServers=ice_servers) if ice_servers else None - - self._session_manager = WebRTCSessionManager( - max_sessions=config.webrtc_max_sessions, - configuration=configuration, - video_codec=config.webrtc_video_codec, - video_bitrate=config.webrtc_video_bitrate, - video_buffer_seconds=config.webrtc_video_buffer_seconds, - data_channel_timeout_seconds=config.webrtc_data_channel_timeout_seconds, - disconnected_grace_seconds=config.webrtc_disconnected_grace_seconds, - ) - - @staticmethod - def _resolve_fps(svc: object, requested_fps: int | None) -> int: - """Use an explicit request FPS or the underlying service default.""" - if requested_fps is not None: - return requested_fps - - service = getattr(svc, "service", svc) - default_fps = getattr(service, "default_fps", None) - if default_fps is not None: - return int(default_fps) - return 24 - - async def handle_offer(self, request: WebRTCOfferRequest) -> WebRTCOfferResponse: - svc = self.api.stream_service - if svc is None or not svc.is_running: - raise HTTPException(status_code=503, detail="Stream service is not running") - - if svc.stream_mode == STREAM_MODE_SERVER_PUSH: - return await self._handle_server_push(svc, request) - elif svc.stream_mode == STREAM_MODE_BIDIRECTIONAL: - return await self._handle_bidirectional(svc, request) - else: - raise HTTPException(status_code=500, detail=f"Unknown stream mode: {svc.stream_mode}") - - async def _handle_server_push(self, svc, request: WebRTCOfferRequest) -> WebRTCOfferResponse: - session_id = request.session_id - task_data = request.model_dump(exclude={"sdp", "type"}, exclude_none=True) - task_data["task_id"] = session_id - - generator = svc.stream_task(task_data) - - try: - answer_sdp, answer_type = await self._session_manager.create_session( - session_id=session_id, - offer_sdp=request.sdp, - offer_type=request.type, - generator=generator, - fps=self._resolve_fps(svc, request.fps), - ) - except RuntimeError as exc: - raise HTTPException(status_code=503, detail=str(exc)) - except Exception as exc: - raise HTTPException(status_code=400, detail=f"SDP negotiation failed: {exc}") - - return WebRTCOfferResponse( - session_id=session_id, - sdp=answer_sdp, - type=answer_type, - ) - - async def _handle_bidirectional(self, svc, request: WebRTCOfferRequest) -> WebRTCOfferResponse: - session_id = request.session_id - request_config = request.model_dump( - exclude={"sdp", "type", "config"}, - exclude_none=True, - exclude_unset=True, - ) - config = {**request.config, **request_config, "session_id": session_id} - - try: - pipeline_session_id = svc.create_session(config) - except (OSError, TypeError, ValueError) as exc: - raise HTTPException(status_code=400, detail=f"Session creation failed: {exc}") from exc - except RuntimeError as exc: - raise HTTPException(status_code=409, detail=f"Session creation conflict: {exc}") from exc - - try: - output_gen = svc.pull_chunks(pipeline_session_id) - answer_sdp, answer_type = await self._session_manager.create_bidirectional_session( - session_id=pipeline_session_id, - offer_sdp=request.sdp, - offer_type=request.type, - output_generator=output_gen, - on_input=lambda sid, chunk: svc.push_chunk(sid, chunk), - on_close=lambda sid: svc.close_session(sid), - fps=self._resolve_fps(svc, request.fps if request.fps is not None else config.get("fps")), - ) - except RuntimeError as exc: - try: - await asyncio.to_thread(svc.close_session, pipeline_session_id) - except Exception: - pass - raise HTTPException(status_code=503, detail=str(exc)) from exc - except Exception as exc: - try: - await asyncio.to_thread(svc.close_session, pipeline_session_id) - except Exception: - pass - raise HTTPException(status_code=400, detail=f"SDP negotiation failed: {exc}") from exc - - return WebRTCOfferResponse( - session_id=pipeline_session_id, - sdp=answer_sdp, - type=answer_type, - ) - - async def close_session(self, session_id: str) -> dict: - closed = await self._session_manager.close_session(session_id, reason="webrtc_session_delete") - if not closed: - raise HTTPException(status_code=404, detail=f"WebRTC session {session_id} not found") - return {"session_id": session_id, "status": "closed"} - - async def cleanup(self) -> None: - await self._session_manager.close_all() - - -def create_router(api_server: ApiServer) -> APIRouter: - router = APIRouter(prefix="/v1/stream/webrtc", tags=["webrtc"]) - routes = WebRTCRoutes(api_server) - - api_server._webrtc_routes = routes - - @router.post("/offer", response_model=WebRTCOfferResponse, summary="WebRTC SDP offer/answer signaling") - async def webrtc_offer(request: WebRTCOfferRequest): - return await routes.handle_offer(request) - - @router.delete("/{session_id}", summary="Close WebRTC session") - async def close_webrtc(session_id: str): - return await routes.close_session(session_id) - - return router - - -def setup_routes(api_server: ApiServer) -> APIRouter: - return create_router(api_server) diff --git a/telefuser/service/api/stream_schema.py b/telefuser/service/api/stream_schema.py index d43b902..e908573 100644 --- a/telefuser/service/api/stream_schema.py +++ b/telefuser/service/api/stream_schema.py @@ -4,12 +4,11 @@ import json import time -import uuid from pydantic import BaseModel, Field # --------------------------------------------------------------------------- -# Wire messages (used over DataChannel and internally by WebRTC) +# Stream status messages published through the transport adapter # --------------------------------------------------------------------------- @@ -31,31 +30,6 @@ class StreamDoneMessage(BaseModel): timestamp: float = Field(default_factory=time.time) -# --------------------------------------------------------------------------- -# WebRTC signaling -# --------------------------------------------------------------------------- - - -class WebRTCOfferRequest(BaseModel): - """Body for POST /v1/stream/webrtc/offer.""" - - session_id: str = Field(default_factory=lambda: str(uuid.uuid4())) - sdp: str = Field(description="SDP offer from browser") - type: str = Field(default="offer", description="SDP type") - task: str = Field(description="Task type, e.g. t2v, i2v") - prompt: str | None = None - fps: int | None = Field(default=None, description="Target video FPS") - config: dict = Field(default_factory=dict, description="Session configuration (bidirectional mode)") - - model_config = {"extra": "allow"} - - -class WebRTCOfferResponse(BaseModel): - session_id: str - sdp: str = Field(description="SDP answer from server") - type: str = Field(default="answer", description="SDP type") - - # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- diff --git a/telefuser/service/core/config.py b/telefuser/service/core/config.py index 9c3c2ac..8d07a82 100644 --- a/telefuser/service/core/config.py +++ b/telefuser/service/core/config.py @@ -187,37 +187,6 @@ class ServerConfig(BaseSettings): default="auto", description="GPU platform for metrics collection" ) - # Stream settings - webrtc_max_sessions: int = Field(default=10, ge=1, le=100, description="Maximum concurrent WebRTC sessions") - webrtc_video_codec: Literal["H264", "VP8"] = Field( - default="H264", - description="Preferred WebRTC video codec. Other supported codecs remain available as fallbacks.", - ) - webrtc_video_bitrate: int = Field( - default=8_000_000, - ge=500_000, - le=50_000_000, - description="Target WebRTC video bitrate in bits per second.", - ) - webrtc_video_buffer_seconds: float = Field( - default=1.0, - ge=0.1, - le=10.0, - description="Maximum buffered WebRTC video duration before oldest frames are dropped.", - ) - webrtc_data_channel_timeout_seconds: float = Field( - default=10.0, - gt=0, - le=120.0, - description="Time to wait for the required bidirectional WebRTC DataChannel to open.", - ) - webrtc_disconnected_grace_seconds: float = Field( - default=5.0, - ge=0, - le=120.0, - description="Grace period for transient WebRTC disconnected states before the session is closed.", - ) - # Pipeline replication settings num_replicas: int = Field( default=1, @@ -226,15 +195,6 @@ class ServerConfig(BaseSettings): description="Number of independent pipeline replicas for concurrent serving.", ) - # WebRTC ICE settings (for public network deployment) - stun_servers: list[str] = Field( - default_factory=lambda: ["stun:stun.l.google.com:19302"], - description="STUN server URLs (e.g. stun:stun.l.google.com:19302)", - ) - turn_server: str | None = Field(default=None, description="TURN server URL (e.g. turn:your-domain.com:3478)") - turn_username: str | None = Field(default=None, description="TURN server username") - turn_credential: str | None = Field(default=None, description="TURN server credential") - @field_validator("port") @classmethod def validate_port(cls: type[ServerConfig], v: int) -> int: diff --git a/telefuser/service/core/container.py b/telefuser/service/core/container.py index c04f607..631a61e 100644 --- a/telefuser/service/core/container.py +++ b/telefuser/service/core/container.py @@ -20,7 +20,6 @@ from .file_service import FileService from .pipeline_pool import PipelinePool from .pipeline_service import PipelineService -from .stream_pipeline_service import StreamPipelineService from .task_manager import TaskManager from .task_service import MediaGenerationService @@ -39,7 +38,6 @@ class ServiceContainer: task_manager: TaskManager file_service: FileService | None = None pipeline_service: PipelineService | None = None - stream_pipeline_service: StreamPipelineService | None = None media_service: MediaGenerationService | None = None cache_service: Any | None = None cache_adapter: Any | None = None # cacheseek.adapters.telefuser.TeleFuserCacheAdapter @@ -246,26 +244,8 @@ def initialize_all( return True - def initialize_stream_service( - self, - pipe_path: str, - skip_validation: bool = False, - gpu_num: int = 1, - ) -> bool: - """Initialize stream pipeline service (alternative to initialize_all).""" - self.stream_pipeline_service = StreamPipelineService( - security_level=self.config.security_level, - config=self.config, - ) - return self.stream_pipeline_service.start_service( - ppl_file=pipe_path, - gpu_num=gpu_num, - skip_validation=skip_validation, - ) - def get_api_app(self, enable_rate_limit: bool = True) -> FastAPI: """Get FastAPI application with all services initialized.""" - route_profile = "stream" if self.stream_pipeline_service and not self.pipeline_service else "request_response" api_server = ApiServer( max_queue_size=self.config.max_queue_size, max_concurrent_tasks=self.config.effective_max_concurrent_tasks, @@ -274,7 +254,6 @@ def get_api_app(self, enable_rate_limit: bool = True) -> FastAPI: enable_rate_limit=enable_rate_limit, enable_logging=False, config=self.config, - route_profile=route_profile, ) if self.file_service and self.pipeline_service: @@ -286,9 +265,6 @@ def get_api_app(self, enable_rate_limit: bool = True) -> FastAPI: file_service=self.file_service, ) - if self.stream_pipeline_service: - api_server.initialize_stream_service(self.stream_pipeline_service) - return api_server.get_app() async def __aenter__(self) -> ServiceContainer: @@ -314,10 +290,6 @@ async def cleanup(self) -> None: await self.pipeline_service.aclose() self.pipeline_service = None - if self.stream_pipeline_service: - await self.stream_pipeline_service.aclose() - self.stream_pipeline_service = None - if self.cache_service is not None: try: self.cache_service.shutdown() diff --git a/telefuser/service/core/stream_pipeline_service.py b/telefuser/service/core/stream_pipeline_service.py index 0ad318a..a5db874 100644 --- a/telefuser/service/core/stream_pipeline_service.py +++ b/telefuser/service/core/stream_pipeline_service.py @@ -6,8 +6,8 @@ Two interaction modes are supported: -* SERVER_PUSH – single request in, continuous chunks out (WebRTC media tracks) -* BIDIRECTIONAL – continuous input & output (WebRTC DataChannel + media tracks) +* SERVER_PUSH – single request in, continuous chunks out +* BIDIRECTIONAL – continuous input and output Pipeline ``serve()`` methods may contain blocking calls (GPU inference, ``time.sleep``, etc.). ``stream_task()`` runs them on a dedicated thread @@ -220,7 +220,7 @@ async def stream_task(self, task_data: dict) -> AsyncGenerator[dict, None]: etc.) do not stall the server's main loop. A ``threading.Event`` stop flag is set when the consumer stops - iterating (e.g. WebRTC disconnect), so the producer thread can + iterating (for example after a transport disconnect), so the producer thread can break out of the pipeline's ``serve()`` loop promptly instead of running to completion. """ diff --git a/telefuser/service/main.py b/telefuser/service/main.py index 2bd88ed..d1b7771 100644 --- a/telefuser/service/main.py +++ b/telefuser/service/main.py @@ -85,37 +85,3 @@ def run_server( logger.info("All services initialized successfully") _run("server", container, enable_rate_limit) - - -def run_stream_server( - pipe_path: str, - port: int, - host: str, - enable_rate_limit: bool = True, - skip_validation: bool = False, - security_level: str | None = None, - gpu_num: int = 1, -) -> None: - """Run the TeleFuser stream server. - - Unlike run_server (request-response), this loads a stream pipeline - that exposes get_service() and serves via WebRTC or WebSocket. - """ - server_config.host = host - server_config.port = port - if security_level is not None: - from .security.security_validator import SecurityLevel - - server_config.security_level = SecurityLevel[security_level.upper()] - - container = ServiceContainer.create(config=server_config) - - if not container.initialize_stream_service( - pipe_path=pipe_path, - gpu_num=gpu_num, - skip_validation=skip_validation, - ): - raise RuntimeError("Failed to initialize stream service") - - logger.info("Stream service initialized successfully") - _run("stream server", container, enable_rate_limit) diff --git a/telefuser/service/webrtc/__init__.py b/telefuser/service/webrtc/__init__.py deleted file mode 100644 index fd732b3..0000000 --- a/telefuser/service/webrtc/__init__.py +++ /dev/null @@ -1,32 +0,0 @@ -"""WebRTC transport for TeleFuser stream server. - -The default TeleFuser installation includes aiortc for WebRTC support. -""" - -from __future__ import annotations - -try: - from .chunk_router import ChunkRouter - from .session_manager import WebRTCSessionManager - from .track import ( - AudioGeneratorTrack, - FrameGeneratorTrack, - IncomingAudioRelay, - IncomingVideoRelay, - ) -except ImportError: - AudioGeneratorTrack = None # type: ignore[assignment,misc] - ChunkRouter = None # type: ignore[assignment,misc] - FrameGeneratorTrack = None # type: ignore[assignment,misc] - IncomingAudioRelay = None # type: ignore[assignment,misc] - IncomingVideoRelay = None # type: ignore[assignment,misc] - WebRTCSessionManager = None # type: ignore[assignment,misc] - -__all__ = [ - "AudioGeneratorTrack", - "ChunkRouter", - "FrameGeneratorTrack", - "IncomingAudioRelay", - "IncomingVideoRelay", - "WebRTCSessionManager", -] diff --git a/telefuser/service/webrtc/chunk_router.py b/telefuser/service/webrtc/chunk_router.py deleted file mode 100644 index 32c98b7..0000000 --- a/telefuser/service/webrtc/chunk_router.py +++ /dev/null @@ -1,143 +0,0 @@ -"""Fan-out adapter: consumes one output generator and routes to tracks + DataChannel. - -The ``ChunkRouter`` reads chunks from a ``BidirectionalService.pull_chunks()`` -async generator exactly once and dispatches them: - -* ``frames_b64`` / ``audio_b64`` → decoded and pushed to the outgoing - ``FrameGeneratorTrack`` / ``AudioGeneratorTrack`` (RTP media) -* Remaining metadata fields → serialised to JSON and sent over the - client-created DataChannel as ``StreamChunkMessage`` / ``StreamDoneMessage`` - -This avoids double-consuming the generator and prevents duplicate sends. -""" - -from __future__ import annotations - -import asyncio -import base64 -from collections.abc import AsyncGenerator, Callable - -import av -import cv2 -import numpy as np -from PIL import Image - -from telefuser.service.api.stream_schema import StreamChunkMessage, StreamDoneMessage, serialisable_chunk -from telefuser.utils.logging import logger - -_MEDIA_KEYS = frozenset({"frames_b64", "audio_b64", "audio_sample_rate", "audio_channels"}) - - -class ChunkRouter: - """Consumes an output generator once, routes media to tracks and metadata to DataChannel.""" - - def __init__( - self, - generator: AsyncGenerator[dict, None], - video_track: object | None, - audio_track: object | None, - data_channel_send: Callable[[str], None] | None, - session_id: str, - on_complete: Callable[[str], None] | None = None, - ) -> None: - self._generator = generator - self._video_track = video_track - self._audio_track = audio_track - self._dc_send = data_channel_send - self._session_id = session_id - self._on_complete = on_complete - self._chunk_count = 0 - - async def run(self) -> None: - """Main loop: consume generator, dispatch chunks.""" - cancelled = False - try: - async for chunk in self._generator: - if self._route_chunk(chunk): - self._chunk_count += 1 - except asyncio.CancelledError: - cancelled = True - raise - except Exception as exc: - logger.error(f"ChunkRouter error: session={self._session_id} {exc}") - finally: - self._send_done() - if self._video_track is not None: - self._video_track.signal_done() - if self._audio_track is not None: - self._audio_track.signal_done() - if not cancelled and self._on_complete is not None: - self._on_complete(self._session_id) - - def _route_chunk(self, chunk: dict) -> bool: - is_nested = "frames_b64" not in chunk and isinstance(chunk.get("data"), dict) - data = chunk.get("data", {}) if is_nested else chunk - - raw_frames = data.get("frames") - frames: list[object] | tuple[object, ...] = raw_frames if isinstance(raw_frames, (list, tuple)) else () - frames_b64: list[str] = data.get("frames_b64", []) - is_video_chunk = bool(frames or frames_b64) - if self._video_track is not None: - if frames: - converted_frames: list[av.VideoFrame] = [] - for source in frames: - if isinstance(source, av.VideoFrame): - frame = source - elif isinstance(source, Image.Image): - rgb = np.ascontiguousarray(source.convert("RGB")) - frame = av.VideoFrame.from_ndarray(rgb, format="rgb24") - else: - logger.warning(f"ChunkRouter dropped unsupported raw video frame: session={self._session_id}") - continue - converted_frames.append(frame) - push_frames = getattr(type(self._video_track), "push_frames", None) - if callable(push_frames): - self._video_track.push_frames(converted_frames) - else: - for frame in converted_frames: - self._video_track.push_frame(frame) - else: - for fb64 in frames_b64: - raw = base64.b64decode(fb64) - np_arr = np.frombuffer(raw, dtype=np.uint8) - bgr = cv2.imdecode(np_arr, cv2.IMREAD_COLOR) - if bgr is None: - logger.warning(f"ChunkRouter dropped undecodable video frame: session={self._session_id}") - continue - rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) - frame = av.VideoFrame.from_ndarray(rgb, format="rgb24") - self._video_track.push_frame(frame) - - audio_b64 = data.get("audio_b64") - if audio_b64 and self._audio_track is not None: - self._audio_track.feed(base64.b64decode(audio_b64)) - - media_keys = _MEDIA_KEYS | {"frames"} if isinstance(raw_frames, (list, tuple)) else _MEDIA_KEYS - metadata = {k: v for k, v in chunk.items() if k not in media_keys} - if is_nested and isinstance(metadata.get("data"), dict): - metadata["data"] = {k: v for k, v in metadata["data"].items() if k not in media_keys} - if not metadata["data"]: - del metadata["data"] - if metadata and self._dc_send is not None: - msg = StreamChunkMessage( - session_id=self._session_id, - index=chunk.get("index"), - data=serialisable_chunk(metadata), - ) - try: - self._dc_send(msg.model_dump_json()) - except Exception as exc: - logger.warning(f"ChunkRouter metadata send failed: session={self._session_id} {exc}") - return is_video_chunk - - def _send_done(self) -> None: - if self._dc_send is None: - return - done = StreamDoneMessage( - session_id=self._session_id, - total_chunks=self._chunk_count, - ) - try: - self._dc_send(done.model_dump_json()) - except Exception as exc: - logger.warning(f"ChunkRouter done send failed: session={self._session_id} {exc}") diff --git a/telefuser/service/webrtc/session_manager.py b/telefuser/service/webrtc/session_manager.py deleted file mode 100644 index 092fb68..0000000 --- a/telefuser/service/webrtc/session_manager.py +++ /dev/null @@ -1,469 +0,0 @@ -"""WebRTC session manager — tracks RTCPeerConnection lifecycle. - -Supports two session types: - -* **_Session** (server-push): output-only video/audio tracks, no DataChannel. -* **_BidirectionalSession**: DataChannel for JSON control, optional incoming - and outgoing media tracks, ChunkRouter for fan-out. -""" - -from __future__ import annotations - -import asyncio -import json -from collections.abc import AsyncGenerator, Callable -from dataclasses import dataclass, field - -from aiortc import RTCConfiguration, RTCPeerConnection, RTCRtpSender, RTCSessionDescription - -from telefuser.utils.logging import logger - -from .chunk_router import ChunkRouter -from .track import ( - AudioGeneratorTrack, - FrameGeneratorTrack, - IncomingAudioRelay, - IncomingVideoRelay, -) - -_SENTINEL = object() - - -@dataclass -class _Session: - pc: RTCPeerConnection - track: FrameGeneratorTrack - audio_track: AudioGeneratorTrack | None = None - - -@dataclass -class _BidirectionalSession: - pc: RTCPeerConnection - on_close: Callable[[str], None] | None = None - session_id: str = "" - data_channel: object | None = None - output_video_track: FrameGeneratorTrack | None = None - output_audio_track: AudioGeneratorTrack | None = None - chunk_router: ChunkRouter | None = None - router_task: asyncio.Task | None = None - relay_tasks: list[asyncio.Task] = field(default_factory=list) - data_channel_timeout_task: asyncio.Task | None = None - disconnect_task: asyncio.Task | None = None - - -class WebRTCSessionManager: - """Creates and manages WebRTC peer connections for stream sessions.""" - - def __init__( - self, - max_sessions: int = 10, - configuration: RTCConfiguration | None = None, - video_codec: str = "H264", - video_bitrate: int = 8_000_000, - video_buffer_seconds: float = 1.0, - terminal_grace_seconds: float = 0.1, - close_on_output_complete: bool = False, - pipeline_close_timeout: float = 30.0, - data_channel_timeout_seconds: float = 10.0, - disconnected_grace_seconds: float = 5.0, - ) -> None: - if video_bitrate < 500_000: - raise ValueError(f"video_bitrate must be at least 500000, got {video_bitrate}") - if video_buffer_seconds <= 0: - raise ValueError("video_buffer_seconds must be positive") - video_codec = video_codec.upper() - if video_codec not in {"H264", "VP8"}: - raise ValueError(f"Unsupported WebRTC video codec: {video_codec}") - if terminal_grace_seconds < 0: - raise ValueError("terminal_grace_seconds must be non-negative") - if pipeline_close_timeout <= 0: - raise ValueError("pipeline_close_timeout must be positive") - if data_channel_timeout_seconds <= 0: - raise ValueError("data_channel_timeout_seconds must be positive") - if disconnected_grace_seconds < 0: - raise ValueError("disconnected_grace_seconds must be non-negative") - self._sessions: dict[str, _Session | _BidirectionalSession | object] = {} - self._max_sessions = max_sessions - self._configuration = configuration - self._video_codec = video_codec - self._video_bitrate = video_bitrate - self._video_buffer_seconds = video_buffer_seconds - self._terminal_grace_seconds = terminal_grace_seconds - self._close_on_output_complete = close_on_output_complete - self._pipeline_close_timeout = pipeline_close_timeout - self._data_channel_timeout_seconds = data_channel_timeout_seconds - self._disconnected_grace_seconds = disconnected_grace_seconds - self._lock = asyncio.Lock() - self._configure_video_bitrate(video_bitrate) - - def _schedule_output_complete(self, session_id: str) -> None: - """Optionally close a completed stream after its terminal message flushes.""" - if not self._close_on_output_complete: - return - asyncio.create_task(self._close_after_output_complete(session_id)) - - async def _close_after_output_complete(self, session_id: str) -> None: - if self._terminal_grace_seconds: - await asyncio.sleep(self._terminal_grace_seconds) - await self.close_session(session_id, reason="output_complete") - - async def _close_if_data_channel_missing(self, session_id: str) -> None: - """Release a reserved pipeline session when the required channel never opens.""" - await asyncio.sleep(self._data_channel_timeout_seconds) - entry = self._sessions.get(session_id) - if isinstance(entry, _BidirectionalSession) and entry.data_channel is None: - logger.warning(f"WebRTC DataChannel did not open before timeout: session={session_id}") - await self.close_session(session_id, reason="data_channel_timeout") - - async def _close_after_disconnect_grace(self, session_id: str) -> None: - """Allow transient ICE disconnects to recover before releasing the session.""" - if self._disconnected_grace_seconds: - await asyncio.sleep(self._disconnected_grace_seconds) - entry = self._sessions.get(session_id) - if isinstance(entry, _BidirectionalSession) and entry.pc.connectionState == "disconnected": - logger.info(f"WebRTC disconnect grace expired: session={session_id}") - await self.close_session(session_id, reason="connection_disconnected") - - @staticmethod - def _configure_video_bitrate(video_bitrate: int) -> None: - """Configure aiortc software encoders before their lazy construction.""" - from aiortc.codecs import h264, vpx - - h264.DEFAULT_BITRATE = video_bitrate - h264.MAX_BITRATE = video_bitrate - vpx.DEFAULT_BITRATE = video_bitrate - vpx.MAX_BITRATE = video_bitrate - - def _set_video_codec_preferences(self, pc: RTCPeerConnection) -> None: - """Prefer the configured codec while retaining interoperable fallbacks.""" - codecs = RTCRtpSender.getCapabilities("video").codecs - preferred_mime = f"video/{self._video_codec}".lower() - preferred = [codec for codec in codecs if codec.mimeType.lower() == preferred_mime] - remaining = [codec for codec in codecs if codec.mimeType.lower() not in {preferred_mime, "video/rtx"}] - rtx = [codec for codec in codecs if codec.mimeType.lower() == "video/rtx"] - ordered = [*preferred, *remaining, *rtx] - for transceiver in pc.getTransceivers(): - if transceiver.kind == "video": - transceiver.setCodecPreferences(ordered) - - # -- Server-push sessions ------------------------------------------------ - - async def create_session( - self, - session_id: str, - offer_sdp: str, - offer_type: str, - generator: AsyncGenerator[dict, None], - fps: int = 24, - ) -> tuple[str, str]: - """Process SDP offer and return (answer_sdp, answer_type).""" - async with self._lock: - if len(self._sessions) >= self._max_sessions: - raise RuntimeError(f"Max WebRTC sessions ({self._max_sessions}) reached") - if session_id in self._sessions: - raise RuntimeError(f"Session {session_id} already exists") - self._sessions[session_id] = _SENTINEL - - try: - pc = RTCPeerConnection(configuration=self._configuration or RTCConfiguration()) - - has_audio_offer = "m=audio" in offer_sdp - audio_track: AudioGeneratorTrack | None = None - if has_audio_offer: - audio_track = AudioGeneratorTrack() - - track = FrameGeneratorTrack( - generator, - fps=fps, - audio_track=audio_track, - max_buffer_seconds=self._video_buffer_seconds, - ) - - @pc.on("connectionstatechange") - async def _on_state_change() -> None: - state = pc.connectionState - if state in ("failed", "disconnected", "closed"): - logger.info(f"WebRTC connection {state}: session={session_id}") - await self.close_session(session_id, reason=f"connection_{state}") - - pc.addTrack(track) - self._set_video_codec_preferences(pc) - if audio_track is not None: - pc.addTrack(audio_track) - - offer = RTCSessionDescription(sdp=offer_sdp, type=offer_type) - await pc.setRemoteDescription(offer) - answer = await pc.createAnswer() - await pc.setLocalDescription(answer) - - async with self._lock: - self._sessions[session_id] = _Session(pc=pc, track=track, audio_track=audio_track) - - audio_status = " + audio" if audio_track else "" - logger.info(f"WebRTC session created: session={session_id}{audio_status}") - return pc.localDescription.sdp, pc.localDescription.type - except BaseException: - async with self._lock: - self._sessions.pop(session_id, None) - raise - - # -- Bidirectional sessions ---------------------------------------------- - - async def create_bidirectional_session( - self, - session_id: str, - offer_sdp: str, - offer_type: str, - output_generator: AsyncGenerator[dict, None], - on_input: Callable[[str, dict], None], - on_close: Callable[[str], None], - fps: int = 24, - ) -> tuple[str, str]: - """Create a bidirectional WebRTC session. - - The client must create a DataChannel named ``"telefuser"`` before - generating the SDP offer. The server reuses that single channel - for both reading input and writing output. - - Args: - session_id: Unique session identifier (from pipeline). - offer_sdp: Client SDP offer string. - offer_type: SDP type (usually ``"offer"``). - output_generator: Async generator from ``pull_chunks()``. - on_input: Callback ``(session_id, chunk_dict)`` for incoming data. - on_close: Callback ``(session_id)`` when client sends ``stop``. - fps: Target video FPS for outgoing media tracks. - - Returns: - ``(answer_sdp, answer_type)`` tuple. - """ - async with self._lock: - if len(self._sessions) >= self._max_sessions: - raise RuntimeError(f"Max WebRTC sessions ({self._max_sessions}) reached") - if session_id in self._sessions: - raise RuntimeError(f"Session {session_id} already exists") - self._sessions[session_id] = _SENTINEL - - try: - pc = RTCPeerConnection(configuration=self._configuration or RTCConfiguration()) - session = _BidirectionalSession(pc=pc, on_close=on_close, session_id=session_id) - async with self._lock: - self._sessions[session_id] = session - - has_audio_offer = "m=audio" in offer_sdp - output_video = FrameGeneratorTrack( - generator=None, - fps=fps, - max_buffer_seconds=self._video_buffer_seconds, - ) - output_audio: AudioGeneratorTrack | None = None - if has_audio_offer: - output_audio = AudioGeneratorTrack() - - session.output_video_track = output_video - session.output_audio_track = output_audio - - pc.addTrack(output_video) - self._set_video_codec_preferences(pc) - if output_audio is not None: - pc.addTrack(output_audio) - - # --- DataChannel (client-created, server reuses) ---------------- - - @pc.on("datachannel") - def _on_datachannel(channel) -> None: - if channel.label != "telefuser": - logger.warning(f"Ignoring unexpected DataChannel: session={session_id} label={channel.label}") - channel.close() - return - if session.data_channel is not None: - logger.warning(f"Ignoring duplicate DataChannel: session={session_id}") - channel.close() - return - session.data_channel = channel - if session.data_channel_timeout_task is not None: - session.data_channel_timeout_task.cancel() - session.data_channel_timeout_task = None - logger.info(f"DataChannel received: session={session_id} label={channel.label}") - - @channel.on("close") - def _on_channel_close() -> None: - asyncio.ensure_future(self.close_session(session_id, reason="data_channel_closed")) - - @channel.on("message") - def _on_message(message) -> None: - try: - data = json.loads(message) if isinstance(message, str) else message - except (json.JSONDecodeError, TypeError) as exc: - logger.warning(f"DataChannel message decode failed: session={session_id} {exc}") - return - if isinstance(data, dict) and data.get("type") == "stop": - asyncio.ensure_future(self.close_session(session_id, reason="client_stop")) - return - try: - on_input(session_id, data) - except Exception as exc: - logger.warning(f"DataChannel input callback failed: session={session_id} {exc}") - - router = ChunkRouter( - generator=output_generator, - video_track=output_video, - audio_track=output_audio, - data_channel_send=channel.send, - session_id=session_id, - on_complete=self._schedule_output_complete if self._close_on_output_complete else None, - ) - session.chunk_router = router - session.router_task = asyncio.ensure_future(router.run()) - - # --- Incoming media tracks (optional) --------------------------- - - @pc.on("track") - def _on_track(track) -> None: - logger.info(f"Incoming track: session={session_id} kind={track.kind}") - if track.kind == "video": - relay = IncomingVideoRelay(track, session_id, on_input) - task = asyncio.ensure_future(relay.run()) - session.relay_tasks.append(task) - elif track.kind == "audio": - relay = IncomingAudioRelay(track, session_id, on_input) - task = asyncio.ensure_future(relay.run()) - session.relay_tasks.append(task) - - # --- Connection state ------------------------------------------- - - @pc.on("connectionstatechange") - async def _on_state_change() -> None: - state = pc.connectionState - if state in ("failed", "closed"): - logger.info(f"WebRTC connection {state}: session={session_id}") - await self.close_session(session_id, reason=f"connection_{state}") - elif state == "disconnected": - if session.disconnect_task is None or session.disconnect_task.done(): - session.disconnect_task = asyncio.create_task(self._close_after_disconnect_grace(session_id)) - elif state == "connected" and session.disconnect_task is not None: - session.disconnect_task.cancel() - session.disconnect_task = None - - # --- SDP exchange ----------------------------------------------- - - offer = RTCSessionDescription(sdp=offer_sdp, type=offer_type) - await pc.setRemoteDescription(offer) - answer = await pc.createAnswer() - await pc.setLocalDescription(answer) - - session.data_channel_timeout_task = asyncio.create_task(self._close_if_data_channel_missing(session_id)) - - logger.info(f"WebRTC bidirectional session created: session={session_id}") - return pc.localDescription.sdp, pc.localDescription.type - except BaseException: - async with self._lock: - self._sessions.pop(session_id, None) - raise - - # -- Session lifecycle --------------------------------------------------- - - async def close_session(self, session_id: str, *, reason: str = "api", notify_pipeline: bool = True) -> bool: - async with self._lock: - entry = self._sessions.pop(session_id, None) - if entry is None or entry is _SENTINEL: - return False - - logger.info(f"Closing WebRTC session: session={session_id} reason={reason}") - - if isinstance(entry, _Session): - entry.track.stop() - if entry.track._task is not None and not entry.track._task.done(): - try: - await entry.track._task - except asyncio.CancelledError: - logger.info(f"WebRTC server-push track cancelled: session={session_id}") - except Exception as exc: - logger.warning(f"WebRTC server-push track close failed: session={session_id} {exc}") - elif isinstance(entry, _BidirectionalSession): - current_task = asyncio.current_task() - for task in (entry.data_channel_timeout_task, entry.disconnect_task): - if task is not None and task is not current_task and not task.done(): - task.cancel() - for task in entry.relay_tasks: - if not task.done(): - task.cancel() - try: - await task - except asyncio.CancelledError: - logger.info(f"WebRTC relay task cancelled: session={session_id}") - except Exception as exc: - logger.warning(f"WebRTC relay task close failed: session={session_id} {exc}") - if entry.router_task is not None and not entry.router_task.done(): - entry.router_task.cancel() - try: - await entry.router_task - except asyncio.CancelledError: - logger.info(f"WebRTC chunk router cancelled: session={session_id}") - except Exception as exc: - logger.warning(f"WebRTC chunk router close failed: session={session_id} {exc}") - if entry.output_video_track is not None: - entry.output_video_track.stop() - if entry.output_audio_track is not None: - entry.output_audio_track.stop() - if notify_pipeline and entry.on_close is not None: - try: - await asyncio.wait_for( - asyncio.to_thread(entry.on_close, entry.session_id), - timeout=self._pipeline_close_timeout, - ) - except asyncio.TimeoutError: - logger.warning(f"WebRTC pipeline close callback timed out: session={session_id}") - except Exception as exc: - logger.warning(f"WebRTC pipeline close callback failed: session={session_id} {exc}") - - try: - await asyncio.wait_for(entry.pc.close(), timeout=5.0) - except asyncio.TimeoutError: - logger.warning(f"Timed out closing WebRTC peer connection: session={session_id}") - except Exception as exc: - logger.warning(f"WebRTC peer connection close failed: session={session_id} {exc}") - logger.info(f"WebRTC session closed: session={session_id} reason={reason}") - return True - - async def close_all(self) -> None: - async with self._lock: - session_ids = list(self._sessions.keys()) - for sid in session_ids: - await self.close_session(sid) - - def has_session(self, session_id: str) -> bool: - entry = self._sessions.get(session_id) - return entry is not None and entry is not _SENTINEL - - @property - def active_session_count(self) -> int: - return sum(1 for v in self._sessions.values() if v is not _SENTINEL) - - @property - def server_push_session_count(self) -> int: - return sum(1 for v in self._sessions.values() if isinstance(v, _Session)) - - @property - def bidirectional_session_count(self) -> int: - return sum(1 for v in self._sessions.values() if isinstance(v, _BidirectionalSession)) - - def session_stats(self) -> dict: - """Single-pass session counts (avoids 3 iterations over the dict).""" - active = server_push = bidirectional = 0 - for v in self._sessions.values(): - if v is _SENTINEL: - continue - active += 1 - if isinstance(v, _Session): - server_push += 1 - elif isinstance(v, _BidirectionalSession): - bidirectional += 1 - return { - "webrtc_active_sessions": active, - "webrtc_server_push_sessions": server_push, - "webrtc_bidirectional_sessions": bidirectional, - "webrtc_max_sessions": self._max_sessions, - "webrtc_video_codec": self._video_codec, - "webrtc_video_bitrate": self._video_bitrate, - } diff --git a/telefuser/service/webrtc/track.py b/telefuser/service/webrtc/track.py deleted file mode 100644 index 6fb590e..0000000 --- a/telefuser/service/webrtc/track.py +++ /dev/null @@ -1,343 +0,0 @@ -"""WebRTC media tracks — bridges async generators to aiortc. - -FrameGeneratorTrack: decodes ``frames_b64`` → ``av.VideoFrame`` at target fps. - In server-push mode it owns an async generator; in bidirectional mode frames - are pushed directly via ``push_frame()``. -AudioGeneratorTrack: receives raw PCM16 bytes → ``av.AudioFrame`` at 20ms pacing. -IncomingVideoRelay / IncomingAudioRelay: consume incoming client media tracks and - forward decoded frames as native Python objects to a callback. -""" - -from __future__ import annotations - -import asyncio -import base64 -import fractions -import time -from collections.abc import AsyncGenerator, Callable - -import av -import cv2 -import numpy as np -from aiortc import MediaStreamTrack -from aiortc.mediastreams import MediaStreamError - -from telefuser.utils.logging import logger - -_RTP_CLOCK_RATE = 90_000 -AUDIO_SAMPLE_RATE = 48_000 -_AUDIO_SAMPLES_PER_FRAME = 960 # 20ms at 48kHz — standard Opus frame - - -class AudioGeneratorTrack(MediaStreamTrack): - """Audio track fed by raw PCM16 bytes pushed from the video track.""" - - kind = "audio" - - def __init__( - self, - sample_rate: int = AUDIO_SAMPLE_RATE, - channels: int = 1, - samples_per_frame: int = _AUDIO_SAMPLES_PER_FRAME, - ) -> None: - super().__init__() - self._sample_rate = sample_rate - self._channels = channels - self._samples_per_frame = samples_per_frame - self._bytes_per_frame = samples_per_frame * channels * 2 # int16 = 2 bytes - self._frame_duration = samples_per_frame / sample_rate - - self._queue: asyncio.Queue[bytes] = asyncio.Queue(maxsize=200) - self._buffer = bytearray() - self._buf_offset = 0 - self._frame_count = 0 - self._start_time: float | None = None - self._finished = False - self.dropped_chunks = 0 - - def feed(self, data: bytes) -> bool: - """Push raw PCM16 bytes (called by FrameGeneratorTrack from its consumer task).""" - try: - self._queue.put_nowait(data) - return True - except asyncio.QueueFull: - self.dropped_chunks += 1 - if self.dropped_chunks == 1 or self.dropped_chunks % 100 == 0: - logger.warning(f"AudioGeneratorTrack queue full; dropped_chunks={self.dropped_chunks}") - return False - - def signal_done(self) -> None: - self._finished = True - - async def recv(self) -> av.AudioFrame: - if self._start_time is None: - self._start_time = time.time() - - target_time = self._start_time + self._frame_count * self._frame_duration - wait = target_time - time.time() - if wait > 0: - await asyncio.sleep(wait) - - buf_avail = len(self._buffer) - self._buf_offset - while buf_avail < self._bytes_per_frame: - if self._finished and self._queue.empty(): - if buf_avail > 0: - self._buffer.extend(b"\x00" * (self._bytes_per_frame - buf_avail)) - buf_avail = self._bytes_per_frame - break - raise MediaStreamError("Audio track ended") - try: - data = await asyncio.wait_for(self._queue.get(), timeout=2.0) - self._buffer.extend(data) - buf_avail = len(self._buffer) - self._buf_offset - except asyncio.TimeoutError: - if self._finished: - raise MediaStreamError("Audio track ended") - self._buffer.extend(b"\x00" * self._bytes_per_frame) - buf_avail = len(self._buffer) - self._buf_offset - break - - end = self._buf_offset + self._bytes_per_frame - pcm_bytes = bytes(self._buffer[self._buf_offset : end]) - self._buf_offset = end - if self._buf_offset > 64_000: - del self._buffer[: self._buf_offset] - self._buf_offset = 0 - - frame = av.AudioFrame(format="s16", layout="mono", samples=self._samples_per_frame) - frame.planes[0].update(pcm_bytes) - frame.sample_rate = self._sample_rate - frame.pts = self._frame_count * self._samples_per_frame - frame.time_base = fractions.Fraction(1, self._sample_rate) - self._frame_count += 1 - return frame - - -class FrameGeneratorTrack(MediaStreamTrack): - """Video track that delivers ``av.VideoFrame`` at a constant FPS. - - Two modes of operation: - - * **Generator mode** (server-push): pass an async generator that yields - chunk dicts with ``frames_b64``. The track spawns a consumer task. - * **Push mode** (bidirectional): pass ``generator=None``. The caller - feeds frames directly via :meth:`push_frame`. - """ - - kind = "video" - - def __init__( - self, - generator: AsyncGenerator[dict, None] | None = None, - fps: int = 24, - audio_track: AudioGeneratorTrack | None = None, - max_buffer_seconds: float = 1.0, - ) -> None: - if fps <= 0: - raise ValueError("fps must be positive") - if max_buffer_seconds <= 0: - raise ValueError("max_buffer_seconds must be positive") - super().__init__() - self._generator = generator - self._fps = fps - self._frame_interval = 1.0 / fps - self._pts_per_frame = _RTP_CLOCK_RATE // fps - self._audio_track = audio_track - self._queue: asyncio.Queue[av.VideoFrame] = asyncio.Queue(maxsize=max(1, round(fps * max_buffer_seconds))) - self._frame_count = 0 - self._task: asyncio.Task | None = None - self._finished = False - self._start_time: float | None = None - self._last_frame: av.VideoFrame | None = None - self._placeholder_width = 640 - self._placeholder_height = 360 - self.dropped_frames = 0 - - def push_frame(self, frame: av.VideoFrame) -> bool: - """Push one frame, evicting the oldest frame when the client falls behind.""" - return self.push_frames((frame,)) - - def push_frames(self, frames: tuple[av.VideoFrame, ...] | list[av.VideoFrame]) -> bool: - """Push a video batch while preserving the newest possible playback window.""" - if not frames: - return True - for frame in frames: - self._placeholder_width = frame.width - self._placeholder_height = frame.height - - capacity = self._queue.maxsize - accepted = list(frames[-capacity:]) - dropped = len(frames) - len(accepted) - while self._queue.qsize() + len(accepted) > capacity: - try: - self._queue.get_nowait() - dropped += 1 - except asyncio.QueueEmpty: # pragma: no cover - qsize is advisory - break - for frame in accepted: - self._queue.put_nowait(frame) - if dropped: - self.dropped_frames += dropped - if self.dropped_frames == dropped or self.dropped_frames % 100 < dropped: - logger.warning(f"FrameGeneratorTrack dropped_old_frames={self.dropped_frames}") - return dropped == 0 - - def signal_done(self) -> None: - self._finished = True - - def _make_placeholder_frame(self) -> av.VideoFrame: - image = np.zeros((self._placeholder_height, self._placeholder_width, 3), dtype=np.uint8) - return av.VideoFrame.from_ndarray(image, format="rgb24") - - async def _consume_generator(self) -> None: - try: - async for chunk in self._generator: - data = chunk if "frames_b64" in chunk else chunk.get("data", {}) - frames_b64: list[str] = data.get("frames_b64", []) - for fb64 in frames_b64: - raw = base64.b64decode(fb64) - np_arr = np.frombuffer(raw, dtype=np.uint8) - bgr = cv2.imdecode(np_arr, cv2.IMREAD_COLOR) - if bgr is None: - continue - rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) - frame = av.VideoFrame.from_ndarray(rgb, format="rgb24") - self.push_frame(frame) - - if self._audio_track is not None: - audio_b64 = data.get("audio_b64") - if audio_b64: - self._audio_track.feed(base64.b64decode(audio_b64)) - except asyncio.CancelledError: - pass - except Exception as exc: - logger.error(f"FrameGeneratorTrack generator error: {exc}") - finally: - self._finished = True - if self._audio_track is not None: - self._audio_track.signal_done() - - async def recv(self) -> av.VideoFrame: - if self._task is None and self._generator is not None: - self._task = asyncio.create_task(self._consume_generator()) - - if self._start_time is None: - self._start_time = time.time() - - target_time = self._start_time + self._frame_count * self._frame_interval - wait = target_time - time.time() - if wait > 0: - await asyncio.sleep(wait) - - try: - frame = self._queue.get_nowait() - self._last_frame = frame - except asyncio.QueueEmpty: - if self._last_frame is not None: - # An idle control stream must keep displaying its last generated - # result. Ending the RTP track here makes browsers render black - # while the peer connection is still open. - frame = self._last_frame - elif self._finished: - raise MediaStreamError("Track ended — no frames received") - elif self._generator is None: - try: - frame = await asyncio.wait_for(self._queue.get(), timeout=0.25) - self._last_frame = frame - except asyncio.TimeoutError: - if self._finished: - raise MediaStreamError("Track ended — no frames received") - frame = self._make_placeholder_frame() - else: - try: - frame = await asyncio.wait_for(self._queue.get(), timeout=10.0) - self._last_frame = frame - except asyncio.TimeoutError: - raise MediaStreamError("Track ended — no frames received") - - frame.pts = self._frame_count * self._pts_per_frame - frame.time_base = fractions.Fraction(1, _RTP_CLOCK_RATE) - self._frame_count += 1 - return frame - - def stop(self) -> None: - if self._task is not None and not self._task.done(): - self._task.cancel() - if self._audio_track is not None: - self._audio_track.stop() - super().stop() - - -# --------------------------------------------------------------------------- -# Incoming media relays (client → server) -# --------------------------------------------------------------------------- - - -class IncomingVideoRelay: - """Consumes an incoming video track and forwards decoded frames to a callback. - - Frames are passed as native numpy arrays — no JPEG/base64 re-encoding. - """ - - def __init__( - self, - track: MediaStreamTrack, - session_id: str, - on_chunk: Callable[[str, dict], None], - ) -> None: - self._track = track - self._session_id = session_id - self._on_chunk = on_chunk - - async def run(self) -> None: - try: - while True: - frame: av.VideoFrame = await self._track.recv() - rgb = frame.to_ndarray(format="rgb24") - self._on_chunk(self._session_id, {"type": "media", "video_frames": [rgb]}) - except MediaStreamError: - logger.info(f"IncomingVideoRelay ended: session={self._session_id}") - except asyncio.CancelledError: - logger.info(f"IncomingVideoRelay cancelled: session={self._session_id}") - raise - except Exception as exc: - logger.error(f"IncomingVideoRelay error: session={self._session_id} {exc}") - - -class IncomingAudioRelay: - """Consumes an incoming audio track and forwards raw PCM bytes to a callback. - - Audio data is passed as raw bytes — no base64 re-encoding. - """ - - def __init__( - self, - track: MediaStreamTrack, - session_id: str, - on_chunk: Callable[[str, dict], None], - ) -> None: - self._track = track - self._session_id = session_id - self._on_chunk = on_chunk - - async def run(self) -> None: - try: - while True: - frame: av.AudioFrame = await self._track.recv() - pcm = frame.to_ndarray().flatten().astype(np.int16).tobytes() - self._on_chunk( - self._session_id, - { - "type": "media", - "audio_pcm": pcm, - "sample_rate": frame.sample_rate, - "channels": len(frame.layout.channels), - }, - ) - except MediaStreamError: - logger.info(f"IncomingAudioRelay ended: session={self._session_id}") - except asyncio.CancelledError: - logger.info(f"IncomingAudioRelay cancelled: session={self._session_id}") - raise - except Exception as exc: - logger.error(f"IncomingAudioRelay error: session={self._session_id} {exc}") diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py deleted file mode 100644 index cfe77be..0000000 --- a/tests/integration/conftest.py +++ /dev/null @@ -1,163 +0,0 @@ -"""Shared fixtures and helpers for stream integration tests.""" - -from __future__ import annotations - -import asyncio -from collections.abc import AsyncGenerator - -import pytest - -pytest.importorskip("fastapi") -pytest.importorskip("httpx") - -from fastapi.testclient import TestClient - -# --------------------------------------------------------------------------- -# Mock services -# --------------------------------------------------------------------------- - - -class MockServerPushService: - """Fake server-push service that yields chunks with optional audio.""" - - def __init__(self, num_chunks: int = 3, include_audio: bool = False): - self._num_chunks = num_chunks - self._include_audio = include_audio - - def start(self) -> None: - pass - - def stop(self) -> None: - pass - - async def serve(self, request: dict) -> AsyncGenerator[dict, None]: - import base64 - - fake_jpeg = base64.b64encode(b"\xff\xd8\xff\xe0fake-jpeg-data").decode() - num = request.get("num_chunks", self._num_chunks) - for i in range(num): - chunk: dict = { - "type": "chunk", - "index": i, - "frames_b64": [fake_jpeg], - "fps": 24, - "prompt": request.get("prompt", ""), - } - if self._include_audio: - import numpy as np - - silence = np.zeros(960, dtype=np.int16) - chunk["audio_b64"] = base64.b64encode(silence.tobytes()).decode() - chunk["audio_sample_rate"] = 48000 - chunk["audio_channels"] = 1 - yield chunk - - -class MockBidirectionalService: - """Fake bidirectional service with in-memory session tracking.""" - - def __init__(self): - self._sessions: dict[str, list[dict]] = {} - self._outputs: dict[str, asyncio.Queue] = {} - - def start(self) -> None: - pass - - def stop(self) -> None: - self._sessions.clear() - self._outputs.clear() - - def create_session(self, config: dict) -> str: - import uuid - - sid = config.get("session_id", str(uuid.uuid4())) - self._sessions[sid] = [] - self._outputs[sid] = asyncio.Queue() - return sid - - def push_chunk(self, session_id: str, chunk: dict) -> None: - if session_id not in self._sessions: - raise KeyError(f"Session {session_id} not found") - self._sessions[session_id].append(chunk) - self._outputs[session_id].put_nowait( - {"type": "chunk", "index": len(self._sessions[session_id]) - 1, "echo": chunk} - ) - - async def pull_chunks(self, session_id: str) -> AsyncGenerator[dict, None]: - if session_id not in self._outputs: - return - q = self._outputs[session_id] - while True: - try: - chunk = await asyncio.wait_for(q.get(), timeout=2.0) - if chunk.get("type") == "done": - break - yield chunk - except asyncio.TimeoutError: - break - - def close_session(self, session_id: str) -> None: - if session_id in self._outputs: - self._outputs[session_id].put_nowait({"type": "done"}) - self._sessions.pop(session_id, None) - - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - - -def make_stream_svc(service, stream_mode: str): - from telefuser.service.core.stream_pipeline_service import StreamPipelineService - - svc = StreamPipelineService.__new__(StreamPipelineService) - svc.is_running = True - svc.service = service - svc.stream_mode = stream_mode - svc.ppl_file = "mock.py" - svc._module = None - svc._module_name = None - svc._startup_measurement = None - svc._runtime_environment = {} - svc.security_level = None - svc.security_validator = None - return svc - - -def make_test_server(stream_svc): - from telefuser.service.api.api_server import ApiServer - from telefuser.service.core.task_manager import TaskManager - - task_manager = TaskManager(max_queue_size=10) - server = ApiServer(max_queue_size=10, task_manager=task_manager, enable_openai_api=False) - server.initialize_stream_service(stream_svc) - return server - - -# --------------------------------------------------------------------------- -# Fixtures -# --------------------------------------------------------------------------- - - -@pytest.fixture -def server_push_client(): - server = make_test_server(make_stream_svc(MockServerPushService(num_chunks=5), "server_push")) - with TestClient(server.app) as client: - yield client - asyncio.run(server.cleanup()) - - -@pytest.fixture -def bidirectional_client(): - server = make_test_server(make_stream_svc(MockBidirectionalService(), "bidirectional")) - with TestClient(server.app) as client: - yield client - asyncio.run(server.cleanup()) - - -@pytest.fixture -def audio_server_push_client(): - server = make_test_server(make_stream_svc(MockServerPushService(num_chunks=3, include_audio=True), "server_push")) - with TestClient(server.app) as client: - yield client - asyncio.run(server.cleanup()) diff --git a/tests/integration/test_stream_api.py b/tests/integration/test_stream_api.py deleted file mode 100644 index 2427ffa..0000000 --- a/tests/integration/test_stream_api.py +++ /dev/null @@ -1,188 +0,0 @@ -"""Integration tests for stream API endpoints.""" - -from __future__ import annotations - -import asyncio - -import pytest - -pytest.importorskip("fastapi") -pytest.importorskip("httpx") - - -def _make_offer() -> dict: - """Create a bidirectional SDP offer with a DataChannel.""" - pytest.importorskip("aiortc") - from aiortc import RTCPeerConnection - - async def _create(): - pc = RTCPeerConnection() - pc.createDataChannel("telefuser") - pc.addTransceiver("video", direction="recvonly") - offer = await pc.createOffer() - await pc.setLocalDescription(offer) - sdp = pc.localDescription.sdp - sdp_type = pc.localDescription.type - await pc.close() - return {"sdp": sdp, "type": sdp_type} - - return asyncio.run(_create()) - - -def _make_server_push_offer() -> dict: - """Create a server-push SDP offer (no DataChannel).""" - pytest.importorskip("aiortc") - from aiortc import RTCPeerConnection - - async def _create(): - pc = RTCPeerConnection() - pc.addTransceiver("video", direction="recvonly") - offer = await pc.createOffer() - await pc.setLocalDescription(offer) - sdp = pc.localDescription.sdp - sdp_type = pc.localDescription.type - await pc.close() - return {"sdp": sdp, "type": sdp_type} - - return asyncio.run(_create()) - - -# --------------------------------------------------------------------------- -# Bidirectional (session management) tests — via WebRTC offer -# --------------------------------------------------------------------------- - - -class TestBidirectionalSessions: - """Tests for bidirectional session management via WebRTC offer endpoint.""" - - @pytest.fixture(autouse=True) - def _skip_without_aiortc(self): - pytest.importorskip("aiortc") - - def test_create_bidirectional_session_via_offer(self, bidirectional_client): - offer = _make_offer() - body = {**offer, "task": "s2v", "config": {"fps": 24}} - resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - data = resp.json() - assert "session_id" in data - assert data["type"] == "answer" - - def test_close_bidirectional_session(self, bidirectional_client): - offer = _make_offer() - body = {**offer, "task": "s2v"} - create_resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - session_id = create_resp.json()["session_id"] - - resp = bidirectional_client.delete(f"/v1/stream/webrtc/{session_id}") - assert resp.status_code == 200 - assert resp.json()["status"] == "closed" - - def test_close_via_stream_sessions_alias(self, bidirectional_client): - """DELETE /v1/stream/sessions/{id} should close both pipeline and WebRTC sessions.""" - offer = _make_offer() - body = {**offer, "task": "s2v"} - create_resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - session_id = create_resp.json()["session_id"] - - resp = bidirectional_client.delete(f"/v1/stream/sessions/{session_id}") - assert resp.status_code == 200 - assert resp.json()["status"] == "closed" - - status_resp = bidirectional_client.get(f"/v1/stream/sessions/{session_id}/status") - assert status_resp.json()["status"] == "unknown" - - -# --------------------------------------------------------------------------- -# Service status / metadata tests -# --------------------------------------------------------------------------- - - -class TestStreamServiceStatus: - """Tests for service health / metadata with stream endpoints.""" - - def test_service_status_with_stream(self, server_push_client): - resp = server_push_client.get("/v1/service/status") - assert resp.status_code == 200 - data = resp.json() - assert "service_status" in data - - def test_health_with_stream(self, server_push_client): - resp = server_push_client.get("/v1/service/health") - assert resp.status_code == 200 - assert resp.json()["status"] == "healthy" - - def test_health_reports_stream_readiness(self, server_push_client): - resp = server_push_client.get("/v1/service/health") - data = resp.json() - assert data["stream_ready"] is True - assert data["stream_mode"] == "server_push" - - def test_metadata_returns_stream_info_without_inference_service(self, server_push_client): - resp = server_push_client.get("/v1/service/metadata") - assert resp.status_code == 200 - data = resp.json() - assert data["service_type"] == "stream" - assert data["stream_mode"] == "server_push" - assert data["runner"] == "StreamPipelineService" - - def test_metrics_json_includes_webrtc_stats(self, server_push_client): - resp = server_push_client.get("/v1/service/metrics/json") - assert resp.status_code == 200 - data = resp.json() - assert "webrtc" in data - - def test_close_server_push_via_stream_sessions_alias(self, server_push_client): - """DELETE /v1/stream/sessions/{id} should close server-push WebRTC sessions.""" - pytest.importorskip("aiortc") - offer = _make_server_push_offer() - body = {**offer, "task": "t2v", "prompt": "test"} - create_resp = server_push_client.post("/v1/stream/webrtc/offer", json=body) - assert create_resp.status_code == 200 - session_id = create_resp.json()["session_id"] - - resp = server_push_client.delete(f"/v1/stream/sessions/{session_id}") - assert resp.status_code == 200 - assert resp.json()["status"] == "closed" - - status_resp = server_push_client.get(f"/v1/stream/sessions/{session_id}/status") - assert status_resp.json()["status"] == "unknown" - - -# --------------------------------------------------------------------------- -# Session status tests -# --------------------------------------------------------------------------- - - -class TestSessionStatusLookup: - """Session status should query both stream service and WebRTC session manager.""" - - @pytest.fixture(autouse=True) - def _skip_without_aiortc(self): - pytest.importorskip("aiortc") - - def test_active_session_returns_active(self, bidirectional_client): - offer = _make_offer() - body = {**offer, "task": "s2v"} - create_resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - session_id = create_resp.json()["session_id"] - - resp = bidirectional_client.get(f"/v1/stream/sessions/{session_id}/status") - assert resp.status_code == 200 - data = resp.json() - assert data["status"] == "active" - - def test_unknown_session_returns_unknown(self, bidirectional_client): - resp = bidirectional_client.get("/v1/stream/sessions/nonexistent/status") - assert resp.status_code == 200 - assert resp.json()["status"] == "unknown" - - def test_closed_session_returns_unknown(self, bidirectional_client): - offer = _make_offer() - body = {**offer, "task": "s2v"} - create_resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - session_id = create_resp.json()["session_id"] - bidirectional_client.delete(f"/v1/stream/webrtc/{session_id}") - - resp = bidirectional_client.get(f"/v1/stream/sessions/{session_id}/status") - assert resp.json()["status"] == "unknown" diff --git a/tests/integration/test_webrtc_api.py b/tests/integration/test_webrtc_api.py deleted file mode 100644 index 111a340..0000000 --- a/tests/integration/test_webrtc_api.py +++ /dev/null @@ -1,180 +0,0 @@ -"""Integration tests for WebRTC signaling endpoints.""" - -from __future__ import annotations - -import asyncio - -import pytest - -pytest.importorskip("fastapi") -pytest.importorskip("httpx") -aiortc = pytest.importorskip("aiortc") - -from fastapi.testclient import TestClient - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - - -def _make_sdp_offer(include_audio: bool = True) -> dict: - """Create a minimal SDP offer using aiortc.""" - from aiortc import RTCPeerConnection - - async def _create(): - pc = RTCPeerConnection() - pc.addTransceiver("video", direction="recvonly") - if include_audio: - pc.addTransceiver("audio", direction="recvonly") - offer = await pc.createOffer() - await pc.setLocalDescription(offer) - sdp = pc.localDescription.sdp - sdp_type = pc.localDescription.type - await pc.close() - return {"sdp": sdp, "type": sdp_type} - - return asyncio.run(_create()) - - -def _make_bidirectional_sdp_offer(include_audio: bool = True) -> dict: - """Create an SDP offer with a DataChannel (for bidirectional mode).""" - from aiortc import RTCPeerConnection - - async def _create(): - pc = RTCPeerConnection() - pc.createDataChannel("telefuser") - pc.addTransceiver("video", direction="recvonly") - if include_audio: - pc.addTransceiver("audio", direction="recvonly") - offer = await pc.createOffer() - await pc.setLocalDescription(offer) - sdp = pc.localDescription.sdp - sdp_type = pc.localDescription.type - await pc.close() - return {"sdp": sdp, "type": sdp_type} - - return asyncio.run(_create()) - - -# --------------------------------------------------------------------------- -# Server-push tests (unchanged) -# --------------------------------------------------------------------------- - - -class TestWebRTCOffer: - """Tests for POST /v1/stream/webrtc/offer (server-push mode).""" - - def test_offer_returns_sdp_answer(self, server_push_client): - offer = _make_sdp_offer() - body = {**offer, "task": "t2v", "prompt": "a sunset"} - resp = server_push_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - data = resp.json() - assert "sdp" in data - assert data["type"] == "answer" - assert "session_id" in data - - def test_offer_rejects_when_service_not_running(self): - from telefuser.service.api.api_server import ApiServer - from telefuser.service.core.task_manager import TaskManager - - server = ApiServer(max_queue_size=10, task_manager=TaskManager(), enable_openai_api=False) - with TestClient(server.app) as client: - offer = _make_sdp_offer() - body = {**offer, "task": "t2v", "prompt": "test"} - resp = client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 503 - - -class TestWebRTCSession: - """Tests for DELETE /v1/stream/webrtc/{session_id}.""" - - def test_close_existing_session(self, server_push_client): - offer = _make_sdp_offer() - body = {**offer, "task": "t2v", "prompt": "test"} - create_resp = server_push_client.post("/v1/stream/webrtc/offer", json=body) - session_id = create_resp.json()["session_id"] - - resp = server_push_client.delete(f"/v1/stream/webrtc/{session_id}") - assert resp.status_code == 200 - assert resp.json()["status"] == "closed" - - def test_close_nonexistent_session(self, server_push_client): - resp = server_push_client.delete("/v1/stream/webrtc/nonexistent") - assert resp.status_code == 404 - - -class TestWebRTCAudio: - """Tests for WebRTC audio track support.""" - - def test_offer_with_audio_and_audio_chunks(self, audio_server_push_client): - offer = _make_sdp_offer(include_audio=True) - body = {**offer, "task": "t2v", "prompt": "audio test"} - resp = audio_server_push_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - data = resp.json() - assert "sdp" in data - assert "m=audio" in data["sdp"] - - def test_offer_with_audio_but_no_audio_chunks(self, server_push_client): - """Backwards compat: client offers audio but pipeline has no audio data.""" - offer = _make_sdp_offer(include_audio=True) - body = {**offer, "task": "t2v", "prompt": "no audio"} - resp = server_push_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - - def test_offer_without_audio_transceiver(self, audio_server_push_client): - """Video-only client: no audio transceiver in SDP offer.""" - offer = _make_sdp_offer(include_audio=False) - body = {**offer, "task": "t2v", "prompt": "video only"} - resp = audio_server_push_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - data = resp.json() - assert "m=audio" not in data["sdp"] - - -# --------------------------------------------------------------------------- -# Bidirectional WebRTC tests -# --------------------------------------------------------------------------- - - -class TestWebRTCBidirectional: - """Tests for WebRTC bidirectional mode via POST /v1/stream/webrtc/offer.""" - - def test_offer_creates_bidirectional_session(self, bidirectional_client): - offer = _make_bidirectional_sdp_offer() - body = {**offer, "task": "s2v", "prompt": "test", "config": {"fps": 24}} - resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - data = resp.json() - assert "sdp" in data - assert data["type"] == "answer" - assert "session_id" in data - - def test_offer_bidirectional_without_audio(self, bidirectional_client): - offer = _make_bidirectional_sdp_offer(include_audio=False) - body = {**offer, "task": "s2v", "prompt": "no audio"} - resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 200 - data = resp.json() - assert "m=audio" not in data["sdp"] - - def test_close_bidirectional_session(self, bidirectional_client): - offer = _make_bidirectional_sdp_offer() - body = {**offer, "task": "s2v", "prompt": "test"} - create_resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - session_id = create_resp.json()["session_id"] - - resp = bidirectional_client.delete(f"/v1/stream/webrtc/{session_id}") - assert resp.status_code == 200 - assert resp.json()["status"] == "closed" - - def test_close_nonexistent_bidirectional(self, bidirectional_client): - resp = bidirectional_client.delete("/v1/stream/webrtc/nonexistent") - assert resp.status_code == 404 - - def test_sdp_failure_rolls_back_pipeline_session(self, bidirectional_client): - """Offer with invalid SDP should not leave orphaned pipeline sessions.""" - body = {"sdp": "invalid-sdp", "type": "offer", "task": "s2v", "prompt": "test"} - resp = bidirectional_client.post("/v1/stream/webrtc/offer", json=body) - assert resp.status_code == 400 diff --git a/tests/unit/service/test_service_routes.py b/tests/unit/service/test_service_routes.py index a854bc3..4188867 100644 --- a/tests/unit/service/test_service_routes.py +++ b/tests/unit/service/test_service_routes.py @@ -1,7 +1,6 @@ from __future__ import annotations import asyncio -import sys import threading import types from pathlib import Path @@ -9,16 +8,11 @@ import pytest from click.testing import CliRunner -from fastapi import HTTPException from telefuser.entrypoints.cli.main import main from telefuser.service.api.api_server import ApiServer from telefuser.service.api.routers.service import ServiceRoutes -from telefuser.service.api.routers.stream import StreamRoutes -from telefuser.service.api.routers.webrtc import WebRTCRoutes -from telefuser.service.api.stream_schema import WebRTCOfferRequest from telefuser.service.core.config import ServerConfig -from telefuser.service.core.container import ServiceContainer from telefuser.service.core.file_service import FileService from telefuser.service.core.pipeline_service import PipelineService from telefuser.service.core.stream_pipeline_service import StreamPipelineService @@ -28,11 +22,6 @@ from telefuser.service_types import MediaType, TaskStatus -def _openapi_paths(app: object) -> set[str]: - openapi = app.openapi() - return set(openapi.get("paths", {})) - - def test_task_status_uses_shared_service_enum() -> None: assert CoreTaskStatus is TaskStatus assert TaskStatus.STREAMING.value == "streaming" @@ -131,364 +120,6 @@ def pool_status(self) -> list[dict]: assert ready.status_code == 200 -def test_stream_route_profile_exposes_only_stream_service_routes() -> None: - server = ApiServer(task_manager=TaskManager(), enable_openai_api=True, route_profile="stream") - paths = _openapi_paths(server.app) - - assert "/v1/service/health" in paths - assert "/v1/stream/sessions/{session_id}/status" in paths - assert "/v1/tasks/create" not in paths - assert "/v1/tasks/form" not in paths - assert "/v1/files/download/{file_path}" not in paths - assert "/v1/images/generations" not in paths - assert "/v1/videos" not in paths - - -def test_request_response_route_profile_excludes_stream_routes() -> None: - server = ApiServer(task_manager=TaskManager(), enable_openai_api=False, route_profile="request_response") - paths = _openapi_paths(server.app) - - assert "/v1/service/health" in paths - assert "/v1/tasks/create" in paths - assert "/v1/files/download/{file_path}" in paths - assert "/v1/stream/sessions/{session_id}/status" not in paths - - -def test_container_stream_app_uses_stream_route_profile() -> None: - container = ServiceContainer.create(config=ServerConfig()) - container.stream_pipeline_service = Mock() - container.stream_pipeline_service.is_running = True - - app = container.get_api_app() - paths = _openapi_paths(app) - - assert "/v1/service/health" in paths - assert "/v1/stream/sessions/{session_id}/status" in paths - assert "/v1/tasks/create" not in paths - assert "/v1/files/download/{file_path}" not in paths - - -def test_stream_session_close_logs_pipeline_close_failure(monkeypatch: pytest.MonkeyPatch) -> None: - class RunningStreamService: - is_running = True - - def close_session(self, session_id: str) -> None: - raise RuntimeError("pipeline close failed") - - class WebRTCSessionManager: - async def close_session(self, session_id: str, **kwargs: object) -> bool: - return True - - server = ApiServer(task_manager=TaskManager(), enable_openai_api=False) - server.stream_service = RunningStreamService() - server._webrtc_routes = Mock(_session_manager=WebRTCSessionManager()) - warnings: list[str] = [] - monkeypatch.setattr("telefuser.service.api.routers.stream.logger.warning", warnings.append) - - result = asyncio.run(StreamRoutes(server).close_session("session-123")) - - assert result == {"session_id": "session-123", "status": "closed"} - assert warnings == ["Failed to close pipeline stream session session-123: pipeline close failed"] - - -def test_stream_session_alias_closes_webrtc_without_pipeline_notify() -> None: - class RunningStreamService: - is_running = True - - def __init__(self) -> None: - self.closed_sessions: list[str] = [] - - def close_session(self, session_id: str) -> None: - self.closed_sessions.append(session_id) - - class WebRTCSessionManager: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - - async def close_session(self, session_id: str, *, reason: str, notify_pipeline: bool) -> bool: - self.calls.append( - { - "session_id": session_id, - "reason": reason, - "notify_pipeline": notify_pipeline, - } - ) - return True - - stream_service = RunningStreamService() - webrtc_manager = WebRTCSessionManager() - server = ApiServer(task_manager=TaskManager(), enable_openai_api=False) - server.stream_service = stream_service - server._webrtc_routes = Mock(_session_manager=webrtc_manager) - - result = asyncio.run(StreamRoutes(server).close_session("session-123")) - - assert result == {"session_id": "session-123", "status": "closed"} - assert stream_service.closed_sessions == ["session-123"] - assert webrtc_manager.calls == [ - { - "session_id": "session-123", - "reason": "stream_session_delete", - "notify_pipeline": False, - } - ] - - -def test_webrtc_delete_uses_session_manager_pipeline_owner() -> None: - class RunningStreamService: - stream_mode = "bidirectional" - - def __init__(self) -> None: - self.closed_sessions: list[str] = [] - - def close_session(self, session_id: str) -> None: - self.closed_sessions.append(session_id) - - class WebRTCSessionManager: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - - async def close_session(self, session_id: str, *, reason: str) -> bool: - self.calls.append({"session_id": session_id, "reason": reason}) - return True - - stream_service = RunningStreamService() - webrtc_manager = WebRTCSessionManager() - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=stream_service) - routes._session_manager = webrtc_manager - - result = asyncio.run(routes.close_session("session-123")) - - assert result == {"session_id": "session-123", "status": "closed"} - assert stream_service.closed_sessions == [] - assert webrtc_manager.calls == [{"session_id": "session-123", "reason": "webrtc_session_delete"}] - - -def test_bidirectional_webrtc_offer_flattens_session_config() -> None: - class RunningStreamService: - def __init__(self) -> None: - self.config: dict | None = None - - def create_session(self, config: dict) -> str: - self.config = config - return "pipeline-session" - - async def pull_chunks(self, session_id: str): - if False: - yield session_id - - def push_chunk(self, session_id: str, chunk: dict) -> None: - pass - - def close_session(self, session_id: str) -> None: - pass - - class WebRTCSessionManager: - def __init__(self) -> None: - self.fps: int | None = None - - async def create_bidirectional_session(self, **kwargs: object) -> tuple[str, str]: - self.fps = int(kwargs["fps"]) - return "answer-sdp", "answer" - - stream_service = RunningStreamService() - webrtc_manager = WebRTCSessionManager() - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=stream_service) - routes._session_manager = webrtc_manager - request = WebRTCOfferRequest( - session_id="request-session", - sdp="offer-sdp", - task="i2v", - prompt="test prompt", - image="data:image/png;base64,test-image", - config={"image_path": "/tmp/input.png", "fps": 16, "frame_num": 81}, - ) - - response = asyncio.run(routes._handle_bidirectional(stream_service, request)) - - assert response.session_id == "pipeline-session" - assert stream_service.config == { - "session_id": "request-session", - "task": "i2v", - "prompt": "test prompt", - "fps": 16, - "image": "data:image/png;base64,test-image", - "image_path": "/tmp/input.png", - "frame_num": 81, - } - assert webrtc_manager.fps == 16 - - -def test_bidirectional_webrtc_offer_uses_service_default_fps_when_omitted() -> None: - class RunningStreamService: - def __init__(self) -> None: - self.config: dict | None = None - self.service = types.SimpleNamespace(default_fps=16) - - def create_session(self, config: dict) -> str: - self.config = config - return "pipeline-session" - - async def pull_chunks(self, session_id: str): - if False: - yield session_id - - def push_chunk(self, session_id: str, chunk: dict) -> None: - pass - - def close_session(self, session_id: str) -> None: - pass - - class WebRTCSessionManager: - def __init__(self) -> None: - self.fps: int | None = None - - async def create_bidirectional_session(self, **kwargs: object) -> tuple[str, str]: - self.fps = int(kwargs["fps"]) - return "answer-sdp", "answer" - - stream_service = RunningStreamService() - webrtc_manager = WebRTCSessionManager() - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=stream_service) - routes._session_manager = webrtc_manager - request = WebRTCOfferRequest(session_id="request-session", sdp="offer-sdp", task="i2v") - - asyncio.run(routes._handle_bidirectional(stream_service, request)) - - assert stream_service.config == {"session_id": "request-session", "task": "i2v"} - assert webrtc_manager.fps == 16 - - -def test_server_push_webrtc_offer_uses_service_default_fps_when_omitted() -> None: - class RunningStreamService: - service = types.SimpleNamespace(default_fps=16) - - async def stream_task(self, task_data: dict): - if False: - yield task_data - - class WebRTCSessionManager: - def __init__(self) -> None: - self.fps: int | None = None - - async def create_session(self, **kwargs: object) -> tuple[str, str]: - self.fps = int(kwargs["fps"]) - return "answer-sdp", "answer" - - stream_service = RunningStreamService() - webrtc_manager = WebRTCSessionManager() - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=stream_service) - routes._session_manager = webrtc_manager - request = WebRTCOfferRequest(session_id="request-session", sdp="offer-sdp", task="i2v") - - asyncio.run(routes._handle_server_push(stream_service, request)) - - assert webrtc_manager.fps == 16 - - -def test_bidirectional_webrtc_offer_prefers_explicit_top_level_config() -> None: - class RunningStreamService: - def __init__(self) -> None: - self.config: dict | None = None - - def create_session(self, config: dict) -> str: - self.config = config - return "pipeline-session" - - async def pull_chunks(self, session_id: str): - if False: - yield session_id - - class WebRTCSessionManager: - async def create_bidirectional_session(self, **kwargs: object) -> tuple[str, str]: - return "answer-sdp", "answer" - - stream_service = RunningStreamService() - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=stream_service) - routes._session_manager = WebRTCSessionManager() - request = WebRTCOfferRequest( - session_id="request-session", - sdp="offer-sdp", - task="i2v", - fps=30, - config={"task": "t2v", "fps": 16}, - ) - - asyncio.run(routes._handle_bidirectional(stream_service, request)) - - assert stream_service.config is not None - assert stream_service.config["task"] == "i2v" - assert stream_service.config["fps"] == 30 - - -@pytest.mark.parametrize( - ("error", "expected_status", "expected_detail"), - [ - (ValueError("image_path is required"), 400, "Session creation failed: image_path is required"), - (RuntimeError("active session exists"), 409, "Session creation conflict: active session exists"), - ], -) -def test_bidirectional_webrtc_offer_maps_session_creation_errors( - error: Exception, - expected_status: int, - expected_detail: str, -) -> None: - class FailingStreamService: - def create_session(self, config: dict) -> str: - raise error - - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=FailingStreamService()) - routes._session_manager = Mock() - request = WebRTCOfferRequest(session_id="request-session", sdp="offer-sdp", task="i2v") - - with pytest.raises(HTTPException) as exc_info: - asyncio.run(routes._handle_bidirectional(routes.api.stream_service, request)) - - assert exc_info.value.status_code == expected_status - assert exc_info.value.detail == expected_detail - - -def test_bidirectional_webrtc_offer_runs_rollback_outside_event_loop() -> None: - callback_threads: list[int] = [] - - class RunningStreamService: - def create_session(self, config: dict) -> str: - return "pipeline-session" - - async def pull_chunks(self, session_id: str): - if False: - yield session_id - - def close_session(self, session_id: str) -> None: - callback_threads.append(threading.get_ident()) - - class FailingWebRTCSessionManager: - async def create_bidirectional_session(self, **kwargs: object) -> tuple[str, str]: - raise ValueError("invalid SDP") - - async def run_offer() -> int: - stream_service = RunningStreamService() - routes = WebRTCRoutes.__new__(WebRTCRoutes) - routes.api = Mock(stream_service=stream_service) - routes._session_manager = FailingWebRTCSessionManager() - request = WebRTCOfferRequest(session_id="request-session", sdp="invalid-sdp", task="i2v") - event_loop_thread = threading.get_ident() - with pytest.raises(HTTPException, match="SDP negotiation failed"): - await routes._handle_bidirectional(stream_service, request) - return event_loop_thread - - event_loop_thread = asyncio.run(run_offer()) - - assert len(callback_threads) == 1 - assert callback_threads[0] != event_loop_thread - - def test_cli_serve_forwards_security_and_skip_validation(monkeypatch: pytest.MonkeyPatch) -> None: captured = {} @@ -513,40 +144,6 @@ def fake_run_server(**kwargs): assert captured["skip_validation"] is True -def test_cli_stream_serve_does_not_force_skip_validation(monkeypatch: pytest.MonkeyPatch) -> None: - captured = {} - - def fake_assert_safe(self, pipe_path): - captured["validated"] = pipe_path - - def fake_run_stream_server(pipe_path, port, host, **kwargs): - captured.update({"pipe_path": pipe_path, "port": port, "host": host, **kwargs}) - - monkeypatch.setattr( - "telefuser.service.security.security_validator.PipelineSecurityValidator.assert_safe", - fake_assert_safe, - ) - monkeypatch.setattr("telefuser.service.main.run_stream_server", fake_run_stream_server) - - result = CliRunner().invoke( - main, - [ - "stream-serve", - "stream_pipeline.py", - "--gpu-num", - "3", - "--security-level", - "none", - ], - ) - - assert result.exit_code == 0 - assert captured["validated"] == "stream_pipeline.py" - assert captured["gpu_num"] == 3 - assert captured["security_level"] == "none" - assert captured["skip_validation"] is False - - def test_run_server_security_level_is_applied(monkeypatch: pytest.MonkeyPatch) -> None: from telefuser.service import main as service_main @@ -683,76 +280,3 @@ def get_service(gpu_num: int) -> FakeBidirectionalService: assert service.start_service("stream_pipeline.py", gpu_num=3, skip_validation=True) assert captured == {"gpu_num": 3, "started": True} - - -def test_webrtc_routes_use_api_server_config(monkeypatch: pytest.MonkeyPatch) -> None: - captured = {} - - class FakeRTCIceServer: - def __init__(self, urls: str, username: str | None = None, credential: str | None = None) -> None: - self.urls = urls - self.username = username - self.credential = credential - - class FakeRTCConfiguration: - def __init__(self, iceServers: list[FakeRTCIceServer]) -> None: - self.iceServers = iceServers - - class FakeWebRTCSessionManager: - def __init__( - self, - max_sessions: int, - configuration: FakeRTCConfiguration, - video_codec: str, - video_bitrate: int, - video_buffer_seconds: float, - data_channel_timeout_seconds: float, - disconnected_grace_seconds: float, - ) -> None: - captured["max_sessions"] = max_sessions - captured["configuration"] = configuration - captured["video_codec"] = video_codec - captured["video_bitrate"] = video_bitrate - captured["video_buffer_seconds"] = video_buffer_seconds - captured["data_channel_timeout_seconds"] = data_channel_timeout_seconds - captured["disconnected_grace_seconds"] = disconnected_grace_seconds - - monkeypatch.setitem( - sys.modules, - "aiortc", - types.SimpleNamespace(RTCConfiguration=FakeRTCConfiguration, RTCIceServer=FakeRTCIceServer), - ) - monkeypatch.setitem( - sys.modules, - "telefuser.service.webrtc.session_manager", - types.SimpleNamespace(WebRTCSessionManager=FakeWebRTCSessionManager), - ) - - config = ServerConfig( - webrtc_max_sessions=3, - stun_servers=["stun:local:3478"], - turn_server="turn:local:3478", - turn_username="user", - turn_credential="secret", - ) - server = ApiServer( - task_manager=TaskManager(), - enable_openai_api=False, - config=config, - route_profile="request_response", - ) - - from telefuser.service.api.routers.webrtc import WebRTCRoutes - - WebRTCRoutes(server) - - assert captured["max_sessions"] == 3 - assert captured["video_codec"] == "H264" - assert captured["video_bitrate"] == 8_000_000 - assert captured["video_buffer_seconds"] == 1.0 - assert captured["data_channel_timeout_seconds"] == 10.0 - assert captured["disconnected_grace_seconds"] == 5.0 - ice_servers = captured["configuration"].iceServers - assert [server.urls for server in ice_servers] == ["stun:local:3478", "turn:local:3478"] - assert ice_servers[1].username == "user" - assert ice_servers[1].credential == "secret" diff --git a/tests/unit/service/test_webrtc_session_manager.py b/tests/unit/service/test_webrtc_session_manager.py deleted file mode 100644 index 3ef714a..0000000 --- a/tests/unit/service/test_webrtc_session_manager.py +++ /dev/null @@ -1,208 +0,0 @@ -from __future__ import annotations - -import asyncio -import json -import threading -from unittest.mock import AsyncMock, MagicMock, patch - -import av -import numpy as np -import pytest -from PIL import Image - -pytest.importorskip("aiortc") - -from aiortc.codecs import h264, vpx - -from telefuser.service.webrtc.chunk_router import ChunkRouter -from telefuser.service.webrtc.session_manager import WebRTCSessionManager, _BidirectionalSession -from telefuser.service.webrtc.track import FrameGeneratorTrack - - -def test_video_quality_configuration_prefers_h264_and_raises_bitrate() -> None: - with ( - patch.object(h264, "DEFAULT_BITRATE", 1_000_000), - patch.object(h264, "MAX_BITRATE", 3_000_000), - patch.object(vpx, "DEFAULT_BITRATE", 500_000), - patch.object(vpx, "MAX_BITRATE", 1_500_000), - ): - manager = WebRTCSessionManager(video_codec="H264", video_bitrate=8_000_000) - - assert h264.DEFAULT_BITRATE == h264.MAX_BITRATE == 8_000_000 - assert vpx.DEFAULT_BITRATE == vpx.MAX_BITRATE == 8_000_000 - - transceiver = MagicMock(kind="video") - peer_connection = MagicMock() - peer_connection.getTransceivers.return_value = [transceiver] - manager._set_video_codec_preferences(peer_connection) - - codecs = transceiver.setCodecPreferences.call_args.args[0] - assert codecs[0].mimeType == "video/H264" - - -def test_chunk_router_prefers_raw_frames_over_jpeg_transport() -> None: - video_track = MagicMock() - image = Image.fromarray(np.full((8, 12, 3), [17, 113, 241], dtype=np.uint8)) - router = ChunkRouter( - generator=MagicMock(), - video_track=video_track, - audio_track=None, - data_channel_send=None, - session_id="quality-test", - ) - - router._route_chunk({"frames": [image], "frames_b64": ["not-used"]}) - - frame = video_track.push_frame.call_args.args[0] - np.testing.assert_array_equal(frame.to_ndarray(format="rgb24"), np.asarray(image)) - - -def test_frame_track_discards_oldest_frames_when_client_falls_behind() -> None: - track = FrameGeneratorTrack(fps=2, max_buffer_seconds=1.0) - frames = [ - av.VideoFrame.from_ndarray(np.full((2, 2, 3), color, dtype=np.uint8), format="rgb24") for color in (10, 20, 30) - ] - - track.push_frames(frames[:2]) - track.push_frame(frames[2]) - - queued = [track._queue.get_nowait(), track._queue.get_nowait()] - assert [frame.to_ndarray(format="rgb24")[0, 0, 0] for frame in queued] == [20, 30] - assert track.dropped_frames == 1 - - -def test_frame_track_holds_last_frame_after_output_finishes() -> None: - async def receive_frames() -> tuple[int, int]: - track = FrameGeneratorTrack(fps=120) - source = av.VideoFrame.from_ndarray(np.full((2, 2, 3), 77, dtype=np.uint8), format="rgb24") - track.push_frame(source) - - first = await track.recv() - track.signal_done() - held = await track.recv() - return first.to_ndarray(format="rgb24")[0, 0, 0], held.to_ndarray(format="rgb24")[0, 0, 0] - - assert asyncio.run(receive_frames()) == (77, 77) - - -def test_chunk_router_counts_only_video_chunks() -> None: - async def output(): - yield {"type": "status", "stage": "ready"} - yield {"type": "chunk", "frames": [Image.new("RGB", (2, 2))]} - - data_channel_send = MagicMock() - router = ChunkRouter( - generator=output(), - video_track=MagicMock(), - audio_track=None, - data_channel_send=data_channel_send, - session_id="count-test", - ) - - asyncio.run(router.run()) - - done = json.loads(data_channel_send.call_args.args[0]) - assert done["type"] == "done" - assert done["total_chunks"] == 1 - - -def test_output_completion_schedules_transport_cleanup() -> None: - manager = WebRTCSessionManager(terminal_grace_seconds=0, close_on_output_complete=True) - close_session = MagicMock() - - async def close(*args, **kwargs): - close_session(*args, **kwargs) - return True - - manager.close_session = close - - async def run() -> None: - manager._schedule_output_complete("completed-session") - await asyncio.sleep(0) - - asyncio.run(run()) - - close_session.assert_called_once_with("completed-session", reason="output_complete") - - -def test_output_completion_keeps_session_open_by_default() -> None: - manager = WebRTCSessionManager(terminal_grace_seconds=0) - close_session = MagicMock() - - async def run() -> None: - manager.close_session = close_session - manager._schedule_output_complete("completed-session") - await asyncio.sleep(0) - - asyncio.run(run()) - - close_session.assert_not_called() - - -def test_missing_data_channel_closes_bidirectional_session_after_timeout() -> None: - manager = WebRTCSessionManager(data_channel_timeout_seconds=0.001) - manager._sessions["missing-channel"] = _BidirectionalSession( - pc=MagicMock(), - session_id="missing-channel", - ) - manager.close_session = AsyncMock(return_value=True) - - asyncio.run(manager._close_if_data_channel_missing("missing-channel")) - - manager.close_session.assert_awaited_once_with("missing-channel", reason="data_channel_timeout") - - -def test_disconnect_grace_does_not_close_a_recovered_session() -> None: - manager = WebRTCSessionManager(disconnected_grace_seconds=0) - peer_connection = MagicMock() - peer_connection.connectionState = "connected" - manager._sessions["recovered"] = _BidirectionalSession(pc=peer_connection, session_id="recovered") - manager.close_session = AsyncMock(return_value=True) - - asyncio.run(manager._close_after_disconnect_grace("recovered")) - - manager.close_session.assert_not_awaited() - - -def test_chunk_router_preserves_numeric_frame_count_as_metadata() -> None: - data_channel_send = MagicMock() - router = ChunkRouter( - generator=MagicMock(), - video_track=MagicMock(), - audio_track=None, - data_channel_send=data_channel_send, - session_id="status-test", - ) - - router._route_chunk({"type": "status", "stage": "chunk_sent", "frames": 13}) - - message = json.loads(data_channel_send.call_args.args[0]) - assert message["data"]["frames"] == 13 - - -def test_close_session_runs_pipeline_callback_outside_event_loop() -> None: - callback_threads: list[int] = [] - - class PeerConnection: - async def close(self) -> None: - pass - - def on_close(session_id: str) -> None: - callback_threads.append(threading.get_ident()) - - async def close_session() -> tuple[bool, int]: - manager = WebRTCSessionManager() - manager._sessions["session-123"] = _BidirectionalSession( - pc=PeerConnection(), - on_close=on_close, - session_id="session-123", - ) - event_loop_thread = threading.get_ident() - closed = await manager.close_session("session-123") - return closed, event_loop_thread - - closed, event_loop_thread = asyncio.run(close_session()) - - assert closed is True - assert len(callback_threads) == 1 - assert callback_threads[0] != event_loop_thread diff --git a/webui/stream_app.py b/webui/stream_app.py deleted file mode 100644 index dbca1a8..0000000 --- a/webui/stream_app.py +++ /dev/null @@ -1,191 +0,0 @@ -"""Gradio web UI for TeleFuser stream video generation. - -Connects to the stream server via WebRTC for real-time playback. -Requires ``pip install telefuser[webrtc]`` on the server side. - -Usage: - # 1. Start the stream server - telefuser stream-serve examples/stream_video_replay.py -p 8088 - - # 2. Launch the Gradio UI - python webui/stream_app.py --server-url http://localhost:8088 -""" - -from __future__ import annotations - -import argparse - -import gradio as gr -import requests - -DEFAULT_SERVER_URL = "http://localhost:8088" -DURATION_S = 30 - - -# --------------------------------------------------------------------------- -# WebRTC transport -# --------------------------------------------------------------------------- - -_WEBRTC_BTN_JS = """ -(prompt, server_url, duration) => { - // Clean up previous connection - if (window._telefuserPC) { - window._telefuserPC.close(); - window._telefuserPC = null; - } - - const container = document.querySelector('#webrtc-container'); - if (!container) return [prompt, server_url, duration]; - - container.innerHTML = ` -
- -
- Connecting... -
-
- `; - - const video = document.getElementById("webrtc-video"); - const overlay = document.getElementById("webrtc-overlay"); - - const unmuteBtn = document.createElement("button"); - unmuteBtn.textContent = "Unmute"; - unmuteBtn.style.cssText = "margin-top:8px;padding:6px 16px;border:none;" + - "border-radius:4px;background:#16a34a;color:#fff;cursor:pointer;font-size:14px;"; - unmuteBtn.onclick = () => { - video.muted = !video.muted; - unmuteBtn.textContent = video.muted ? "Unmute" : "Mute"; - }; - container.appendChild(unmuteBtn); - - const pc = new RTCPeerConnection(); - window._telefuserPC = pc; - - pc.addTransceiver("video", { direction: "recvonly" }); - pc.addTransceiver("audio", { direction: "recvonly" }); - pc.ontrack = (evt) => { - if (evt.track.kind === "video") { - video.srcObject = evt.streams[0]; - overlay.textContent = ""; - } - }; - pc.onconnectionstatechange = () => { - if (pc.connectionState === "failed") overlay.textContent = "Connection failed"; - }; - - (async () => { - try { - const offer = await pc.createOffer(); - await pc.setLocalDescription(offer); - const r = await fetch(server_url.replace(/\\/$/, "") + "/v1/stream/webrtc/offer", { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ - sdp: pc.localDescription.sdp, - type: pc.localDescription.type, - task: "t2v", - prompt: prompt, - duration_s: duration, - fps: 24, - }), - }); - if (!r.ok) { - const e = await r.json(); - overlay.textContent = "Error: " + (e.detail || r.statusText); - return; - } - const ans = await r.json(); - window._telefuserSessionId = ans.session_id; - await pc.setRemoteDescription(new RTCSessionDescription({ sdp: ans.sdp, type: ans.type })); - } catch(e) { - overlay.textContent = "Error: " + e.message; - } - })(); - - return [prompt, server_url, duration]; -} -""" - - -def _generate_webrtc(prompt: str, server_url: str, duration: float): - """WebRTC mode: JS handles the connection, Python just updates status.""" - if not prompt.strip(): - yield gr.skip(), "Please enter a prompt." - return - - server_url_clean = server_url.rstrip("/") - - try: - health = requests.get(f"{server_url_clean}/v1/service/health", timeout=3) - if health.status_code != 200: - yield gr.skip(), f"Server not healthy: {health.status_code}" - return - except requests.ConnectionError: - yield gr.skip(), f"Cannot reach server at {server_url_clean}" - return - - yield gr.skip(), (f"**WebRTC** stream started — prompt: *{prompt}*, duration: {duration:.0f}s") - - -# --------------------------------------------------------------------------- -# Gradio app -# --------------------------------------------------------------------------- - - -def build_app(server_url: str = DEFAULT_SERVER_URL) -> gr.Blocks: - with gr.Blocks(title="TeleFuser Stream Video", theme=gr.themes.Soft()) as app: - gr.Markdown("# TeleFuser Stream Video Generator") - gr.Markdown("Enter a prompt and click **Generate** to stream video from the server via WebRTC.") - - with gr.Row(): - with gr.Column(scale=1): - prompt_input = gr.Textbox( - label="Prompt", - placeholder="Describe the video you want to generate...", - lines=2, - ) - duration_input = gr.Slider( - minimum=5, - maximum=60, - value=DURATION_S, - step=5, - label="Duration (seconds)", - ) - server_url_input = gr.Textbox( - label="Server URL", - value=server_url, - ) - generate_btn = gr.Button("Generate", variant="primary", size="lg") - - with gr.Column(scale=2): - status_text = gr.Markdown("Ready.") - webrtc_html = gr.HTML( - value='
', - ) - - generate_btn.click( - fn=_generate_webrtc, - inputs=[prompt_input, server_url_input, duration_input], - outputs=[webrtc_html, status_text], - js=_WEBRTC_BTN_JS, - ) - - return app - - -def main() -> None: - parser = argparse.ArgumentParser(description="TeleFuser Stream Video Web UI") - parser.add_argument("--server-url", default=DEFAULT_SERVER_URL, help="Stream server base URL") - parser.add_argument("--port", type=int, default=7860, help="Gradio server port") - parser.add_argument("--share", action="store_true", help="Create public share link") - args = parser.parse_args() - - app = build_app(server_url=args.server_url) - app.launch(server_name="0.0.0.0", server_port=args.port, share=args.share) - - -if __name__ == "__main__": - main() From b06f32b455364e9dd6cdefc020dcc9fa1dc79ebc Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 10:47:58 +0000 Subject: [PATCH 06/11] chore(benchmark): drop unsupported WebRTC adapters Remove aiortc-specific LingBot stream benchmark contracts, configs, data, and helper scripts. Keep AIPerf documentation explicit that a validated LiveKit client adapter is required before streaming benchmarks are restored. Verification: - bash -n scripts/setup_aiperf_repo.sh - no references to removed stream benchmark assets - git diff --cached --check --- .../baseline/sglang_lingbot_stream/README.md | 53 ----- .../benchmark_contract.yaml | 104 --------- .../stream_lingbot_world_fast_compare.json | 34 --- .../stream_lingbot_world_fast_quick.json | 33 --- .../scripts/run_service.sh | 47 ---- .../scripts/run_stream_bench.sh | 49 ---- .../stream_lingbot_world_fast_compare.json | 42 ---- .../stream_lingbot_world_fast_quick.json | 42 ---- .../data/stream_lingbot_controls.json | 60 ----- .../scripts/run_stream_bench.sh | 66 ------ .../stream_benchmark_contract.yaml | 79 ------- docs/en/benchmark_aiperf.md | 188 +++------------ docs/zh/benchmark_aiperf.md | 162 +++---------- docs/zh/benchmark_aiperf_design.md | 214 +++--------------- scripts/setup_aiperf_repo.sh | 2 +- 15 files changed, 95 insertions(+), 1080 deletions(-) delete mode 100644 benchmarks/baseline/sglang_lingbot_stream/README.md delete mode 100644 benchmarks/baseline/sglang_lingbot_stream/benchmark_contract.yaml delete mode 100644 benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_compare.json delete mode 100644 benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_quick.json delete mode 100755 benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh delete mode 100755 benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh delete mode 100644 benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json delete mode 100644 benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_quick.json delete mode 100644 benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json delete mode 100755 benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh delete mode 100644 benchmarks/telefuser_aiperf/stream_benchmark_contract.yaml diff --git a/benchmarks/baseline/sglang_lingbot_stream/README.md b/benchmarks/baseline/sglang_lingbot_stream/README.md deleted file mode 100644 index 633d3d9..0000000 --- a/benchmarks/baseline/sglang_lingbot_stream/README.md +++ /dev/null @@ -1,53 +0,0 @@ -# SGLang-Diffusion LingBot Stream Baseline - -This target compares TeleFuser LingBot streaming with the diffusion runtime in -`sgl-project/sglang` (`sglang.multimodal_gen`). It uses WebSocket + MessagePack while -TeleFuser uses WebRTC + DataChannel; AIPerf normalizes both into the same session and -control timeline. - -## Requirements - -- a version-pinned SGLang checkout that provides `LingBotWorldCausalDMDPipeline`; -- `robbyant/lingbot-world-fast-diffusers` or an equivalent local model path; -- the AIPerf checkout prepared by `scripts/setup_aiperf_repo.sh`. - -The launcher does not monkeypatch SGLang internals. Missing dependencies or incompatible -CUDA kernels must be fixed in the SGLang environment or recorded as a failed -qualification, not hidden behind an unversioned shim. - -## Start the target - -```bash -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh -``` - -Common overrides: - -```bash -SGLANG_PYTHON=/path/to/venv/bin/python \ -SGLANG_LINGBOT_MODEL_PATH=/path/to/model \ -SGLANG_LINGBOT_NUM_GPUS=1 \ -SGLANG_LINGBOT_ULYSSES_DEGREE=1 \ - bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh -``` - -The default address is `http://127.0.0.1:30000`; readiness is checked at `/health`. - -## Run AIPerf - -```bash -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh - -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh \ - benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_compare.json -``` - -The baseline reuses -`benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json`. AIPerf maps the shared -directional controls to SGLang camera actions and keeps implementation-specific fields -as raw evidence. - -For a valid performance comparison, use the same accelerator count, prompt, first -frame, FPS, session window, control trace, dtype, attention/cache geometry, and offload -policy. Record the exact SGLang and model revisions. Mock, native fallback, CPU offload, -and layerwise offload runs require separate qualifications. diff --git a/benchmarks/baseline/sglang_lingbot_stream/benchmark_contract.yaml b/benchmarks/baseline/sglang_lingbot_stream/benchmark_contract.yaml deleted file mode 100644 index 488c60b..0000000 --- a/benchmarks/baseline/sglang_lingbot_stream/benchmark_contract.yaml +++ /dev/null @@ -1,104 +0,0 @@ -# Benchmark contract example for a WebSocket stream-world baseline. -# Enum comments use "one-of". List comments use "one-or-more". -contract_version: v1 # one-of: v1. Bump only when the contract schema changes. -name: sglang_lingbot_world_stream # Stable benchmark target id used in reports and automation. -mode: stream_world # one-of: batch_video, stream_world. -implementation: sglang_diffusion # Example values: telefuser, diffusers, sglang_diffusion. -model_family: lingbot_world_fast # Example values: wan_video, lingbot_world_fast, hunyuan_video, ltx_video. -model: robbyant/lingbot-world-fast-diffusers # Concrete model or service profile under test. -supported_tasks: # one-or-more: t2v, i2v, ti2v, bidirectional. - - bidirectional -transport: websocket # one-of: http, http_polling, websocket, sse, webrtc. -adapter: sglang_websocket # Built-in AIPerf adapter; target repository owns no transport implementation. -endpoint: - health_path: /health # Service readiness path. - metadata_path: /v1/models # Optional SGLang model and pipeline identity snapshot. - websocket_path: /v1/realtime_video/generate # WebSocket path for realtime video sessions. - models_path: /v1/models # Optional model metadata path. -request_encoding: - message_format: msgpack # one-of for this baseline: msgpack. - init_required_fields: # one-or-more. Required fields in the first WebSocket message. - - type - - prompt - - first_frame - - size - - fps - - num_frames - init_parameters: # Mapping from benchmark parameter names to SGLang init payload fields. - model: model - prompt: prompt - image_path: first_frame - size: size - fps: fps - num_frames: num_frames - control_channel: - transport: websocket_message # one-of for WebSocket: websocket_message. - message_type: event # Runtime control messages use type=event. - kind: camera_actions # one-of for LingBot realtime controls: camera_actions, prompt. - payload_mode: state # one-of for camera_actions: script, state. - action_tokens: # one-or-more. SGLang LingBot camera action tokens. - - w - - a - - s - - d - - i - - j - - k - - l - key_mapping: - ArrowUp: w - ArrowDown: s - ArrowLeft: a - ArrowRight: d -result_delivery: - media: websocket_frame_batch # one-of for this baseline: websocket_frame_batch. - metadata: websocket_chunk_stats # one-of for this baseline: websocket_chunk_stats. - session_log: sessions.jsonl # Per-session result records. - event_log: events/{phase}_{logical_session_index}_{session_id}.jsonl # Per-session event trace template. -workload: - mode: bidirectional # one-of: server_push, bidirectional. - task: bidirectional # one-of for this service: bidirectional. - size: 832x480 - fps: 16 - session_count: 1 - warmup_sessions: 1 - session_duration_s: 90.0 - control_trace: benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json # Timed control-event workload. - request_extra: # Service-specific request config passed through to SGLang. - num_frames: 9 - num_inference_steps: 4 - guidance_scale: 1.0 - realtime_causal_sink_size: 6 - realtime_causal_kv_cache_num_frames: 9 - realtime_output_format: webp # one-of: webp, jpeg, raw. - output_compression: 95 - max_chunks: 8 -metrics: # one-or-more. Choose all metrics emitted by this benchmark mode. - - connected_latency_ms - - first_frame_latency_ms - - first_metadata_latency_ms - - stream_fps - - session_runtime_s - - frames_received - - control_ack_latency_ms - - control_to_next_frame_latency_ms - - chunk_request_prepare_seconds - - chunk_compute_seconds - - chunk_encode_seconds - - chunk_output_pacing_seconds - - chunk_output_header_write_seconds - - chunk_output_payload_write_seconds - - chunk_output_write_seconds - - chunk_total_seconds - - chunk_compute_fps - - chunk_raw_output_bytes - - chunk_wire_output_bytes - - chunk_output_batches - - chunk_peak_reserved_bytes - - success_rate -limits: - active_sessions: 1 # Current SGLang realtime endpoint accepts one active session for this target. -artifacts: # Paths consumed by automation and documentation. - config: benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_compare.json - control_trace: benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json - runner: aiperf profile --stream-config diff --git a/benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_compare.json b/benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_compare.json deleted file mode 100644 index 76597d4..0000000 --- a/benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_compare.json +++ /dev/null @@ -1,34 +0,0 @@ -{ - "contract": "benchmarks/baseline/sglang_lingbot_stream/benchmark_contract.yaml", - "server_url": "http://127.0.0.1:30000", - "mode": "bidirectional", - "task": "bidirectional", - "prompt": "walk forward through the scene", - "image_path": "examples/data/1.png", - "fps": 16, - "session_count": 1, - "warmup_sessions": 1, - "warmup_chunks": 1, - "session_duration_s": 90.0, - "stagger_s": 0.0, - "control_trace_path": "benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json", - "request_extra": { - "num_frames": 9, - "num_inference_steps": 4, - "guidance_scale": 1.0, - "realtime_causal_sink_size": 6, - "realtime_causal_kv_cache_num_frames": 9, - "realtime_output_format": "webp", - "output_compression": 95, - "max_chunks": 8, - "realtime_output_pacing": false - }, - "transport": { - "connect_timeout_s": 60.0, - "message_timeout_s": 180.0 - }, - "server_metrics": { - "enabled": false - }, - "artifacts_dir": "artifacts/sglang_lingbot_stream/stream_lingbot_compare" -} diff --git a/benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_quick.json b/benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_quick.json deleted file mode 100644 index 6ad373c..0000000 --- a/benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_quick.json +++ /dev/null @@ -1,33 +0,0 @@ -{ - "contract": "benchmarks/baseline/sglang_lingbot_stream/benchmark_contract.yaml", - "server_url": "http://127.0.0.1:30000", - "mode": "bidirectional", - "task": "bidirectional", - "prompt": "walk forward through the scene", - "image_path": "examples/data/1.png", - "fps": 16, - "session_count": 1, - "warmup_sessions": 0, - "session_duration_s": 12.0, - "stagger_s": 0.0, - "control_trace_path": "benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json", - "request_extra": { - "num_frames": 9, - "num_inference_steps": 4, - "guidance_scale": 1.0, - "realtime_causal_sink_size": 6, - "realtime_causal_kv_cache_num_frames": 9, - "realtime_output_format": "webp", - "output_compression": 95, - "max_chunks": 8, - "realtime_output_pacing": false - }, - "transport": { - "connect_timeout_s": 30.0, - "message_timeout_s": 120.0 - }, - "server_metrics": { - "enabled": false - }, - "artifacts_dir": "artifacts/sglang_lingbot_stream/stream_lingbot_quick" -} diff --git a/benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh b/benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh deleted file mode 100755 index fe34921..0000000 --- a/benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh +++ /dev/null @@ -1,47 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail - -ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd)" -cd "${ROOT_DIR}" - -if [[ -z "${CUDA_HOME:-}" && -d /usr/local/cuda ]]; then - export CUDA_HOME=/usr/local/cuda -fi - -if [[ -n "${SGLANG_EXTRA_PYTHONPATH:-}" ]]; then - export PYTHONPATH="${SGLANG_EXTRA_PYTHONPATH}${PYTHONPATH:+:${PYTHONPATH}}" -fi - -SGLANG_BIN="${SGLANG_BIN:-sglang}" -SGLANG_PYTHON="${SGLANG_PYTHON:-}" -SERVICE_PORT="${SGLANG_LINGBOT_PORT:-30000}" -MODEL_PATH="${SGLANG_LINGBOT_MODEL_PATH:-robbyant/lingbot-world-fast-diffusers}" -MODEL_ID="${SGLANG_LINGBOT_MODEL_ID:-lingbot-world-fast-diffusers}" -MODEL_TYPE="${SGLANG_LINGBOT_MODEL_TYPE:-diffusion}" -PIPELINE_CLASS="${SGLANG_LINGBOT_PIPELINE_CLASS:-LingBotWorldCausalDMDPipeline}" -PERFORMANCE_MODE="${SGLANG_LINGBOT_PERFORMANCE_MODE:-speed}" -ATTENTION_BACKEND_CONFIG="${SGLANG_LINGBOT_ATTENTION_BACKEND_CONFIG:-VSA_sparsity=0.0}" -NUM_GPUS="${SGLANG_LINGBOT_NUM_GPUS:-1}" -ULYSSES_DEGREE="${SGLANG_LINGBOT_ULYSSES_DEGREE:-1}" -DIT_CPU_OFFLOAD="${SGLANG_LINGBOT_DIT_CPU_OFFLOAD:-false}" -TEXT_ENCODER_CPU_OFFLOAD="${SGLANG_LINGBOT_TEXT_ENCODER_CPU_OFFLOAD:-false}" - -if [[ -n "${SGLANG_PYTHON}" ]]; then - SGLANG_CMD=("${SGLANG_PYTHON}" -c "from sglang.cli.main import main; main()") -else - read -r -a SGLANG_CMD <<< "${SGLANG_BIN}" -fi - -exec "${SGLANG_CMD[@]}" serve \ - --model-type "${MODEL_TYPE}" \ - --model-path "${MODEL_PATH}" \ - --model-id "${MODEL_ID}" \ - --pipeline-class-name "${PIPELINE_CLASS}" \ - --performance-mode "${PERFORMANCE_MODE}" \ - --attention-backend-config "${ATTENTION_BACKEND_CONFIG}" \ - --port "${SERVICE_PORT}" \ - --num-gpus "${NUM_GPUS}" \ - --ulysses-degree "${ULYSSES_DEGREE}" \ - --dit-cpu-offload "${DIT_CPU_OFFLOAD}" \ - --text-encoder-cpu-offload "${TEXT_ENCODER_CPU_OFFLOAD}" \ - "$@" diff --git a/benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh b/benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh deleted file mode 100755 index 1385d59..0000000 --- a/benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh +++ /dev/null @@ -1,49 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail - -ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd)" -cd "${ROOT_DIR}" - -AIPERF_DIR="${ROOT_DIR}/benchmarks/aiperf" -UV_BIN="${AIPERF_UV_BIN:-uv}" -CONFIG_PATH="${1:-benchmarks/baseline/sglang_lingbot_stream/configs/stream_lingbot_world_fast_quick.json}" -if [[ $# -gt 0 ]]; then - shift -fi - -if [[ ! -f "${AIPERF_DIR}/pyproject.toml" ]]; then - echo "AIPerf checkout not found. Run: bash scripts/setup_aiperf_repo.sh" >&2 - exit 1 -fi -if ! command -v "${UV_BIN}" >/dev/null 2>&1; then - echo "uv is required: https://docs.astral.sh/uv/getting-started/installation/" >&2 - exit 1 -fi - -SERVER_URL="${SGLANG_STREAM_BENCH_URL:-http://127.0.0.1:30000}" -SERVER_ARGS=(--stream-server-url "${SERVER_URL}") -for argument in "$@"; do - if [[ "${argument}" == "--stream-server-url" || "${argument}" == --stream-server-url=* ]]; then - SERVER_ARGS=() - break - fi -done - -RESOURCE_ARGS=() -RESOURCE_HISTORY_URL="${AIPERF_HISTORY_URL:-}" -RESOURCE_TARGET_PID="${AIPERF_RESOURCE_TARGET_PID:-${SGLANG_STREAM_BENCH_PID:-}}" -if [[ -n "${RESOURCE_HISTORY_URL}" || -n "${RESOURCE_TARGET_PID}" ]]; then - if [[ -z "${RESOURCE_HISTORY_URL}" || -z "${RESOURCE_TARGET_PID}" ]]; then - echo "AIPERF_HISTORY_URL and AIPERF_RESOURCE_TARGET_PID must be set together" >&2 - exit 2 - fi - RESOURCE_ARGS+=(--stream-resource-history-url "${RESOURCE_HISTORY_URL}") - RESOURCE_ARGS+=(--stream-resource-target-pid "${RESOURCE_TARGET_PID}") -fi - -exec "${UV_BIN}" run --frozen --no-dev --project "${AIPERF_DIR}" \ - aiperf profile \ - --stream-config "${CONFIG_PATH}" \ - "${SERVER_ARGS[@]}" \ - "${RESOURCE_ARGS[@]}" \ - "$@" diff --git a/benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json b/benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json deleted file mode 100644 index e922f1f..0000000 --- a/benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json +++ /dev/null @@ -1,42 +0,0 @@ -{ - "contract": "benchmarks/telefuser_aiperf/stream_benchmark_contract.yaml", - "server_url": "http://127.0.0.1:8088", - "mode": "bidirectional", - "task": "bidirectional", - "prompt": "walk forward through the scene", - "image_path": "examples/data/1.png", - "fps": 16, - "session_count": 1, - "warmup_sessions": 1, - "warmup_chunks": 1, - "session_duration_s": 90.0, - "stagger_s": 0.0, - "control_trace_path": "benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json", - "request_extra": { - "chunk_size": 3, - "frame_num": 81, - "sample_shift": 5.0, - "control_mode": "cam", - "show_control_hud": false, - "benchmark_metrics": true - }, - "transport": { - "connect_timeout_s": 60.0, - "frame_timeout_s": 180.0, - "ice_gather_timeout_s": 5.0, - "shutdown_timeout_s": 5.0, - "receive_audio": false - }, - "server_metrics": { - "enabled": true, - "urls": [ - "http://127.0.0.1:8088/v1/service/metrics" - ], - "collection_interval_s": 1.0, - "export_raw_jsonl": true - }, - "observability": { - "mapping": "builtin:telefuser" - }, - "artifacts_dir": "artifacts/telefuser_aiperf/stream_lingbot_compare" -} diff --git a/benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_quick.json b/benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_quick.json deleted file mode 100644 index 31ff159..0000000 --- a/benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_quick.json +++ /dev/null @@ -1,42 +0,0 @@ -{ - "contract": "benchmarks/telefuser_aiperf/stream_benchmark_contract.yaml", - "server_url": "http://127.0.0.1:8088", - "mode": "bidirectional", - "task": "bidirectional", - "prompt": "walk forward through the scene", - "image_path": "examples/data/1.png", - "fps": 16, - "session_count": 1, - "warmup_sessions": 0, - "warmup_chunks": 1, - "session_duration_s": 30.0, - "stagger_s": 0.0, - "control_trace_path": "benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json", - "request_extra": { - "chunk_size": 3, - "frame_num": 81, - "sample_shift": 5.0, - "control_mode": "cam", - "show_control_hud": false, - "benchmark_metrics": true - }, - "transport": { - "connect_timeout_s": 30.0, - "frame_timeout_s": 60.0, - "ice_gather_timeout_s": 5.0, - "shutdown_timeout_s": 5.0, - "receive_audio": false - }, - "server_metrics": { - "enabled": true, - "urls": [ - "http://127.0.0.1:8088/v1/service/metrics" - ], - "collection_interval_s": 1.0, - "export_raw_jsonl": true - }, - "observability": { - "mapping": "builtin:telefuser" - }, - "artifacts_dir": "artifacts/telefuser_aiperf/stream_lingbot_quick" -} diff --git a/benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json b/benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json deleted file mode 100644 index d70b5dd..0000000 --- a/benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json +++ /dev/null @@ -1,60 +0,0 @@ -{ - "events": [ - { - "delay_s": 1.0, - "message": { - "type": "control", - "key": "ArrowUp", - "action": "press" - } - }, - { - "delay_s": 1.8, - "message": { - "type": "control", - "key": "ArrowUp", - "action": "release" - } - }, - { - "delay_s": 2.8, - "message": { - "type": "control", - "key": "ArrowLeft", - "action": "press" - } - }, - { - "delay_s": 3.6, - "message": { - "type": "control", - "key": "ArrowLeft", - "action": "release" - } - }, - { - "delay_s": 4.6, - "message": { - "type": "control", - "key": "ArrowRight", - "action": "press" - } - }, - { - "delay_s": 5.4, - "message": { - "type": "control", - "key": "ArrowRight", - "action": "release" - } - }, - { - "delay_s": 6.4, - "message": { - "type": "control", - "key": "ArrowUp", - "action": "press" - } - } - ] -} diff --git a/benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh b/benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh deleted file mode 100755 index 92e9947..0000000 --- a/benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh +++ /dev/null @@ -1,66 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail - -ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)" -cd "${ROOT_DIR}" - -AIPERF_DIR="${ROOT_DIR}/benchmarks/aiperf" -UV_BIN="${AIPERF_UV_BIN:-uv}" -CONFIG_PATH="${1:-benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_quick.json}" -if [[ $# -gt 0 ]]; then - shift -fi - -if [[ ! -f "${AIPERF_DIR}/pyproject.toml" ]]; then - echo "AIPerf checkout not found. Run: bash scripts/setup_aiperf_repo.sh" >&2 - exit 1 -fi -if ! command -v "${UV_BIN}" >/dev/null 2>&1; then - echo "uv is required: https://docs.astral.sh/uv/getting-started/installation/" >&2 - exit 1 -fi - -SERVER_URL="${TELEFUSER_STREAM_BENCH_URL:-http://127.0.0.1:8088}" -SERVER_ARGS=(--stream-server-url "${SERVER_URL}") -for argument in "$@"; do - if [[ "${argument}" == "--stream-server-url" || "${argument}" == --stream-server-url=* ]]; then - SERVER_ARGS=() - break - fi -done -ICE_HOST_IPS="${TELEFUSER_STREAM_BENCH_ICE_HOST_IPS:-}" -ICE_HOST_ARGS=() -if [[ -n "${ICE_HOST_IPS}" ]]; then - IFS=',' read -r -a _ICE_HOST_IP_ARRAY <<< "${ICE_HOST_IPS}" - for ice_host_ip in "${_ICE_HOST_IP_ARRAY[@]}"; do - if [[ -n "${ice_host_ip}" ]]; then - ICE_HOST_ARGS+=(--stream-ice-host-ip "${ice_host_ip}") - fi - done -fi - -METRICS_ARGS=() -if [[ -n "${TELEFUSER_STREAM_BENCH_METRICS_URL:-}" ]]; then - METRICS_ARGS+=(--stream-server-metrics-url "${TELEFUSER_STREAM_BENCH_METRICS_URL}") -fi - -RESOURCE_ARGS=() -RESOURCE_HISTORY_URL="${AIPERF_HISTORY_URL:-}" -RESOURCE_TARGET_PID="${AIPERF_RESOURCE_TARGET_PID:-${TELEFUSER_STREAM_BENCH_PID:-}}" -if [[ -n "${RESOURCE_HISTORY_URL}" || -n "${RESOURCE_TARGET_PID}" ]]; then - if [[ -z "${RESOURCE_HISTORY_URL}" || -z "${RESOURCE_TARGET_PID}" ]]; then - echo "AIPERF_HISTORY_URL and AIPERF_RESOURCE_TARGET_PID must be set together" >&2 - exit 2 - fi - RESOURCE_ARGS+=(--stream-resource-history-url "${RESOURCE_HISTORY_URL}") - RESOURCE_ARGS+=(--stream-resource-target-pid "${RESOURCE_TARGET_PID}") -fi - -exec "${UV_BIN}" run --frozen --no-dev --project "${AIPERF_DIR}" --extra streaming-webrtc \ - aiperf profile \ - --stream-config "${CONFIG_PATH}" \ - "${SERVER_ARGS[@]}" \ - "${ICE_HOST_ARGS[@]}" \ - "${METRICS_ARGS[@]}" \ - "${RESOURCE_ARGS[@]}" \ - "$@" diff --git a/benchmarks/telefuser_aiperf/stream_benchmark_contract.yaml b/benchmarks/telefuser_aiperf/stream_benchmark_contract.yaml deleted file mode 100644 index be088d2..0000000 --- a/benchmarks/telefuser_aiperf/stream_benchmark_contract.yaml +++ /dev/null @@ -1,79 +0,0 @@ -# Benchmark contract example for a WebRTC stream-world target. -# Enum comments use "one-of". List comments use "one-or-more". -contract_version: v1 # one-of: v1. Bump only when the contract schema changes. -name: telefuser_lingbot_world_fast_stream # Stable benchmark target id used in reports and automation. -mode: stream_world # one-of: batch_video, stream_world. -implementation: telefuser # Example values: telefuser, diffusers, sglang_diffusion. -model_family: lingbot_world_fast # Example values: wan_video, lingbot_world_fast, hunyuan_video, ltx_video. -model: LingBot-World-Fast # Concrete model or service profile under test. -supported_tasks: # one-or-more: t2v, i2v, ti2v, bidirectional. - - bidirectional -transport: webrtc # one-of: http, http_polling, websocket, sse, webrtc. -adapter: telefuser_webrtc # Built-in AIPerf adapter; target repository owns no transport implementation. -endpoint: - health_path: /v1/service/health # Service readiness path. - metadata_path: /v1/service/metadata # Optional target environment and startup phase facts. - offer_path: /v1/stream/webrtc/offer # WebRTC SDP offer/answer path. - delete_path_template: /v1/stream/webrtc/{session_id} # Session cleanup path template. -request_encoding: - offer_content_type: application/json # one-of: application/json. - offer_required_fields: # one-or-more. Required JSON fields in the offer payload. - - sdp - - type - - task - offer_parameters: # Mapping from benchmark parameter names to wire payload fields. - session_id: session_id - prompt: prompt - fps: fps - image_path: image_path - config: config - control_channel: - transport: datachannel # one-of for WebRTC: datachannel. - label: telefuser # DataChannel label expected by the service. - message_types: # one-or-more. Message types understood by the stream benchmark semantics. - - control - - status - - chunk - - done - - error -result_delivery: - media: rtp_video_track # one-of for WebRTC: rtp_video_track. - metadata: datachannel # one-of for WebRTC: datachannel. - session_log: sessions.jsonl # Per-session result records. - event_log: events/{phase}_{logical_session_index}_{session_id}.jsonl # Per-session event trace template. -workload: - mode: bidirectional # one-of: server_push, bidirectional. - task: bidirectional # one-of for this service: bidirectional. - fps: 16 - session_count: 1 - warmup_sessions: 1 - session_duration_s: 90.0 - control_trace: benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json # Timed control-event workload. - request_extra: # Service-specific request config passed through to the stream service. - chunk_size: 3 - frame_num: 81 - sample_shift: 5.0 - control_mode: cam # one-of for LingBotWorldFast: cam. - show_control_hud: false - benchmark_metrics: true # Emit synchronized runtime/chunk facts for AIPerf aggregation. -metrics: # one-or-more. Choose all metrics emitted by this benchmark mode. - - offer_rtt_ms - - connected_latency_ms - - first_frame_latency_ms - - first_metadata_latency_ms - - stream_fps - - session_runtime_s - - frames_received - - control_ack_latency_ms - - control_to_next_frame_latency_ms - - pipeline_init_seconds - - runtime_creation_seconds - - chunk_compute_seconds - - chunk_compute_fps - - success_rate -limits: - active_sessions: 1 # Service limit for this target. The harness still exposes session_count. -artifacts: # Paths consumed by automation and documentation. - config: benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json - control_trace: benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json - runner: aiperf profile --stream-config diff --git a/docs/en/benchmark_aiperf.md b/docs/en/benchmark_aiperf.md index 9c9f1c9..2cf98c0 100644 --- a/docs/en/benchmark_aiperf.md +++ b/docs/en/benchmark_aiperf.md @@ -1,48 +1,40 @@ # TeleFuser and AIPerf -TeleFuser exposes raw target-side facts; AIPerf owns workload execution, aggregation, -resource collection, artifacts, GreptimeDB history, and visualization. This separation -keeps the same benchmark and dashboard reusable across TeleFuser, SGLang-Diffusion, and -future targets. +TeleFuser exposes raw target-side facts; AIPerf owns workload execution, aggregation, resource collection, artifacts, +GreptimeDB history, and visualization. The checked-in integration currently covers batch video generation through +the OpenAI-compatible `/v1/videos` API. -The included assets cover: +The former LingBot streaming adapter and its SGLang comparison assets were removed with the legacy transport backend. +A LiveKit benchmark adapter has not been added to AIPerf yet, so this repository does not present an +unsupported stream benchmark as runnable. Target compute metrics emitted by streaming services remain available to +future LiveKit-aware benchmark clients. -- Wan2.1 image-to-video through the OpenAI-compatible `/v1/videos` API; -- LingBot-World-Fast sessions through WebRTC and DataChannel; -- a LingBot SGLang-Diffusion baseline through WebSocket and MessagePack. - -## Repository layout +## Repository boundary ```text benchmarks/ -├── telefuser_aiperf/ # TeleFuser contracts, configs, data, launchers -├── baseline/sglang_lingbot_stream/ # Stream baseline -└── aiperf/ # Ignored external AIPerf checkout +├── telefuser_aiperf/ # Batch contracts, configs, data, and launcher +└── aiperf/ # Ignored external AIPerf checkout ``` -The AIPerf implementation is not vendored into TeleFuser. The setup script always uses -`/benchmarks/aiperf`; neither the setup script nor the launchers accept a -checkout-path override. Install -[uv](https://docs.astral.sh/uv/getting-started/installation/), then run this once from -the TeleFuser repository root: +The AIPerf implementation is not vendored. Install +[uv](https://docs.astral.sh/uv/getting-started/installation/) and run from the TeleFuser repository root: ```bash bash scripts/setup_aiperf_repo.sh ``` -The script clones AIPerf, creates its isolated runtime environment with WebRTC support, -and creates `/artifacts` for benchmark output and History imports. The -dashboard is bundled, so runtime users do not need Node.js or a separate frontend -process. Pin a commit for reproducible runs: +The script clones AIPerf into `benchmarks/aiperf`, installs its non-development runtime, and creates `artifacts/`. +Pin a commit for reproducible runs: ```bash AIPERF_REF= bash scripts/setup_aiperf_repo.sh ``` -`AIPERF_REPO_URL`, `AIPERF_BRANCH`, and `AIPERF_REF` may select the source and revision, -but never change the checkout location. +`AIPERF_REPO_URL`, `AIPERF_BRANCH`, and `AIPERF_REF` may select the source and revision, but never change the checkout +location. -## Batch video +## Batch video benchmark Start the fixed Wan2.1 I2V target: @@ -53,7 +45,7 @@ telefuser serve \ --task i2v ``` -Run a smoke profile or the fixed comparison workload: +Run a smoke profile or fixed comparison workload: ```bash bash benchmarks/telefuser_aiperf/scripts/run_video_bench.sh @@ -62,92 +54,20 @@ bash benchmarks/telefuser_aiperf/scripts/run_video_bench.sh \ benchmarks/telefuser_aiperf/configs/video_generation_wan21_i2v_480p_compare.yaml ``` -The launcher checks `/v1/service/health` before profiling. Common overrides include -`TELEFUSER_AIPERF_URL`, `TELEFUSER_AIPERF_CONCURRENCY`, -`TELEFUSER_AIPERF_REQUESTS`, `TELEFUSER_AIPERF_SIZE`, and +The launcher checks `/v1/service/health` before profiling. Common overrides include `TELEFUSER_AIPERF_URL`, +`TELEFUSER_AIPERF_CONCURRENCY`, `TELEFUSER_AIPERF_REQUESTS`, `TELEFUSER_AIPERF_SIZE`, and `TELEFUSER_AIPERF_SECONDS`. -## LingBot stream - -Start TeleFuser: - -```bash -telefuser stream-serve \ - examples/lingbot/lingbot_world_fast_image_to_video_h100.py \ - -p 8088 \ - --skip-validation -``` - -Then run: - -```bash -bash benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh - -bash benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh \ - benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json -``` - -The stream config enables `benchmark_metrics`. TeleFuser then reports synchronized raw -facts for runtime creation, actor-graph chunk compute, cache geometry, and environment -identity. Allocator peaks are omitted because generation runs in child actors and the -service-process allocator cannot represent the complete graph; active AIPerf resource -telemetry supplies process-tree GPU-memory curves instead. The native WebRTC path does -not report a separate payload encoding duration because encoding happens after the target -chunk fact. AIPerf computes warmup-aware summaries and keeps client delivery separate -from target compute. - -The shared timed control trace is stored at -`benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json`. - -## SGLang-Diffusion baseline - -Use a compatible, version-pinned `sgl-project/sglang` environment. The TeleFuser tree -does not patch SGLang modules at import time. - -```bash -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh -``` - -The baseline uses the same prompt, first frame, FPS target, session window, and control -trace. The adapter translates only transport semantics. For a performance comparison, -record the exact SGLang commit and model revision, use GPU-resident speed mode, and keep -offload and fallback settings identical. An OOM is a result for that configuration; do -not replace it with a mock or offloaded result under the same label. - -Both documented launch commands default to one GPU. Override both targets explicitly -when comparing another accelerator count. - -## Configs - | Config | Purpose | |---|---| -| `video_generation_quick.yaml` | Batch connectivity and latency smoke test | -| `video_generation_e2e.yaml` | Batch warmup, trace, records, and server metrics | -| `video_generation_rate.yaml` | Poisson-arrival Batch load | +| `video_generation_quick.yaml` | Connectivity and latency smoke test | +| `video_generation_e2e.yaml` | Warmup, trace, records, and server metrics | +| `video_generation_rate.yaml` | Poisson-arrival load | | `video_generation_wan21_i2v_480p_compare.yaml` | Fixed Wan I2V comparison | -| `stream_lingbot_world_fast_quick.json` | Bounded Stream smoke test | -| `stream_lingbot_world_fast_compare.json` | Fixed LingBot Stream comparison | - -SGLang equivalents are under `benchmarks/baseline/sglang_lingbot_stream/configs`. - -## Metric interpretation - -The most important distinction is scope: - -| Metric | Meaning | -|---|---| -| `stream_fps` | Frames received by the client divided by client session time | -| `chunk_compute_fps` | Frames divided by compute time for one target chunk | -| `chunk_compute_fps_weighted` | `sum(frames) / sum(compute_seconds)` after warmup exclusion | - -AIPerf presents metrics under five stable dimensions: delivery, latency, throughput, -target execution, and resources. Implementation-specific fields remain raw evidence and -map into these canonical leaves; they do not become separate top-level metrics. ## Active resource history -Docker provides the shortest persistent GreptimeDB setup: +Start persistent GreptimeDB storage: ```bash docker volume create aiperf-greptime-data @@ -160,9 +80,7 @@ docker run -d --name aiperf-greptime --restart unless-stopped \ --data-home /greptimedb_data ``` -Pin the image tag or digest for production. The named volume keeps history across -container restarts. Then start the bundled AIPerf API and frontend from the TeleFuser -repository root; the `artifacts` root matches the benchmark launchers: +Then start the bundled AIPerf history API and dashboard: ```bash uv run --frozen --no-dev --project benchmarks/aiperf aiperf history serve \ @@ -173,54 +91,14 @@ uv run --frozen --no-dev --project benchmarks/aiperf aiperf history serve \ --port 8095 ``` -Verify the stack: - -```bash -curl --fail http://127.0.0.1:8095/api/v1/history/health -curl --fail -X POST 'http://127.0.0.1:4000/v1/sql?db=public' \ - --data-urlencode 'sql=SELECT 1 AS ready' -``` - -Enable active collection on the target host: - -```bash -export AIPERF_HISTORY_URL=http://:8095 -export AIPERF_RESOURCE_TARGET_PID= - -bash benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh -``` - -The agent recursively observes the target process tree. It samples every second, -uploads every 15 seconds, timestamps samples at the source, and flushes at termination. -It reports process, container-when-detectable, and machine facts for CPU, memory, GPU, -VRAM, Ethernet, and RDMA. Capacity is kept separate from usage. - -GreptimeDB is mandatory for History and active reporting. Startup, query, or final-flush -failure is surfaced; there is no SQLite, in-memory, or direct-file query fallback. - -Open `http://127.0.0.1:8095/` for the Chinese desktop dashboard. It supports two run -groups, canonical metric-tree selection, aggregate curves, resource timelines, and -cross-run comparison. - -For a remote benchmark host, keep the default loopback bind and forward it securely: - -```bash -ssh -L 8095:127.0.0.1:8095 user@benchmark-host -``` - -## Artifacts and reproducibility - -Batch and stream launchers write timestamped artifacts below `artifacts/`. Stream -artifacts include summaries, session and event JSONL, target metadata, normalized -metrics, and a standalone HTML report. +For active process-tree collection, export `AIPERF_HISTORY_URL` and `AIPERF_RESOURCE_TARGET_PID` before running the +batch launcher. GreptimeDB is required for History and active reporting; failures do not silently fall back to an +in-memory or file-only database. -Every performance result should retain: +## Reproducibility -- TeleFuser or SGLang commit and model revision; -- accelerator model/count, driver, CUDA, PyTorch, and dtype; -- workload config and control trace; -- warmup policy and successful/failed session counts; -- offload, cache, attention, and fallback settings. +Every result should retain the TeleFuser and AIPerf commits, model revision, accelerator model/count, driver, CUDA, +PyTorch, dtype, complete workload config, warmup rule, success/failure counts, and offload/cache/attention settings. +Dynamic results belong in GreptimeDB and replayable artifacts, not stable documentation. -See the Chinese [benchmark design](/TeleFuser/zh/benchmark_aiperf_design/) for protocol and -ownership details. +The stable responsibilities and metric boundary are summarized above; dynamic results remain outside this guide. diff --git a/docs/zh/benchmark_aiperf.md b/docs/zh/benchmark_aiperf.md index 6fb8d04..f3ef071 100644 --- a/docs/zh/benchmark_aiperf.md +++ b/docs/zh/benchmark_aiperf.md @@ -1,40 +1,35 @@ # TeleFuser 与 AIPerf -TeleFuser 只暴露目标侧原始事实;AIPerf 统一负责 workload 执行、指标聚合、资源采集、产物、GreptimeDB -历史服务和前端展示。这样同一套 benchmark 与界面可以复用于 TeleFuser、SGLang-Diffusion 和后续实现。 +TeleFuser 只暴露目标侧原始事实;AIPerf 负责 workload 执行、聚合、资源采集、产物、GreptimeDB 历史服务和 +展示。仓库内当前集成只覆盖通过 OpenAI 兼容 `/v1/videos` API 执行的 batch 视频生成。 -当前资产覆盖: - -- 通过 OpenAI 兼容 `/v1/videos` API 测试 Wan2.1 图生视频; -- 通过 WebRTC 与 DataChannel 测试 LingBot-World-Fast; -- 通过 WebSocket 与 MessagePack 测试 SGLang-Diffusion LingBot baseline。 +随旧传输后端一起删除的内容包括 LingBot 直接 WebRTC adapter 和 SGLang 对比资产。AIPerf 目前 +尚未集成 LiveKit benchmark adapter,因此本仓库不会把不受支持的 stream benchmark 标记为可运行。流服务 +输出的 target compute 指标仍可供未来 LiveKit-aware benchmark client 使用。 ## 仓库边界 ```text benchmarks/ -├── telefuser_aiperf/ # TeleFuser contract、配置、数据和启动器 -├── baseline/sglang_lingbot_stream/ # Stream baseline -└── aiperf/ # 被 Git 忽略的外部 AIPerf checkout +├── telefuser_aiperf/ # Batch contract、配置、数据和 launcher +└── aiperf/ # 被 Git 忽略的外部 AIPerf checkout ``` -TeleFuser 不 vendoring AIPerf 实现。安装脚本与所有 launcher 都固定使用 -`/benchmarks/aiperf`,不提供 checkout 路径覆盖。先安装 -[uv](https://docs.astral.sh/uv/getting-started/installation/),然后在 TeleFuser 仓库根目录执行一次: +TeleFuser 不 vendoring AIPerf。安装 [uv](https://docs.astral.sh/uv/getting-started/installation/) 后,在仓库 +根目录执行: ```bash bash scripts/setup_aiperf_repo.sh ``` -脚本会 clone AIPerf、创建包含 WebRTC 支持的隔离运行环境,并创建用于 benchmark 输出与 History -导入的 `/artifacts`。前端产物已经内置,普通用户不需要安装 Node.js,也不需要单独启动 -前端进程。正式实验应固定 AIPerf commit: +脚本把 AIPerf clone 到 `benchmarks/aiperf`,安装其非开发运行环境,并创建 `artifacts/`。正式实验应固定 +commit: ```bash AIPERF_REF= bash scripts/setup_aiperf_repo.sh ``` -`AIPERF_REPO_URL`、`AIPERF_BRANCH` 和 `AIPERF_REF` 只控制来源与 revision,不改变 checkout 位置。 +`AIPERF_REPO_URL`、`AIPERF_BRANCH` 和 `AIPERF_REF` 可以选择来源与 revision,但不改变 checkout 位置。 ## Batch 视频测试 @@ -47,7 +42,7 @@ telefuser serve \ --task i2v ``` -执行快速测试或固定对比 workload: +执行 smoke profile 或固定对比 workload: ```bash bash benchmarks/telefuser_aiperf/scripts/run_video_bench.sh @@ -56,84 +51,20 @@ bash benchmarks/telefuser_aiperf/scripts/run_video_bench.sh \ benchmarks/telefuser_aiperf/configs/video_generation_wan21_i2v_480p_compare.yaml ``` -启动器会先检查 `/v1/service/health`。常用覆盖变量包括 -`TELEFUSER_AIPERF_URL`、`TELEFUSER_AIPERF_CONCURRENCY`、 -`TELEFUSER_AIPERF_REQUESTS`、`TELEFUSER_AIPERF_SIZE` 和 +Launcher 会先检查 `/v1/service/health`。常用覆盖变量包括 `TELEFUSER_AIPERF_URL`、 +`TELEFUSER_AIPERF_CONCURRENCY`、`TELEFUSER_AIPERF_REQUESTS`、`TELEFUSER_AIPERF_SIZE` 和 `TELEFUSER_AIPERF_SECONDS`。 -## LingBot Stream 测试 - -启动 TeleFuser: - -```bash -telefuser stream-serve \ - examples/lingbot/lingbot_world_fast_image_to_video_h100.py \ - -p 8088 \ - --skip-validation -``` - -执行测试: - -```bash -bash benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh - -bash benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh \ - benchmarks/telefuser_aiperf/configs/stream_lingbot_world_fast_compare.json -``` - -Stream 配置通过 `benchmark_metrics: true` 开启目标侧原始事实。TeleFuser 同步记录 runtime 创建、actor graph -chunk 计算、cache 几何和运行环境。生成工作位于子 actor 中,服务进程的 allocator 无法代表完整 actor graph, -因此不对 LingBot 上报不完整的 allocator 峰值;完整进程树显存曲线由 AIPerf 主动资源采集提供。原生 WebRTC -的编码位于 chunk fact 之后,因此当前不伪造独立编码耗时。AIPerf 负责跳过 warmup 并生成聚合结果。 - -TeleFuser 与 SGLang 共用 -`benchmarks/telefuser_aiperf/data/stream_lingbot_controls.json` 中的定时控制 trace。 - -## SGLang-Diffusion baseline - -使用兼容且固定版本的 `sgl-project/sglang` 环境。TeleFuser 仓库不会在 import 时 monkeypatch SGLang -内部模块。 - -```bash -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_service.sh -bash benchmarks/baseline/sglang_lingbot_stream/scripts/run_stream_bench.sh -``` - -Baseline 固定 prompt、首帧、FPS、session 时长和控制 trace,只由 adapter 转换 transport 语义。正式性能 -对比必须记录 SGLang commit 与模型 revision,使用 GPU-resident speed mode,并保持 offload、fallback、cache -和 attention 设置一致。某个配置 OOM 就应记录为该配置失败,不能用 mock 或 offload 结果替代。 - -本文两条启动命令默认都使用 1 张 GPU;比较其他卡数时必须同时显式覆盖两个 target。 - -## 配置清单 - | 配置 | 用途 | |---|---| -| `video_generation_quick.yaml` | Batch 连通性与延迟 smoke test | -| `video_generation_e2e.yaml` | Batch warmup、trace、records 和服务指标 | +| `video_generation_quick.yaml` | 连通性和时延 smoke test | +| `video_generation_e2e.yaml` | Warmup、trace、records 和服务指标 | | `video_generation_rate.yaml` | Poisson 到达负载 | -| `video_generation_wan21_i2v_480p_compare.yaml` | 固定 Wan I2V 对比 workload | -| `stream_lingbot_world_fast_quick.json` | 有界 Stream smoke test | -| `stream_lingbot_world_fast_compare.json` | 固定 LingBot Stream 对比 workload | - -SGLang 对应配置位于 `benchmarks/baseline/sglang_lingbot_stream/configs`。 - -## 指标解释 - -必须区分指标 scope: - -| 指标 | 含义 | -|---|---| -| `stream_fps` | 客户端收到帧数除以客户端 session 时间 | -| `chunk_compute_fps` | 单个目标 chunk 的帧数除以计算时间 | -| `chunk_compute_fps_weighted` | 排除 warmup 后的 `sum(frames) / sum(compute_seconds)` | - -AIPerf 按交付、时延、吞吐、目标执行和资源五个稳定维度展示指标。不同实现的细分上报先保留为原始证据, -再映射到 canonical leaf,不扩张成新的顶层指标。 +| `video_generation_wan21_i2v_480p_compare.yaml` | 固定 Wan I2V 对比 | ## 主动资源上报与历史曲线 -使用 Docker 可以直接启动带持久化卷的 GreptimeDB: +启动持久化 GreptimeDB: ```bash docker volume create aiperf-greptime-data @@ -146,8 +77,7 @@ docker run -d --name aiperf-greptime --restart unless-stopped \ --data-home /greptimedb_data ``` -生产环境应固定镜像 tag 或 digest。命名卷会在容器重启后保留历史数据。随后在 TeleFuser 仓库根目录 -启动内置的 AIPerf 后端和中文前端;`artifacts` 与 benchmark launcher 的输出目录一致: +再启动 AIPerf history API 与内置 dashboard: ```bash uv run --frozen --no-dev --project benchmarks/aiperf aiperf history serve \ @@ -158,51 +88,13 @@ uv run --frozen --no-dev --project benchmarks/aiperf aiperf history serve \ --port 8095 ``` -检查前后端和数据库: - -```bash -curl --fail http://127.0.0.1:8095/api/v1/history/health -curl --fail -X POST 'http://127.0.0.1:4000/v1/sql?db=public' \ - --data-urlencode 'sql=SELECT 1 AS ready' -``` - -在 target 所在机器开启主动采集: - -```bash -export AIPERF_HISTORY_URL=http://:8095 -export AIPERF_RESOURCE_TARGET_PID= - -bash benchmarks/telefuser_aiperf/scripts/run_stream_bench.sh -``` - -Agent 递归观测目标进程树,默认每 1 秒采样、每 15 秒上报,并在任务结束时 flush。它采集 CPU、内存、 -GPU、显存、Ethernet 和 RDMA 的进程、可探测容器与整机事实。用量曲线与整机容量始终分开;CPU 以一个 -逻辑核为 100%,多核和多卡允许超过 100%。 - -GreptimeDB 是 History 与主动上报的强依赖。启动、查询或最终 flush 失败会直接暴露,不会切换到 SQLite、 -内存索引或文件直查。 - -打开 `http://127.0.0.1:8095/` 查看中文桌面界面。页面支持左右两组 Run、按 canonical 指标树选择图表、 -同时展示 avg/P95/P99、资源时间折线和跨实验对比。 - -远端实验机建议保持默认 loopback 监听,并通过 SSH 安全转发: - -```bash -ssh -L 8095:127.0.0.1:8095 user@benchmark-host -``` - -## 产物与复现要求 - -Batch 和 Stream 启动器默认将带时间戳的结果写入 `artifacts/`。Stream 产物包含 summary、session/event -JSONL、目标 metadata、normalized metrics 和独立 HTML 报告。 +如需采集目标进程树,在执行 batch launcher 前设置 `AIPERF_HISTORY_URL` 和 +`AIPERF_RESOURCE_TARGET_PID`。History 与主动上报强依赖 GreptimeDB;失败时不会静默回退到内存或文件数据库。 -正式性能结果至少应保留: +## 复现要求 -- TeleFuser 或 SGLang commit 与模型 revision; -- GPU 型号/数量、driver、CUDA、PyTorch 和 dtype; -- workload 配置与 control trace; -- warmup 规则及成功/失败 session 数; -- offload、cache、attention 和 fallback 设置。 +每个结果都应保留 TeleFuser/AIPerf commit、模型 revision、加速器型号/数量、driver、CUDA、PyTorch、dtype、 +完整 workload、warmup 规则、成功/失败数量,以及 offload/cache/attention 设置。动态结果保存在 GreptimeDB +和可重放产物中,不写入稳定文档。 -协议和职责边界见 [TeleFuser 与 AIPerf Benchmark 设计](benchmark_aiperf_design.md)。动态实验数值保存在 -GreptimeDB 和可重放产物中,不写入稳定用户文档。 +稳定职责与指标边界见 [AIPerf benchmark 设计](benchmark_aiperf_design.md)。 diff --git a/docs/zh/benchmark_aiperf_design.md b/docs/zh/benchmark_aiperf_design.md index 983196e..eb50995 100644 --- a/docs/zh/benchmark_aiperf_design.md +++ b/docs/zh/benchmark_aiperf_design.md @@ -1,214 +1,68 @@ # TeleFuser 与 AIPerf Benchmark 设计 -本文定义稳定的职责、协议和指标语义。具体实验数值、机器地址和运行状态不属于设计文档,应保存在 -GreptimeDB 与可重放产物中。 +本文定义稳定职责、协议和指标语义。具体实验数值、机器地址和运行状态应保存在 GreptimeDB 与可重放产物中。 -## 1. 目标与非目标 - -设计目标: - -- 同一 workload 可以比较 TeleFuser、SGLang-Diffusion 和后续实现; -- 客户端交付性能、目标侧计算性能和资源使用互不混淆; -- Batch 与 Stream 共用产物、历史查询和展示维度; -- target 只上报原始、有限、有时间戳的事实; -- AIPerf 统一负责采集生命周期、warmup、聚合、映射、存储和界面。 - -非目标: - -- 不在 TeleFuser 内复制 AIPerf、GreptimeDB client 或历史前端; -- 不用 mock、offload 或 fallback 结果替代正式 GPU-resident 结果; -- 不把实现私有字段全部提升为用户可选的顶层指标。 - -## 2. 仓库与依赖边界 +## 职责边界 ```mermaid flowchart LR - TF[TeleFuser target] -->|raw phase/chunk/runtime facts| AP[AIPerf] - SG[SGLang target] -->|native stream facts| AP + TF[TeleFuser target] -->|raw phase/runtime facts| AP[AIPerf] AP -->|canonical artifacts| GT[GreptimeDB] - AP --> UI[Vue history dashboard] + AP --> UI[History dashboard] AG[AIPerf resource agent] -->|timestamped batches| AP TF -. no database dependency .-> GT - SG -. no AIPerf package dependency .-> AP ``` | 组件 | 归属 | 职责 | |---|---|---| -| TeleFuser runtime | TeleFuser | 同步测量 target phase/chunk,暴露环境与 cache 原始事实 | -| Target adapter | AIPerf | 将 HTTP、WebRTC、WebSocket 等 wire event 转成统一 session timeline | -| 聚合与语义映射 | AIPerf | warmup、percentile、weighted FPS、canonical metric | +| TeleFuser runtime | TeleFuser | 同步测量 target phase,暴露环境与 cache 原始事实 | +| Target adapter | AIPerf | 将 `/v1/videos` HTTP 事件转成统一请求时间线 | +| 聚合与语义映射 | AIPerf | Warmup、percentile、throughput 和 canonical metric | | Resource agent | AIPerf | 采样目标进程树、cgroup、机器和设备资源并主动上报 | -| History API/UI | AIPerf | GreptimeDB schema、查询、左右 Run 对比和图表 | -| Contract/config/data | Target 仓库 | 固定 target 能力、workload 和可复现入口 | +| History API/UI | AIPerf | GreptimeDB schema、查询、跨 Run 对比和图表 | +| Contract/config/data | TeleFuser | 固定 target 能力、workload 和可复现入口 | -`benchmarks/aiperf/` 是被 Git 忽略的固定外部 checkout。`scripts/setup_aiperf_repo.sh` 与 benchmark launcher -不接受路径覆盖,AIPerf 必须位于 `/benchmarks/aiperf`。repo URL、branch 和 ref 可以调整,正式运行 -仍必须固定 commit,不能只记录可移动分支名。 +当前仓库只维护 batch video adapter 资产。流服务使用 LiveKit,但 AIPerf 尚无对应 adapter;在具备经过验证的 +LiveKit client adapter 之前,不用 mock 或已删除的直接 WebRTC adapter 代替真实 stream transport。 -## 3. 场景与实现 +## 原始事实协议 -| 场景 | TeleFuser | Baseline | 公平性边界 | -|---|---|---|---| -| Batch Video | OpenAI 兼容 `/v1/videos` | 兼容相同 contract 的外部 target | prompt、输入图、尺寸、帧数、steps、seed 一致 | -| Stream World | WebRTC media + DataChannel | SGLang WebSocket + MessagePack | prompt、首帧、FPS、session、control trace、GPU/offload 策略一致 | -| Transport Mock | WebRTC mock | WebSocket mock | 只比较 transport/harness,不解释模型性能 | +- Duration 使用单调时钟;跨进程或跨机器样本同时携带源端 UTC 时间戳。 +- CUDA phase 在开始和结束边界同步目标设备。 +- 数值必须有限且非负;不可用字段省略或为 `null`,不能伪造为零。 +- Memory 在线协议中使用 bytes,显示层再转换为 MB/GB。 +- Target 不排除 warmup、不计算 percentile、不生成跨 Run 结论。 -Transport 可以不同,但 adapter 输出的逻辑事件必须一致:连接、首帧、控制发送、控制确认、下一帧、chunk -事实和 session 结束。 - -## 4. Target 原始事实协议 - -### 4.1 通用规则 - -- duration 使用单调时钟;跨进程/跨机器样本同时携带源端 UTC 时间戳; -- CUDA phase 在开始和结束边界同步目标设备; -- 数值必须有限且非负,不可用字段省略或为 `null`,不能伪造为零; -- memory 使用 bytes 作为线协议单位;显示层再转换为 MB/GB; -- target 不排除 warmup、不计算 percentile、不生成跨 Run 结论。 - -### 4.2 Phase fact +Phase fact 示例: ```json { "name": "pipeline_init", "seconds": 12.3, - "memory": [ - { - "device": "cuda:0", - "peak_allocated_bytes": 123, - "peak_reserved_bytes": 456 - } - ] -} -``` - -TeleFuser Stream metadata 可以提供 `pipeline_init`;首次 LingBot chunk 之前提供 `runtime_creation`。 - -### 4.3 Chunk fact - -```json -{ - "index": 3, - "frames": 3, - "compute_seconds": 0.45, - "memory": [] + "memory": [{"device":"cuda:0","peak_allocated_bytes":123,"peak_reserved_bytes":456}] } ``` -TeleFuser 的 `compute_seconds` 从 chunk 提交 actor graph 前开始,覆盖 encode、denoise、decode 与目标内调度, -结束于原始帧返回。`encode_seconds` 是 AIPerf 支持的可选事实,仅在 target 能给出有界 payload 编码阶段时 -上报;当前 TeleFuser 原生 WebRTC 编码发生在 chunk fact 之后,因此省略该字段。客户端网络接收、播放 pacing -和 UI 渲染不进入 target compute。当前 LingBot 生成位于子 actor 中,服务进程的 CUDA allocator 统计不能覆盖 -完整 actor graph,因此 `memory` 保持为空;进程树显存时序由 AIPerf resource telemetry 采集,不能拿它替代 -reset-scoped allocator peak。 - -### 4.4 Runtime fact - -LingBot runtime 只上报稳定几何信息: - -- width、height、latent frames 和 frame tokens; -- chunk size 与 max attention size; -- local attention、sink 和 KV cache capacity。 - -软件环境至少包含 TeleFuser commit、Python、PyTorch、CUDA 以及可见 GPU 型号、compute capability 和显存容量。 +软件环境至少包含 TeleFuser commit、Python、PyTorch、CUDA,以及可见 GPU 型号、compute capability 和显存。 -## 5. AIPerf 聚合语义 +## 聚合语义 -### 5.1 Scope - -| Scope | 示例 | 聚合规则 | +| Scope | 示例 | 规则 | |---|---|---| -| Event | control ack、frame arrival | 保留单事件时间线 | -| Chunk | `chunk_compute_fps` | `frames / compute_seconds` | -| Session | `stream_fps`、first-frame latency | 每个 session 独立计算 | -| Run | `chunk_compute_fps_weighted` | warmup 后 `sum(frames) / sum(compute_seconds)` | - -`stream_fps`、`chunk_compute_fps` 和 `chunk_compute_fps_weighted` 不得合并。avg、P95、P99 是同一个 canonical -指标的统计曲线,不是三项独立指标。 - -### 5.2 五个核心维度 - -| 维度 | Canonical leaf 示例 | -|---|---| -| 交付 | success rate、frames received、stream FPS | -| 时延 | request、first frame、control ack、control-to-frame | -| 吞吐 | request throughput、weighted compute FPS | -| 目标执行 | pipeline/runtime phase、chunk compute、可选 encode、allocator peak | -| 资源 | CPU、内存、GPU、显存、网络 | - -TeleFuser 与 SGLang 的私有字段先保存为 raw point,再由版本化 mapping 映射到这些 leaf。无法等价的字段保持 -私有或 unavailable,不通过改名制造可比性。 - -## 6. 主动资源上报 - -Resource agent 与 target PID 同机运行: - -1. 注册 run 与 source identity; -2. 每 1 秒采样; -3. 每 15 秒有界批量上报; -4. 任务完成、失败或取消时立即 final flush; -5. 注册、批次或 final flush 未获确认时使启用资源采集的 benchmark 失败。 - -每个点包含 `run_id`、metric、subject、source timestamp、value、unit 和区分设备/网卡/cgroup 的 labels。 - -资源 subject: - -- `process_used`:目标 PID 及其后代; -- `container_used`:能可靠解析的 cgroup charged usage; -- `machine_used`:整机或物理设备使用; -- `machine_total`:整机/设备容量; -- `container_total`:有限容器上限,仅作为容量事实。 - -采集规则: - -- CPU 以一个逻辑核为 100%,多核进程和整机允许超过 100%; -- GPU/显存按物理设备保留 labels,同一 Run、同一 subject 内才允许堆叠; -- Ethernet 使用网卡 byte counter;RDMA 使用 active-port counter; -- 机器网络 counter 不能可靠归因到进程,因此不伪造 `process_used`; -- cgroup v1/v2 无法解析的上限保持 unavailable,不拿整机容量替代; -- 通用 cgroup 没有可移植网络带宽上限,因此不生成容器网络容量。 - -## 7. GreptimeDB 与前端 - -GreptimeDB 是 History 的唯一在线存储。服务启动或建表失败直接失败,查询失败返回 503;不存在 SQLite、 -内存索引或文件直查 fallback。JSON/JSONL 是可重放导入源,不是在线查询后端。 - -部署时由 TeleFuser setup 在固定 `benchmarks/aiperf` checkout 内创建独立的无 dev 依赖运行环境,并创建 -固定 `artifacts` 导入根目录;Vue 构建产物随 AIPerf 提供,API 与前端由同一 History 进程服务,不引入 -第二个 Node.js 运行进程。GreptimeDB 使用独立持久化卷,开发机默认只监听 loopback;远端查看通过 SSH -tunnel 或带认证的反向代理完成。 - -前端使用与指标树相同的五维顺序: - -- 左侧固定树只列 canonical leaf; -- TeleFuser/SGLang 和 avg/P95/P99 作为图中曲线,不拆成树节点; -- 右侧支持左右两组 Run,并按固定维度和 leaf 顺序排列卡片; -- 勾选只控制整项指标的显示/隐藏,不改变卡片顺序; -- 百分比用量允许堆叠,容量使用独立单线小图; -- 连续折线节点只在悬浮时显示;tooltip 位于浏览器 top layer,不能被卡片裁剪; -- 内存、显存和带宽按量级显示 KB/MB/GB/TB 与 KB/s/MB/s/GB/s/TB/s。 - -## 8. 产物与失败语义 - -每个 Run 至少保留: +| Event | response arrival | 保留单事件时间线 | +| Request | first output、request latency | 每个请求独立计算 | +| Run | success rate、throughput、percentile | 排除 warmup 后聚合 | -- resolved config、contract 和 AIPerf commit; -- summary、session/event JSONL 和 normalized points; -- target metadata 与资源 source identity; -- 成功、失败、超时、OOM 和取消数量; -- 独立 HTML 报告。 +AIPerf 按交付、时延、吞吐、目标执行和资源五个稳定维度展示。实现私有字段先保留为 raw point,再由版本化 +mapping 映射;无法等价的字段保持 private 或 unavailable。 -正式对比只使用 workload 与资源策略一致的成功 Run。OOM、超时和连接失败本身也是结果,不从分母删除; -mock、native fallback 或 offload Run 必须使用不同 qualification,不能伪装成正式配置。 +## 资源与历史 -## 9. 验证要求 +Resource agent 递归观测目标进程树,时间戳在采样端产生。CPU、内存、GPU、显存和网络用量与整机容量分开; +网络接口分类和容器 attribution 必须保留可验证证据。GreptimeDB 是主动上报和 History 查询的唯一数据库 +边界,失败必须显式暴露。 -提交前至少覆盖: +## 可复现性 -- runtime timer、设备去重、finite value 与 allocator fact 单元测试; -- TeleFuser phase/chunk event 及 API 指标透传; -- TeleFuser WebRTC 与 SGLang WebSocket mock contract; -- warmup 排除、weighted FPS 和 canonical mapping; -- resource schema、cgroup、Ethernet/RDMA 与 final flush; -- GreptimeDB 强依赖、幂等导入、API 查询和 Vue production build; -- 一次固定配置的真实 TeleFuser Run;SGLang 公平配置无法完成时保留明确失败证据。 +正式比较必须固定 TeleFuser、AIPerf、模型和数据 revision,记录完整软硬件环境、workload、warmup、并发、 +offload/cache/attention 设置及失败请求。OOM 是该配置的结果,不能用 mock 或 offload 结果替代。 diff --git a/scripts/setup_aiperf_repo.sh b/scripts/setup_aiperf_repo.sh index a7a5f4c..785362b 100755 --- a/scripts/setup_aiperf_repo.sh +++ b/scripts/setup_aiperf_repo.sh @@ -89,7 +89,7 @@ if [[ -n "${AIPERF_REF}" ]]; then git -C "${AIPERF_DIR}" checkout "${AIPERF_REF}" fi -"${UV_BIN}" sync --no-dev --project "${AIPERF_DIR}" --extra streaming-webrtc +"${UV_BIN}" sync --no-dev --project "${AIPERF_DIR}" mkdir -p "${ROOT_DIR}/artifacts" echo "AIPerf ready: ${AIPERF_DIR}" From 7a3702ba5db16388338c1cdbee1698880cad8ae3 Mon Sep 17 00:00:00 2001 From: lzx1413 Date: Mon, 27 Jul 2026 10:48:44 +0000 Subject: [PATCH 07/11] docs(streaming): document the LiveKit-only workflow Rewrite the English and Chinese stream guides around LiveKit sessions, data topics, deployment, and lifecycle. Add a root README quickstart that launches coturn, LiveKit, TeleFuser, and the browser controller with VS Code port forwarding, health checks, troubleshooting, and shutdown order. Verification: - .venv/bin/mkdocs build --strict --site-dir /tmp/telefuser-docs-livekit-rewrite - no aiortc, old SDP route, or legacy Demo references remain - git diff --cached --check --- CLAUDE.md | 11 +- README.md | 100 +-- README_zh.md | 99 +-- docs/en/index.md | 17 +- docs/en/service.md | 28 +- docs/en/stream_scheduler.md | 3 +- docs/en/stream_server.md | 1084 +++++------------------------- docs/zh/index.md | 17 +- docs/zh/service.md | 27 +- docs/zh/stream_scheduler.md | 3 +- docs/zh/stream_server.md | 1074 +++++------------------------ examples/README.md | 5 +- examples/lingbot/README.md | 171 ++--- examples/stream_server/README.md | 47 ++ 14 files changed, 610 insertions(+), 2076 deletions(-) create mode 100644 examples/stream_server/README.md diff --git a/CLAUDE.md b/CLAUDE.md index 3700178..2fee9f9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -16,6 +16,7 @@ pre-commit run --all-files # Linting checks pytest tests/ # Run tests bash scripts/run_ci_tests.sh # Full CI suite telefuser serve /path/to/pipeline.py --port 8000 # Start API server +telefuser stream-serve /path/to/pipeline.py --port 8088 # LiveKit-backed streaming ``` ## Troubleshooting @@ -65,7 +66,7 @@ telefuser/ ├── orchestrator/ # Request orchestration and actor-based streaming scheduler ├── worker/ # Distributed worker management ├── entrypoints/ # CLI entry points -├── service/ # FastAPI service +├── service/ # FastAPI request-response and LiveKit-backed streaming └── client/ # Python SDK ``` @@ -94,6 +95,14 @@ telefuser/ and memory without retaining duration-sized tensor lists. - Interpret chunk period as output cadence: real-time operation requires p95 to stay below the media duration represented by one chunk, with margin for transport and encoding. +- Treat `telefuser stream-serve` as the only streaming entrypoint. It must accept both + `ServerPushService` and `BidirectionalService` while preserving native frame/audio payloads. +- LiveKit workers own one room and one pipeline session. Browser reconnects must not mutate + pipeline caches, and controller messages must remain on the reliable `tf.control` topic; + status uses reliable `tf.status`, while bounded telemetry uses unreliable `tf.metrics`. +- Keep transport metrics distinct from model facts: output cadence measures adjacent emitted + chunks, pipeline residence measures actor admission to output, and applied-control latency + is bound to the control snapshot consumed by that output chunk. ### LingBot-Video Single-Process Runtime diff --git a/README.md b/README.md index 596f267..c6d3e35 100644 --- a/README.md +++ b/README.md @@ -17,6 +17,8 @@ TeleFuser is a high-performance runtime for world model inference and multimodal ## News 📰 +- ✨ **2026-07-27**: Unified streaming on LiveKit with room sessions, worker admission, reconnect-friendly browser + transport, and support for both server-push and bidirectional pipeline contracts. - ✨ **2026-07-22**: **NEW** Added [**LingBot-Video**](examples/lingbot_video/README.md) support for Dense and MoE T2I/T2V/TI2V generation, native four-GPU CFG/SP execution, and in-memory MoE refinement. - ✨ **2026-07-15**: Added [**LingBot-World v2**](https://github.com/Robbyant/lingbot-world-v2) support for offline generation, interactive WebRTC streaming, and multi-GPU inference. @@ -39,10 +41,12 @@ The project treats a world model as more than a function that returns a single c - **World-model-oriented runtime**: Support for continuous video generation, interactive sessions, and bidirectional control loops. - **ADF (AI Dev First)**: Repository layers, pipeline contracts, examples, and docs are structured for coding agents to discover capabilities, follow project conventions, and extend pipelines efficiently. - **Streaming pipeline scheduler**: Actor-owned stateful stages, bounded artifact edges, per-session ordering, backpressure, lifecycle cleanup, and explicit resource groups. -- **Streaming transport**: WebRTC-based streaming with media tracks plus DataChannel control for real-time inference. +- **Streaming transport**: LiveKit-backed WebRTC for server-push media and resilient bidirectional sessions, with + room lifecycle, reconnect handling, participant roles, and reliable controls. - **Scalable GPU runtime**: Multi-GPU execution with tensor parallelism, sequence parallelism, optional Ray workers, and distributed service replicas. - **Inference optimization stack**: Triton kernels, optimized attention backends, quantization, offload, feature caching, and CacheSeek latent cache integration. -- **Unified serving**: Local Python API, `telefuser serve` for task APIs, and `telefuser stream-serve` for continuous streaming services. +- **Unified serving**: Local Python API, `telefuser serve` for task APIs, and `telefuser stream-serve` for LiveKit + rooms and media. ## Quick Start @@ -62,7 +66,8 @@ TeleFuser does not require `tf-kernel` to run. No prebuilt tf-kernel package is extension is built from source with the Makefile under `tf-kernel/`. See the [tf-kernel README](tf-kernel/README.md) and [installation and usage guide](docs/en/tf_kernel.md) for build, verification, and compatibility details. -WebRTC streaming support is included in the default installation through `aiortc`. +The base installation includes the LiveKit Python SDKs used by `telefuser stream-serve`. A LiveKit Cloud project or +self-hosted LiveKit Server is operated separately. ### 1. Batch Video Inference @@ -84,65 +89,76 @@ video = pipe( ) ``` -### 2. Real-Time World Model Demo +### 2. Real-Time World Model WebRTC Demo -TeleFuser includes a bidirectional WebRTC demo for `LingBot-World v2`. -LingBot-World v2 uses camera control and its v2 PPL defaults; its streaming example caps a session at two minutes. +TeleFuser streams `LingBot-World v2` through LiveKit. LingBot-World v2 uses camera control and its v2 PPL defaults; +its streaming example caps a session at two minutes. LingBot streaming uses the actor-based scheduler for both offline and service execution. Encode, DiT, and decode may overlap even on the same GPU; move stages only when memory placement requires it. See the [streaming scheduler guide](docs/en/stream_scheduler.md). -For a laptop browser connected through VS Code Remote SSH, coturn is the only additional system package required; -no extra Python package is needed. On Debian or Ubuntu, install it with: +The checked-in browser page forces a TCP TURN relay so the same setup works through VS Code Remote SSH. The complete +local development stack therefore has four processes: coturn, LiveKit Server, TeleFuser, and the browser page. +Install the LiveKit Server and your platform's `coturn` package once: ```bash +# Debian/Ubuntu; use the equivalent coturn package on other platforms. sudo apt-get update sudo apt-get install -y coturn + +# Install LiveKit Server once. +curl -sSL https://get.livekit.io | bash +``` + +Then run each command below in a separate terminal from the repository root. + +Terminal 1 — start the development-only TCP TURN relay: + +```bash +turnserver -n -m 1 \ + --listening-ip=127.0.0.1 --relay-ip=127.0.0.1 \ + --listening-port=3478 --min-port=49160 --max-port=49200 \ + --user=livekit-demo:livekit-demo-password \ + --realm=livekit.local --fingerprint --lt-cred-mech \ + --no-tls --no-dtls --no-cli --allow-loopback-peers +``` + +Terminal 2 — start LiveKit with its development credentials (`devkey` / `secret`): + +```bash +livekit-server --dev ``` -The package provides both `turnserver` and the `turnutils_uclient` verification tool. Skip this step when both -commands already exist, or when the browser and GPU service run on the same physical machine. +Terminal 3 — load the four-GPU LingBot-World v2 service: ```bash TF_MODEL_ZOO_PATH=/path/to/model_zoo \ CUDA_VISIBLE_DEVICES=0,1,2,3 \ -TELEFUSER_TURN_SERVER='turn:127.0.0.1:3478?transport=tcp' \ -TELEFUSER_TURN_USERNAME=telefuser \ -TELEFUSER_TURN_CREDENTIAL=telefuser-turn \ telefuser stream-serve examples/lingbot/lingbot_world_v2_image_to_video_h100.py \ - --gpu-num 4 -p 8088 --host 0.0.0.0 --skip-validation - -python examples/stream_server/webrtc_bidirectional_demo.py \ - --server-url http://127.0.0.1:8088 \ - --port 8091 \ - --image-path examples/data/lingbot_world_fast/image.jpg \ - --turn-url 'turn:localhost:3478?transport=tcp' \ - --turn-username telefuser --turn-credential telefuser-turn \ - --force-turn-relay --ice-gather-timeout-ms 30000 --no-open + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey --livekit-api-secret secret \ + --num-workers 1 --worker-gpu-map 0,1,2,3 \ + --port 8088 --skip-validation ``` -This starts a continuous session where the client sends control messages over a WebRTC DataChannel and receives -generated video frames over media tracks. When the browser runs on a laptop through VS Code Remote SSH, configure -TURN over TCP and forward ports `8091` and `3478`; port `8088` does not need forwarding because the demo proxies -signaling requests. Keep local port `3478` equal to remote port `3478`; the forwarded 8091 port may use any available -local port. Without VS Code, run the equivalent tunnel from a terminal on the laptop: +Terminal 4 — serve the browser controller and proxy its session API: ```bash -ssh -N -o ExitOnForwardFailure=yes -o ServerAliveInterval=30 \ - -L 8091:127.0.0.1:8091 \ - -L 3478:127.0.0.1:3478 \ - USER@SERVER_HOST +python examples/stream_server/livekit_bidirectional_demo.py \ + --server-url http://127.0.0.1:8088 --port 8092 --no-open ``` -Then open `http://localhost:8091`. The TURN command and credentials above are development examples. See the -[stream server guide](docs/en/stream_server.md) and the -[LingBot example README](examples/lingbot/README.md) for coturn startup and the tested four-H100 setup. +For VS Code Remote SSH, forward remote TCP ports `8092`, `7880`, and `3478` to the same local ports; `8088` does not +need forwarding because the page proxies the TeleFuser API. Open `http://127.0.0.1:8092`, select an initial image, +click **Start**, and use the on-page controls or `W/A/S/D` and arrow keys. A successful connection shows a video +track plus `control_state`, generation-stage, and chunk status messages. -If the browser runs on the same physical machine as TeleFuser, no SSH tunnel or TURN server is needed. Unset all -`TELEFUSER_TURN_*` variables, start the service on `127.0.0.1:8088`, run the demo without any `--turn-*` or -`--force-turn-relay` arguments, and open `http://localhost:8091`. This does not apply when only the shell is on the -server through SSH but the browser still runs on a laptop. +Check the server independently with `curl http://127.0.0.1:8088/v1/service/health`. To stop the stack, stop the +browser session or close the page first, then press Ctrl+C in terminals 4, 3, 2, and 1. These loopback addresses, +static credentials, disabled TURN TLS, and `--allow-loopback-peers` are for trusted development only. See the +[stream server guide](docs/en/stream_server.md) for LiveKit Cloud, production networking, session APIs, and +troubleshooting. ### 3. Batch Service Mode @@ -162,7 +178,7 @@ See [docs/en/service.md](docs/en/service.md) for full API details. TeleFuser uses a layered runtime architecture that maps cleanly to the repository structure: -1. **Access layer**: FastAPI task APIs and WebRTC streaming entrypoints. +1. **Access layer**: FastAPI task APIs and LiveKit-backed stream room/session entrypoints. 2. **Service layer**: request routing, task management, stream sessions, replica pools, and integration with pipeline execution. 3. **Pipeline abstraction layer**: model-specific `BasePipeline` / `BaseStage` components, with an actor-based streaming orchestrator for bounded dataflow, session ordering, metrics, and cleanup. 4. **Model and optimization layer**: model loading, attention selection, quantization, offload, LoRA, and cache integration. @@ -172,7 +188,7 @@ Relevant directories: ```text telefuser/ -├── service/ # REST APIs, streaming APIs, WebRTC integration +├── service/ # REST APIs and LiveKit-backed streaming ├── orchestrator/ # Request orchestration and actor-based streaming scheduler ├── pipelines/ # Model-specific pipelines ├── distributed/ # TP / SP / FSDP / Ray utilities @@ -188,7 +204,7 @@ telefuser/ | Pipeline | Task | Notes | |----------|------|-------| -| `LingBot-World v2` | Bidirectional world-model streaming | Interactive WebRTC control loop via [examples/lingbot/lingbot_world_v2_image_to_video_h100.py](examples/lingbot/lingbot_world_v2_image_to_video_h100.py) | +| `LingBot-World v2` | Bidirectional world-model streaming | LiveKit control loop via [examples/lingbot/lingbot_world_v2_image_to_video_h100.py](examples/lingbot/lingbot_world_v2_image_to_video_h100.py) | | `LiveAct` | S2V | Speech-driven talking head generation via [examples/liveact/liveact_s2v_h100.py](examples/liveact/liveact_s2v_h100.py) | | `FlashVSR` | VSR | Streaming video super-resolution via [examples/flashvsr/README.md](examples/flashvsr/README.md) | @@ -215,7 +231,7 @@ See [examples/README.md](examples/README.md) for the example runner and baseline ## Documentation - [docs/en/service.md](docs/en/service.md): REST serving, task APIs, OpenAI-compatible APIs -- [docs/en/stream_server.md](docs/en/stream_server.md): continuous streaming and WebRTC protocols +- [docs/en/stream_server.md](docs/en/stream_server.md): LiveKit streaming, session APIs, data topics, and deployment - [docs/en/stream_scheduler.md](docs/en/stream_scheduler.md): actor-based stage scheduling, backpressure, lifecycle, metrics, and LingBot placement - [docs/en/parallel.md](docs/en/parallel.md): distributed inference architecture - [docs/en/latent_cache.md](docs/en/latent_cache.md): CacheSeek latent cache integration diff --git a/README_zh.md b/README_zh.md index 9f2fd1d..2dfa87a 100644 --- a/README_zh.md +++ b/README_zh.md @@ -17,6 +17,8 @@ TeleFuser 是一个面向世界模型推理与多模态生成的高性能运行 ## News 📰 +- ✨ **2026-07-27**:统一使用 LiveKit 流式后端,支持 room 会话、worker 准入、浏览器自动重连,以及 + server-push 和 bidirectional 两种 pipeline contract。 - ✨ **2026-07-22**:**NEW** 新增 [**LingBot-Video**](examples/lingbot_video/README.md) 支持,覆盖 Dense/MoE T2I、T2V、TI2V、原生四卡 CFG/SP 推理与内存直传 MoE refiner。 - ✨ **2026-07-15**:新增 [**LingBot-World v2**](https://github.com/Robbyant/lingbot-world-v2) 支持,支持离线生成、交互式 WebRTC 流和多卡推理。 @@ -39,10 +41,12 @@ TeleFuser 是一个面向世界模型推理与多模态生成的高性能运行 - **面向世界模型的运行时**:支持连续视频生成、交互式会话和双向控制闭环。 - **ADF (AI Dev First)**:通过清晰的仓库分层、Pipeline Contract、示例和文档约束,让 AI Agent 能理解能力边界、遵循项目开发流程,并高效扩展 Pipeline。 - **流式 Pipeline 调度器**:基于 actor 管理有状态 Stage,提供有界 artifact edge、session 顺序、backpressure、生命周期清理和显式 resource group。 -- **流式传输能力**:基于 WebRTC 的媒体流传输,并结合 DataChannel 实现实时控制。 +- **流式传输能力**:LiveKit-backed WebRTC 同时支持 server-push 媒体和稳定的双向会话,提供 room 生命周期、 + 重连、参与者角色和可靠控制消息。 - **可扩展 GPU 运行时**:支持多 GPU、张量并行、序列并行、Ray 部署和分布式工作节点编排。 - **推理优化栈**:包含 Triton Kernel、优化注意力后端、量化、卸载、特征缓存和 CacheSeek latent cache 集成。 -- **统一服务方式**:既支持本地 Python 调用,也支持 `telefuser serve` 和 `telefuser stream-serve` 两种服务模式。 +- **统一服务方式**:支持本地 Python 调用、任务 API `telefuser serve`,以及基于 LiveKit room/media 的 + `telefuser stream-serve`。 ## 快速开始 @@ -62,7 +66,8 @@ TeleFuser 不依赖 `tf-kernel` 也能运行。目前没有发布 tf-kernel 预 `tf-kernel/` 下的 Makefile 从源码构建。编译方法见 [tf-kernel README](tf-kernel/README_zh.md),安装验证、 支持配置和常见问题见 [tf-kernel 安装与使用指南](docs/zh/tf_kernel.md)。 -默认安装已通过 `aiortc` 包含 WebRTC 流式服务能力。 +基础安装已包含 `telefuser stream-serve` 使用的 LiveKit Python SDK;LiveKit Cloud 项目或自托管 LiveKit +Server 需要单独运行。 ### 1. 批量视频推理 @@ -84,64 +89,74 @@ video = pipe( ) ``` -### 2. 实时世界模型 Demo +### 2. 实时世界模型 WebRTC Demo -TeleFuser 当前提供了 `LingBot-World v2` 的双向 WebRTC Demo。 -LingBot-World v2 使用相机控制和 v2 PPL 默认值;其流式示例将单个会话上限设为两分钟。 +TeleFuser 通过 LiveKit 传输 `LingBot-World v2`。LingBot-World v2 使用相机控制和 v2 PPL 默认值;其流式 +示例将单个会话上限设为两分钟。 LingBot 的离线与服务执行共用 actor scheduler。即使位于同一张 GPU,encode、DiT 和 decode 也可以重叠; 仅在显存放置需要时移动 Stage。详见[流式调度器指南](docs/zh/stream_scheduler.md)。 -通过 VS Code Remote SSH 从笔记本浏览器访问时,coturn 是唯一需要额外安装的系统软件,不需要增加 -Python 包。在 Debian 或 Ubuntu 上执行: +仓库内浏览器页面强制使用 TCP TURN relay,以便同一套配置可通过 VS Code Remote SSH 工作。因此完整的本地 +开发环境包含四个进程:coturn、LiveKit Server、TeleFuser 和浏览器页面。先安装一次 LiveKit Server,并 +通过操作系统的包管理器安装 `coturn`: ```bash +# Debian/Ubuntu;其他平台请安装对应的 coturn 软件包。 sudo apt-get update sudo apt-get install -y coturn + +# LiveKit Server 只需安装一次。 +curl -sSL https://get.livekit.io | bash +``` + +然后从仓库根目录在四个独立终端中依次运行以下命令。 + +终端 1——启动仅供开发使用的 TCP TURN relay: + +```bash +turnserver -n -m 1 \ + --listening-ip=127.0.0.1 --relay-ip=127.0.0.1 \ + --listening-port=3478 --min-port=49160 --max-port=49200 \ + --user=livekit-demo:livekit-demo-password \ + --realm=livekit.local --fingerprint --lt-cred-mech \ + --no-tls --no-dtls --no-cli --allow-loopback-peers +``` + +终端 2——使用开发凭据(`devkey` / `secret`)启动 LiveKit: + +```bash +livekit-server --dev ``` -该软件包同时提供 `turnserver` 和用于验证的 `turnutils_uclient`。如果这两个命令已经存在,或者浏览器 -和 GPU 服务运行在同一台物理机器上,则可以跳过安装。 +终端 3——加载四卡 LingBot-World v2 服务: ```bash TF_MODEL_ZOO_PATH=/path/to/model_zoo \ CUDA_VISIBLE_DEVICES=0,1,2,3 \ -TELEFUSER_TURN_SERVER='turn:127.0.0.1:3478?transport=tcp' \ -TELEFUSER_TURN_USERNAME=telefuser \ -TELEFUSER_TURN_CREDENTIAL=telefuser-turn \ telefuser stream-serve examples/lingbot/lingbot_world_v2_image_to_video_h100.py \ - --gpu-num 4 -p 8088 --host 0.0.0.0 --skip-validation - -python examples/stream_server/webrtc_bidirectional_demo.py \ - --server-url http://127.0.0.1:8088 \ - --port 8091 \ - --image-path examples/data/lingbot_world_fast/image.jpg \ - --turn-url 'turn:localhost:3478?transport=tcp' \ - --turn-username telefuser --turn-credential telefuser-turn \ - --force-turn-relay --ice-gather-timeout-ms 30000 --no-open + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey --livekit-api-secret secret \ + --num-workers 1 --worker-gpu-map 0,1,2,3 \ + --port 8088 --skip-validation ``` -该流程会启动一个持续运行的会话:客户端通过 WebRTC DataChannel 发送控制消息,服务端通过媒体轨道 -持续回传生成视频。当浏览器运行在笔记本上,并通过 VS Code Remote SSH 访问远端服务器时,需要配置 -TCP TURN,并转发 `8091` 和 `3478` 端口。由于 demo 会代理信令请求,因此不需要转发 `8088`。本地 -`3478` 应保持映射到远端 `3478`;8091 可以映射到任意可用的本地端口。不使用 VS Code 时,可以在 -笔记本终端中建立等效的 OpenSSH 隧道: +终端 4——启动浏览器控制页面及其 session API 代理: ```bash -ssh -N -o ExitOnForwardFailure=yes -o ServerAliveInterval=30 \ - -L 8091:127.0.0.1:8091 \ - -L 3478:127.0.0.1:3478 \ - USER@SERVER_HOST +python examples/stream_server/livekit_bidirectional_demo.py \ + --server-url http://127.0.0.1:8088 --port 8092 --no-open ``` -然后打开 `http://localhost:8091`。上面的 TURN 账号密码仅作为开发配置示例。coturn 启动方式、生产环境 -注意事项及四张 H100 的实测配置见 [流服务文档](docs/zh/stream_server.md) -和 [LingBot example README](examples/lingbot/README.md)。 +使用 VS Code Remote SSH 时,把远端 TCP `8092`、`7880` 和 `3478` 映射到相同本地端口;页面会代理 +TeleFuser API,因此无需映射 `8088`。打开 `http://127.0.0.1:8092`,选择初始图片,点击 **Start**, +再使用页面按钮或 `W/A/S/D` 和方向键控制相机。成功连接后会显示视频轨道以及 `control_state`、生成 Stage +和 chunk 状态消息。 -如果浏览器和 TeleFuser 服务运行在同一台物理机器上,则不需要 SSH 隧道或 TURN 服务。清除所有 -`TELEFUSER_TURN_*` 环境变量,让服务监听 `127.0.0.1:8088`,启动 demo 时不要传入任何 `--turn-*` -或 `--force-turn-relay` 参数,然后打开 `http://localhost:8091`。如果只是通过 SSH 登录服务器、浏览器 -仍然运行在笔记本上,则不属于本机访问,仍需使用上述端口转发和 TURN 配置。 +可用 `curl http://127.0.0.1:8088/v1/service/health` 独立检查服务。停止时先结束浏览器 session 或关闭页面, +再按终端 4、3、2、1 的顺序按 Ctrl+C。Loopback 地址、静态凭据、禁用 TURN TLS 和 +`--allow-loopback-peers` 仅适用于可信开发环境。LiveKit Cloud、生产网络、session API 和故障排查见 +[流服务文档](docs/zh/stream_server.md)。 ### 3. 批处理服务模式 @@ -161,7 +176,7 @@ TeleFuser 对外提供: TeleFuser 采用分层运行时架构,并与仓库目录结构保持一致: -1. **接入层**:FastAPI 任务接口与 WebRTC 流式入口。 +1. **接入层**:FastAPI 任务接口和 LiveKit-backed stream room/session 入口。 2. **服务层**:请求路由、任务管理、流式会话、副本池,以及与 Pipeline 执行过程的集成。 3. **Pipeline 抽象层**:模型相关的 `BasePipeline` / `BaseStage` 组件;actor-based streaming orchestrator 提供有界数据流、session 顺序、指标和清理。 4. **模型与优化层**:模型加载、注意力选择、量化、offload、LoRA、cache 集成。 @@ -171,7 +186,7 @@ TeleFuser 采用分层运行时架构,并与仓库目录结构保持一致: ```text telefuser/ -├── service/ # REST API、流式 API、WebRTC 集成 +├── service/ # REST API 和 LiveKit-backed 流服务 ├── orchestrator/ # 请求编排与基于 actor 的流式调度 ├── pipelines/ # 模型级 Pipeline 实现 ├── distributed/ # TP / SP / FSDP / Ray 等并行能力 @@ -187,7 +202,7 @@ telefuser/ | Pipeline | 任务 | 说明 | |----------|------|------| -| `LingBot-World v2` | 双向世界模型流式推理 | 交互式 WebRTC 控制闭环,见 [examples/lingbot/lingbot_world_v2_image_to_video_h100.py](examples/lingbot/lingbot_world_v2_image_to_video_h100.py) | +| `LingBot-World v2` | 双向世界模型流式推理 | LiveKit 控制闭环,见 [examples/lingbot/lingbot_world_v2_image_to_video_h100.py](examples/lingbot/lingbot_world_v2_image_to_video_h100.py) | | `LiveAct` | S2V | 语音驱动数字人视频生成,见 [examples/liveact/liveact_s2v_h100.py](examples/liveact/liveact_s2v_h100.py) | | `FlashVSR` | VSR | 流式视频超分,见 [examples/flashvsr/README.md](examples/flashvsr/README.md) | @@ -214,7 +229,7 @@ telefuser/ ## 文档 - [docs/zh/service.md](docs/zh/service.md):REST 服务、任务 API、OpenAI 兼容接口 -- [docs/zh/stream_server.md](docs/zh/stream_server.md):连续流式推理与 WebRTC 协议 +- [docs/zh/stream_server.md](docs/zh/stream_server.md):LiveKit 流服务、session API、data topic 和部署 - [docs/zh/stream_scheduler.md](docs/zh/stream_scheduler.md):基于 actor 的 Stage 调度、backpressure、生命周期、指标和 LingBot 卡位 - [docs/zh/parallel.md](docs/zh/parallel.md):分布式推理架构 - [docs/zh/latent_cache.md](docs/zh/latent_cache.md):CacheSeek latent cache 集成 diff --git a/docs/en/index.md b/docs/en/index.md index db389e3..d2d6ef6 100644 --- a/docs/en/index.md +++ b/docs/en/index.md @@ -36,7 +36,7 @@ Compile-aware ops with eager CUDA Triton kernels and PyTorch native fallbacks.
**Streaming Service** -FastAPI batch serving plus WebRTC media tracks and DataChannel control. +FastAPI batch serving and LiveKit-backed rooms for server-push and resilient interactive WebRTC.
**Feature Cache** @@ -56,8 +56,8 @@ Reusable stages, model configs, schedulers, and pipeline orchestration. | Model | Tasks | Description | |-------|-------|-------------| -| LingBot-World v2 | Bidirectional streaming | Camera-controlled interactive world model via WebRTC | -| LingBot-World-Fast | Bidirectional streaming | Legacy/causal-fast interactive world model via WebRTC DataChannel | +| LingBot-World v2 | Bidirectional streaming | Camera-controlled interactive world model via LiveKit | +| LingBot-World-Fast | Bidirectional streaming | Legacy/causal-fast model via LiveKit reliable data messages | ### Video Generation @@ -88,17 +88,20 @@ pip install telefuser # Batch serving telefuser serve /path/to/pipeline.py --port 8000 -# Stream serving (WebRTC support is included in the default install) -telefuser stream-serve examples/lingbot/lingbot_world_fast_image_to_video_h100.py -p 8088 +# LiveKit-backed streaming (Python SDK included in the base install) +telefuser stream-serve examples/lingbot/lingbot_world_fast_image_to_video_h100.py \ + --livekit-url ws://127.0.0.1:7880 \ + --livekit-api-key devkey --livekit-api-secret secret \ + -p 8088 ``` ## Documentation Sections