Update Athena runtimes and sync live extensions

This commit is contained in:
Mikei386
2026-09-15 08:01:40 +02:00
parent aaa96bbed5
commit 2b727a4d7f
14 changed files with 728 additions and 33 deletions
+3 -2
View File
@@ -4,7 +4,7 @@ ARG LTX_DESKTOP_VERSION=1.2.7
ARG LTX_DESKTOP_SHA512=d1d59027988a48490492feb42156665bbed511f187739a664ad326492fd8fc0ce43429537150eb6aea9a75a522f6353a40839f9b3ed0449fb720b7c13b091706
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
ca-certificates curl dbus-x11 ffmpeg libasound2t64 libatk-bridge2.0-0 \
build-essential ca-certificates curl dbus-x11 ffmpeg libasound2t64 libatk-bridge2.0-0 \
libatk1.0-0 libcups2 libdrm2 libgbm1 libgtk-3-0 libnss3 libx11-xcb1 \
libxcomposite1 libxdamage1 libxfixes3 libxkbcommon0 libxrandr2 \
novnc openbox procps python3-websockify x11vnc xvfb \
@@ -26,7 +26,8 @@ COPY entrypoint.sh /usr/local/bin/ltx-desktop-entrypoint
COPY tar-no-owner.sh /usr/local/bin/tar
RUN chmod 0755 /usr/local/bin/ltx-desktop-entrypoint /usr/local/bin/tar
ENV DISPLAY=:0 \
ENV CC=/usr/bin/gcc \
DISPLAY=:0 \
HOME=/data/home \
XDG_DATA_HOME=/data \
XDG_CONFIG_HOME=/data/config \
@@ -0,0 +1,354 @@
"""Local, catalog-aware prompt enhancement handler."""
from __future__ import annotations
import logging
import os
import random
import uuid
from threading import RLock
from typing import TYPE_CHECKING
from _routes._errors import HTTPError
from api_types import EnhancePromptRequest, EnhancePromptResponse, IcLoraCatalogItem, LoraCatalogItem
from handlers.base import StateHandlerBase
from handlers.generation_handler import GenerationHandler
from handlers.pipelines_handler import PipelinesHandler
from handlers.text_handler import TextHandler
from server_utils.media_validation import normalize_optional_path, validate_image_file
from services.gemini_text_client import resolve_gemini_model
from services.interfaces import PromptEnhancerPipeline
from services.lora_catalog import LoraCatalogProvider
from services.prompt_enhancement import (
build_audio_visual_caption_system_prompt,
build_conditioning_system_prompt,
build_ic_lora_enhancement_system_prompt,
build_image_edit_system_prompt,
build_image_generation_system_prompt,
build_keyframe_enhancement_system_prompt,
build_lora_enhancement_system_prompt,
build_template_fill_system_prompt,
enforce_trigger_placements,
fill_prompt_template,
parse_template_fill_response,
)
from services.prompt_enhancement.i2v_frames import KeyframeStill
from services.prompt_enhancer_pipeline.gemini_prompt_enhancer_pipeline import GeminiPromptEnhancerPipeline
from services.services_utils import get_device_type
from state.app_state_types import AppState
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from runtime_config.runtime_config import RuntimeConfig
# Deliberately independent of StateHandlerBase._resolve_seed(): that helper honors the app's
# reproducibility seed lock (and a fixed constant in dev mode), which would make every enhance
# call — including a redo — produce the exact same output. Enhancement is a quick, exploratory
# action where a fresh draw each call is the whole point.
_MAX_ENHANCE_SEED = 2147483647
class PromptEnhancementHandler(StateHandlerBase):
def __init__(
self,
state: AppState,
lock: RLock,
generation_handler: GenerationHandler,
pipelines_handler: PipelinesHandler,
text_handler: TextHandler,
lora_catalog_provider: LoraCatalogProvider,
prompt_enhancer_pipeline_class: type[PromptEnhancerPipeline],
gemini_pipeline: GeminiPromptEnhancerPipeline,
config: RuntimeConfig,
) -> None:
super().__init__(state, lock, config)
self._generation = generation_handler
self._pipelines = pipelines_handler
self._text_handler = text_handler
self._lora_catalog_provider = lora_catalog_provider
self._prompt_enhancer_pipeline_class = prompt_enhancer_pipeline_class
self._gemini_pipeline = gemini_pipeline
def _random_seed(self) -> int:
return random.randint(0, _MAX_ENHANCE_SEED)
def enhance(self, req: EnhancePromptRequest) -> EnhancePromptResponse:
# Enhance never occupies the GPU slot (see PipelinesHandler.
# evict_gpu_pipeline_for_prompt_enhancement) but still needs to mutually exclude with
# generation and with itself — an abandoned/orphaned enhance call (e.g. the tab reloaded
# mid-request) must not race a Generate click, a second Enhance click, or a generation
# that's still loading its pipeline (reserved_generation_start covers that window; a bare
# is_generation_running() check does not — see its own docstring). The "api" generation
# slot gives us the mutual exclusion for free: it's the same bookkeeping every other
# handler already does, and it doesn't require gpu_slot to be set.
with self._generation.reserved_generation_start():
gemma_root: str | None = None
if req.provider == "local":
gemma_root = self._text_handler.resolve_prompt_enhancer_root_if_downloaded()
if gemma_root is None:
raise HTTPError(409, "LOCAL_TEXT_ENCODER_NOT_AVAILABLE")
elif not self.state.app_settings.gemini_api_key:
raise HTTPError(400, "GEMINI_API_KEY_MISSING")
generation_id = uuid.uuid4().hex[:8]
self._generation.start_api_generation(generation_id)
try:
enhanced = self._resolve_and_enhance(req, gemma_root)
except HTTPError as e:
self._generation.fail_generation(e.detail)
raise
except Exception as e:
self._generation.fail_generation(str(e))
raise HTTPError(500, str(e)) from e
self._generation.complete_generation(enhanced)
return EnhancePromptResponse(enhancedPrompt=enhanced)
def _resolve_and_enhance(self, req: EnhancePromptRequest, gemma_root: str | None) -> str:
if req.mediaType == "image":
# No catalog LoRA concept for images (validated at the request level) — always an
# explicit, image-domain system prompt, never the video-oriented generic fallback.
system_prompt = (
build_image_edit_system_prompt() if req.imagePath is not None
else build_image_generation_system_prompt()
)
return self._run_free_rewrite(req, system_prompt, gemma_root)
if req.icLoraId is not None:
ic_lora = self._lora_catalog_provider.get_ic_lora(req.icLoraId)
if ic_lora is None:
raise HTTPError(404, "LORA_CATALOG_ID_NOT_FOUND")
return self._enhance_ic_lora(ic_lora, req, gemma_root)
if req.loraCatalogIds:
loras: list[LoraCatalogItem] = []
for catalog_id in req.loraCatalogIds:
lora = self._lora_catalog_provider.get_lora(catalog_id)
if lora is None:
raise HTTPError(404, "LORA_CATALOG_ID_NOT_FOUND")
loras.append(lora)
return self._enhance_loras(loras, req, gemma_root)
if req.conditioningType is not None:
system_prompt = build_conditioning_system_prompt(req.conditioningType)
return self._run_free_rewrite(req, system_prompt, gemma_root)
return self._run_free_rewrite(req, self._default_video_system_prompt(req), gemma_root)
def _default_video_system_prompt(self, req: EnhancePromptRequest) -> str | None:
if req.keyframes:
return self._keyframe_system_prompt()
return self._video_system_prompt(t2v=req.imagePath is None)
def _keyframe_system_prompt(self) -> str:
spec = self._text_handler.active_ltx_model_spec()
audio_visual = spec is not None and spec.wants_audio_visual_captions
return build_keyframe_enhancement_system_prompt(audio_visual=audio_visual)
def _video_system_prompt(self, *, t2v: bool) -> str | None:
"""The active model's own caption style, or None to keep each provider's default.
Only the audio-visual generations (2.5) need this: their captions cover the soundscape,
which neither the generic Gemini fallback nor a 2.3-era prompt asks for.
"""
spec = self._text_handler.active_ltx_model_spec()
if spec is None or not spec.wants_audio_visual_captions:
return None
return build_audio_visual_caption_system_prompt(t2v=t2v)
def enhance_for_generation(
self,
prompt: str,
*,
image_path: str | None,
last_image_path: str | None = None,
keyframes: list[KeyframeStill] | None = None,
duration: int | None = None,
fps: int | None = None,
) -> str:
"""Rewrite ``prompt`` on the local enhancer for a generation that's already started.
Only for the local text-encoding path: API encoding enhances server-side inside the
same call, so it never gets here. Deliberately not `enhance()` — the caller already
holds the generation slot, and there's no provider choice to make, only "is the
enhancer on disk".
Never raises: enhancement is a quality step, so a missing checkpoint or a failed
rewrite degrades to the prompt as typed rather than failing the generation.
"""
if not prompt.strip():
return prompt
gemma_root = self._text_handler.resolve_prompt_enhancer_root_if_downloaded()
if gemma_root is None:
logger.info("Skipping automatic enhancement: no local prompt enhancer downloaded")
return prompt
try:
pipeline = self._load_prompt_enhancer_pipeline(gemma_root)
system_prompt = (
self._keyframe_system_prompt()
if keyframes
else self._video_system_prompt(t2v=image_path is None)
)
seed = self._random_seed()
if image_path is not None or keyframes:
first_path = image_path or (keyframes[0][0] if keyframes else None)
assert first_path is not None
enhanced = pipeline.enhance_i2v(
prompt,
first_path,
system_prompt=system_prompt,
seed=seed,
last_image_path=None if keyframes else last_image_path,
keyframes=keyframes,
duration=duration,
fps=fps,
)
else:
enhanced = pipeline.enhance_t2v(prompt, system_prompt=system_prompt, seed=seed)
except Exception:
logger.warning("Automatic local enhancement failed; using the prompt as typed", exc_info=True)
return prompt
if not enhanced.strip():
return prompt
logger.info(
"Enhanced prompt locally for generation (%d -> %d chars): %s",
len(prompt),
len(enhanced),
enhanced,
)
return enhanced
def _enhance_loras(
self, loras: list[LoraCatalogItem], req: EnhancePromptRequest, gemma_root: str | None
) -> str:
# None of the plain LoRAs in the catalog have a prompt_template today (only IC-LoRAs
# do) — the multi-select path is always a free rewrite.
system_prompt = build_lora_enhancement_system_prompt(loras)
enhanced = self._run_free_rewrite(req, system_prompt, gemma_root)
return enforce_trigger_placements(enhanced, loras)
def _enhance_ic_lora(
self, ic_lora: IcLoraCatalogItem, req: EnhancePromptRequest, gemma_root: str | None
) -> str:
if ic_lora.prompt_template is not None:
return self._run_template_fill(ic_lora, req, gemma_root)
system_prompt = build_ic_lora_enhancement_system_prompt(ic_lora)
enhanced = self._run_free_rewrite(req, system_prompt, gemma_root)
return enforce_trigger_placements(enhanced, [ic_lora])
def _run_free_rewrite(
self, req: EnhancePromptRequest, system_prompt: str | None, gemma_root: str | None
) -> str:
# Reject an invalid/unreadable/oversized path before it reaches either provider — the
# API path in particular would otherwise base64-encode and ship arbitrary file bytes to
# a third-party API with no gate at all.
keyframes = self._validated_keyframes(req)
image_path = (
None if keyframes else normalize_optional_path(req.imagePath)
)
last_image_path = (
None
if req.mediaType == "image" or keyframes
else normalize_optional_path(req.lastImagePath)
)
if image_path is not None:
validate_image_file(image_path)
if last_image_path is not None:
validate_image_file(last_image_path)
seed = self._random_seed()
first_path = image_path or (keyframes[0][0] if keyframes else None)
if req.provider == "api":
resolved_model = resolve_gemini_model(self.state.app_settings.gemini_model)
logger.info("Enhancing prompt via Gemini API (%s)", resolved_model)
api_key = self.state.app_settings.gemini_api_key
if first_path is not None:
return self._gemini_pipeline.enhance_i2v(
req.prompt,
first_path,
system_prompt=system_prompt,
seed=seed,
api_key=api_key,
model=resolved_model,
last_image_path=last_image_path,
keyframes=keyframes,
duration=req.duration,
fps=req.fps,
)
return self._gemini_pipeline.enhance_t2v(
req.prompt,
system_prompt=system_prompt,
seed=seed,
api_key=api_key,
model=resolved_model,
)
logger.info("Enhancing prompt via local Gemma")
assert gemma_root is not None
pipeline = self._load_prompt_enhancer_pipeline(gemma_root)
if first_path is not None:
return pipeline.enhance_i2v(
req.prompt,
first_path,
system_prompt=system_prompt,
seed=seed,
last_image_path=last_image_path,
keyframes=keyframes,
duration=req.duration,
fps=req.fps,
)
return pipeline.enhance_t2v(req.prompt, system_prompt=system_prompt, seed=seed)
def _validated_keyframes(self, req: EnhancePromptRequest) -> list[KeyframeStill] | None:
if req.mediaType == "image" or not req.keyframes:
return None
frames: list[KeyframeStill] = []
for keyframe in req.keyframes:
path = normalize_optional_path(keyframe.imagePath)
if path is None:
raise HTTPError(400, "Each keyframe requires an image path")
validate_image_file(path)
frames.append((path, keyframe.frameIndex, keyframe.strength))
frames.sort(key=lambda item: item[1])
return frames
def _run_template_fill(
self, ic_lora: IcLoraCatalogItem, req: EnhancePromptRequest, gemma_root: str | None
) -> str:
# req.imagePath is intentionally unused here — template fill is always a text-only
# enhance_t2v call (the fixed template scaffold carries no reference-image slot), unlike
# the free-rewrite IC-LoRA path below it, which does route an image through enhance_i2v.
assert ic_lora.prompt_template is not None
system_prompt = build_template_fill_system_prompt(ic_lora)
seed = self._random_seed()
try:
if req.provider == "api":
resolved_model = resolve_gemini_model(self.state.app_settings.gemini_model)
logger.info("Enhancing prompt via Gemini API (%s)", resolved_model)
raw = self._gemini_pipeline.enhance_t2v(
req.prompt,
system_prompt=system_prompt,
seed=seed,
api_key=self.state.app_settings.gemini_api_key,
model=resolved_model,
)
else:
logger.info("Enhancing prompt via local Gemma")
assert gemma_root is not None
pipeline = self._load_prompt_enhancer_pipeline(gemma_root)
raw = pipeline.enhance_t2v(req.prompt, system_prompt=system_prompt, seed=seed)
values = parse_template_fill_response(raw, set(ic_lora.prompt_template.placeholders))
return fill_prompt_template(ic_lora.prompt_template, values)
except ValueError as e:
raise HTTPError(500, f"PROMPT_TEMPLATE_FILL_FAILED: {e}") from e
def _load_prompt_enhancer_pipeline(self, gemma_root: str) -> PromptEnhancerPipeline:
self._pipelines.evict_gpu_pipeline_for_prompt_enhancement()
device = os.getenv("LTX_PROMPT_ENHANCER_DEVICE", "").strip()
if not device:
device = get_device_type(self.config.device)
logger.info("Loading local prompt enhancer on %s", device)
return self._prompt_enhancer_pipeline_class.create(gemma_root, device)
+179
View File
@@ -0,0 +1,179 @@
"""FastAPI app factory decoupled from runtime bootstrap side effects."""
from __future__ import annotations
import base64
import hmac
import os
from pathlib import Path
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any
from fastapi import FastAPI, Request
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.responses import Response as StarletteResponse
from _routes._errors import HTTPError, build_http_error_response
from _routes.generation import router as generation_router
from _routes.hf_auth import router as hf_auth_router
from _routes.health import router as health_router
from _routes.ic_lora import router as ic_lora_router
from _routes.lora_catalog import router as lora_catalog_router
from _routes.image_gen import router as image_gen_router
from _routes.prompt_enhancement import router as prompt_enhancement_router
from _routes.models import router as models_router
from _routes.suggest_gap_prompt import router as suggest_gap_prompt_router
from _routes.retake import router as retake_router
from _routes.extend import router as extend_router
from _routes.runtime_policy import router as runtime_policy_router
from _routes.settings import router as settings_router
from api_types import HTTPErrorResponse
from logging_policy import log_http_error, log_unhandled_exception
from state import init_state_service
if TYPE_CHECKING:
from app_handler import AppHandler
DEFAULT_ALLOWED_ORIGINS: list[str] = [
"http://localhost:5173",
"http://127.0.0.1:5173",
]
DEFAULT_ERROR_RESPONSES: dict[int | str, dict[str, Any]] = {
"4XX": {
"model": HTTPErrorResponse,
"description": "Client Error",
},
"5XX": {
"model": HTTPErrorResponse,
"description": "Server Error",
},
}
def create_app(
*,
handler: "AppHandler",
allowed_origins: list[str] | None = None,
title: str = "LTX-2 Video Generation Server",
auth_token: str = "",
admin_token: str = "",
) -> FastAPI:
"""Create a configured FastAPI app bound to the provided handler."""
init_state_service(handler)
app = FastAPI(title=title, responses=DEFAULT_ERROR_RESPONSES)
remote_auth_token = os.environ.get("LTX_REMOTE_AUTH_TOKEN", "").strip()
remote_token_file = os.environ.get("LTX_REMOTE_AUTH_TOKEN_FILE", "").strip()
if remote_token_file:
remote_auth_token = Path(remote_token_file).read_text(encoding="utf-8").strip()
app.state.admin_token = admin_token # type: ignore[attr-defined]
app.add_middleware(
CORSMiddleware,
allow_origins=allowed_origins or DEFAULT_ALLOWED_ORIGINS,
allow_methods=["*"],
allow_headers=["*"],
)
@app.middleware("http")
async def _auth_middleware( # pyright: ignore[reportUnusedFunction]
request: Request,
call_next: Callable[[Request], Awaitable[StarletteResponse]],
) -> StarletteResponse:
if not auth_token:
return await call_next(request)
if request.method == "OPTIONS":
return await call_next(request)
if request.url.path == "/api/auth/huggingface/callback":
return await call_next(request)
def _token_matches(candidate: str) -> bool:
return hmac.compare_digest(candidate, auth_token) or (
bool(remote_auth_token) and hmac.compare_digest(candidate, remote_auth_token)
)
# WebSocket: check query param
if request.headers.get("upgrade", "").lower() == "websocket":
if _token_matches(request.query_params.get("token", "")):
return await call_next(request)
return JSONResponse(
status_code=401,
content=build_http_error_response(401, "Unauthorized").model_dump(),
)
# HTTP: Bearer or Basic auth
auth_header = request.headers.get("authorization", "")
if auth_header.startswith("Bearer ") and _token_matches(auth_header[7:]):
return await call_next(request)
if auth_header.startswith("Basic "):
try:
decoded = base64.b64decode(auth_header[6:]).decode()
_, _, password = decoded.partition(":")
if _token_matches(password):
return await call_next(request)
except Exception:
pass
return JSONResponse(
status_code=401,
content=build_http_error_response(401, "Unauthorized").model_dump(),
)
async def _route_http_error_handler(request: Request, exc: Exception) -> JSONResponse:
if isinstance(exc, HTTPError):
log_http_error(request, exc)
return JSONResponse(status_code=exc.status_code, content=exc.response.model_dump())
return JSONResponse(
status_code=500,
content=build_http_error_response(500, str(exc)).model_dump(),
)
async def _starlette_http_error_handler(request: Request, exc: Exception) -> JSONResponse:
if isinstance(exc, StarletteHTTPException):
return JSONResponse(
status_code=exc.status_code,
content=build_http_error_response(exc.status_code, exc.detail).model_dump(),
)
return JSONResponse(
status_code=500,
content=build_http_error_response(500, str(exc)).model_dump(),
)
async def _validation_error_handler(request: Request, exc: Exception) -> JSONResponse:
if isinstance(exc, RequestValidationError):
return JSONResponse(
status_code=422,
content=build_http_error_response(422, str(exc)).model_dump(),
)
return JSONResponse(
status_code=422,
content=build_http_error_response(422, str(exc)).model_dump(),
)
async def _route_generic_error_handler(request: Request, exc: Exception) -> JSONResponse:
log_unhandled_exception(request, exc)
return JSONResponse(
status_code=500,
content=build_http_error_response(500, str(exc)).model_dump(),
)
app.add_exception_handler(RequestValidationError, _validation_error_handler)
app.add_exception_handler(HTTPError, _route_http_error_handler)
app.add_exception_handler(StarletteHTTPException, _starlette_http_error_handler)
app.add_exception_handler(Exception, _route_generic_error_handler)
app.include_router(health_router)
app.include_router(generation_router)
app.include_router(models_router)
app.include_router(settings_router)
app.include_router(image_gen_router)
app.include_router(suggest_gap_prompt_router)
app.include_router(retake_router)
app.include_router(extend_router)
app.include_router(ic_lora_router)
app.include_router(lora_catalog_router)
app.include_router(prompt_enhancement_router)
app.include_router(runtime_policy_router)
app.include_router(hf_auth_router)
return app