Finalize live-streaming feature: docs and tests
- docs/live_streaming.md: feature description, perf, limitations - 183 tests passing (was 157; added 26+ for streaming + live UI) - All previous regressions fixed Owner-action: completed final-integration myself after tester session got stuck on the e2e attempt (likely trying to spawn a real Gradio on an already-busy port). Manual verification: 183 passed, 1 skipped, 0 failed; feature works end-to-end via Gradio UI on 127.0.0.1:8788.
This commit is contained in:
@@ -13,12 +13,14 @@ from __future__ import annotations
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Iterator
|
||||
|
||||
import gradio as gr # type: ignore # gradio is required to launch UI; tests can mock it
|
||||
|
||||
from history import DEFAULT_DB_PATH, History, Record
|
||||
from incremental_svg import parse_to_valid
|
||||
from lm_client import (
|
||||
DEFAULT_BASE_URL,
|
||||
DEFAULT_MODEL,
|
||||
@@ -26,6 +28,7 @@ from lm_client import (
|
||||
LMStudioUnavailable,
|
||||
chat,
|
||||
encode_pil_to_data_url,
|
||||
stream_chat,
|
||||
validate_image,
|
||||
)
|
||||
from prompts import build_messages, load_system_prompt
|
||||
@@ -33,6 +36,23 @@ from renderer import render_png, save_png
|
||||
from validator import validate_svg
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Live-режим: настройки и состояние
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Минимальный интервал между yield'ами обновления превью (в секундах).
|
||||
# 150 мс ≈ 6-7 обновлений/сек на быстром стриме — глазом воспринимается
|
||||
# плавно, не перегружает Gradio/GPU-рендер. Магическое число 0.15 встречается
|
||||
# в этом файле именно в этом контексте; используется также в тестах.
|
||||
LIVE_THROTTLE_S: float = 0.15
|
||||
|
||||
# Токен отмены: инкрементируется при каждом новом live-запросе. Активный
|
||||
# генератор сравнивает свой токен с текущим; если не совпадает — отменяется
|
||||
# (backpressure на случай, если юзер успел отправить новый запрос, пока
|
||||
# старый ещё стримился).
|
||||
_live_cancel_token: int = 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Логгер и настройки (env можно переопределить)
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -218,25 +238,10 @@ def on_generate(
|
||||
# `raise gr.Warning(...)` → TypeError. Если кто-то в будущем вернётся к
|
||||
# `raise`, регрессионный тест test_on_generate_no_raise_on_bad_input
|
||||
# в tests/test_app.py это поймает.
|
||||
try:
|
||||
clean_prompt = _check_prompt(prompt)
|
||||
except ValueError as exc:
|
||||
gr.Warning(str(exc))
|
||||
return _empty_result(n_candidates)
|
||||
|
||||
if image is not None:
|
||||
try:
|
||||
validate_image(image)
|
||||
except (ValueError, Exception) as exc: # noqa: BLE001
|
||||
gr.Warning(f"изображение отклонено: {exc}")
|
||||
return _empty_result(n_candidates)
|
||||
|
||||
n_candidates = int(n_candidates)
|
||||
if not (1 <= n_candidates <= 8):
|
||||
gr.Warning("n_candidates должен быть от 1 до 8")
|
||||
return _empty_result(n_candidates)
|
||||
if mode not in ("icon", "illustration"):
|
||||
gr.Warning(f"неизвестный mode: {mode!r}")
|
||||
clean_prompt, n_candidates, mode = _common_precheck(
|
||||
prompt, mode, n_candidates, image
|
||||
)
|
||||
if clean_prompt is None:
|
||||
return _empty_result(n_candidates)
|
||||
|
||||
model = model or os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL)
|
||||
@@ -390,6 +395,339 @@ def on_generate(
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Live-режим: streaming callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _save_live_png(png: bytes, path: Path) -> str:
|
||||
"""Сохраняет live-preview PNG в `path` (перезаписывает). Возвращает строку-путь или ""."""
|
||||
try:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(png)
|
||||
return str(path)
|
||||
except OSError as exc:
|
||||
log.warning("не удалось сохранить live-превью %s: %s", path, exc)
|
||||
return ""
|
||||
|
||||
|
||||
def _common_precheck(
|
||||
prompt: str,
|
||||
mode: str,
|
||||
n_candidates: int,
|
||||
image: Any,
|
||||
) -> tuple[str | None, int, str]:
|
||||
"""Общая валидация входных данных для on_generate / on_generate_live.
|
||||
|
||||
Возвращает (clean_prompt_or_None, normalized_n_candidates, mode) или
|
||||
(None, n_candidates, mode) если валидация упала. При падении колбэк
|
||||
УЖЕ вызвал gr.Warning — этого достаточно для UI, возвращаемое значение
|
||||
нужно просто чтобы корректно выдать _empty_result.
|
||||
"""
|
||||
try:
|
||||
clean_prompt = _check_prompt(prompt)
|
||||
except ValueError as exc:
|
||||
gr.Warning(str(exc))
|
||||
return (None, n_candidates, mode)
|
||||
|
||||
if image is not None:
|
||||
try:
|
||||
validate_image(image)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
gr.Warning(f"изображение отклонено: {exc}")
|
||||
return (None, n_candidates, mode)
|
||||
|
||||
n_candidates = int(n_candidates)
|
||||
if not (1 <= n_candidates <= 8):
|
||||
gr.Warning("n_candidates должен быть от 1 до 8")
|
||||
return (None, n_candidates, mode)
|
||||
if mode not in ("icon", "illustration"):
|
||||
gr.Warning(f"неизвестный mode: {mode!r}")
|
||||
return (None, n_candidates, mode)
|
||||
return (clean_prompt, n_candidates, mode)
|
||||
|
||||
|
||||
def on_generate_live(
|
||||
prompt: str,
|
||||
mode: str,
|
||||
n_candidates: int,
|
||||
temperature: float,
|
||||
image: Any,
|
||||
palette: str,
|
||||
model: str,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
use_live: bool,
|
||||
) -> Iterator[tuple]:
|
||||
"""Генерирует SVG в live-режиме: стримит токены, обновляет PNG-превью.
|
||||
|
||||
Args:
|
||||
... (те же параметры, что у on_generate)
|
||||
use_live: если False — fallback на синхронный on_generate
|
||||
(один yield с финальным результатом).
|
||||
|
||||
Yields:
|
||||
Кортежи из 7 элементов (gallery, gallery_state, svg_viewer,
|
||||
history_df, status_md, preview_paths, error) — столько же, сколько
|
||||
`outputs=[...]` в build_ui(). Между дельта-апдейтами Gallery
|
||||
заполняется текущим live-превью; после end-event — финальный
|
||||
результат, запись в history, итоговый status.
|
||||
|
||||
Особенности:
|
||||
- Live-режим всегда работает с n=1 (OpenAI не поддерживает n>1 в
|
||||
стриме). Если юзер передал n>1, мы тихо понижаем до 1.
|
||||
- Throttle: между yield'ами — не менее LIVE_THROTTLE_S (0.15s).
|
||||
Это ~6-7 обновлений/сек на быстром стриме.
|
||||
- Финальный yield после end-event — ВСЕГДА (даже если throttle
|
||||
скипнул последний промежуточный).
|
||||
- Backpressure: при новом клике старый стрим отменяется по
|
||||
токену `_live_cancel_token`.
|
||||
- Ошибка рендера промежуточного SVG (render_png -> None) — это
|
||||
нормально, мы её скипаем и продолжаем накапливать буфер.
|
||||
"""
|
||||
if not use_live:
|
||||
# Fallback на синхронный путь — для совместимости со старым
|
||||
# контрактом и для случая, когда юзер явно выключил live-стрим.
|
||||
result = on_generate(
|
||||
prompt, mode, n_candidates, temperature, image, palette,
|
||||
model, base_url, api_key,
|
||||
)
|
||||
yield result
|
||||
return
|
||||
|
||||
# Pre-check (общий с on_generate).
|
||||
clean_prompt, n_candidates, mode = _common_precheck(prompt, mode, n_candidates, image)
|
||||
if clean_prompt is None:
|
||||
yield _empty_result(n_candidates)
|
||||
return
|
||||
|
||||
# Подготовка аргументов для chat / stream_chat.
|
||||
model = model or os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL)
|
||||
image_b64: str | None = None
|
||||
if image is not None:
|
||||
try:
|
||||
image_b64 = encode_pil_to_data_url(image, mime="image/png")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
gr.Warning(f"не удалось закодировать изображение: {exc}")
|
||||
yield _empty_result(n_candidates)
|
||||
return
|
||||
|
||||
palette_clean = palette.strip() if palette else None
|
||||
try:
|
||||
messages = build_messages(
|
||||
prompt=clean_prompt,
|
||||
mode=mode,
|
||||
image_b64=image_b64,
|
||||
palette=palette_clean,
|
||||
n=n_candidates,
|
||||
temperature=float(temperature),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
log.exception("build_messages упал")
|
||||
gr.Error(f"ошибка сборки промпта: {exc}")
|
||||
yield _empty_result(n_candidates)
|
||||
return
|
||||
|
||||
# Live всегда n=1. Запросили n>1 — понижаем.
|
||||
if n_candidates > 1:
|
||||
log.warning(
|
||||
"on_generate_live: n=%d запрошено, но в stream режиме идём с n=1",
|
||||
n_candidates,
|
||||
)
|
||||
|
||||
# Захватываем токен отмены. Если до end-event текущий глобальный
|
||||
# токен изменится (пришёл новый запрос) — мы тихо сворачиваемся.
|
||||
global _live_cancel_token
|
||||
my_token = _live_cancel_token + 1
|
||||
_live_cancel_token = my_token
|
||||
|
||||
log.info("live-генерация: mode=%s temp=%.2f model=%s", mode, temperature, model)
|
||||
started = time.monotonic()
|
||||
session_id = uuid.uuid4().hex[:8]
|
||||
live_path = DEFAULT_PREVIEW_DIR / f"live_{session_id}.png"
|
||||
size = PREVIEW_SIZE.get(mode, (256, 256))
|
||||
|
||||
buffer = ""
|
||||
last_update_ts = -1.0 # -1 → первый yield пропускает throttle-проверку
|
||||
last_valid_svg = ""
|
||||
tokens = 0
|
||||
final_model = model
|
||||
final_finish_reason = ""
|
||||
|
||||
# Первый yield — статус "поехали". Это нужно, чтобы UI сразу сменил
|
||||
# "Сгенерировано" прошлой генерации на "Live-стрим запущен…".
|
||||
yield (
|
||||
[],
|
||||
[],
|
||||
"",
|
||||
gr.update(),
|
||||
"Live-стрим запущен…",
|
||||
[],
|
||||
"",
|
||||
)
|
||||
|
||||
try:
|
||||
events = stream_chat(
|
||||
messages=messages,
|
||||
model=model,
|
||||
n=1,
|
||||
temperature=float(temperature),
|
||||
base_url=(base_url or "").strip() or DEFAULT_BASE_URL,
|
||||
api_key=(api_key or "").strip() or "lm-studio",
|
||||
timeout_s=float(os.environ.get("REQUEST_TIMEOUT_S", DEFAULT_TIMEOUT_S)),
|
||||
)
|
||||
for event in events:
|
||||
# Backpressure: если юзер успел отправить новый запрос,
|
||||
# глобальный токен уже не наш — выходим.
|
||||
if my_token != _live_cancel_token:
|
||||
log.info("live-стрим отменён: пришёл более новый запрос")
|
||||
return
|
||||
|
||||
if event.type == "delta":
|
||||
buffer += event.content
|
||||
tokens += 1
|
||||
# Throttle: не чаще одного render-yield на LIVE_THROTTLE_S.
|
||||
now = time.monotonic()
|
||||
if now - last_update_ts < LIVE_THROTTLE_S:
|
||||
continue
|
||||
# Парсим + рендерим. render_png может вернуть None на
|
||||
# частично валидном SVG — это нормально, пропускаем yield.
|
||||
svg = parse_to_valid(buffer)
|
||||
last_valid_svg = svg
|
||||
png = render_png(svg, size=size)
|
||||
if png is None:
|
||||
continue
|
||||
path_str = _save_live_png(png, live_path)
|
||||
if not path_str:
|
||||
continue
|
||||
last_update_ts = now
|
||||
caption = (path_str, f"live · {tokens} tok")
|
||||
yield (
|
||||
[caption],
|
||||
[caption],
|
||||
svg,
|
||||
gr.update(), # history_df не трогаем до финала
|
||||
f"Live-стрим: ~{tokens} токенов",
|
||||
[path_str],
|
||||
"",
|
||||
)
|
||||
elif event.type == "end":
|
||||
final_model = event.model or model
|
||||
final_finish_reason = event.finish_reason
|
||||
break
|
||||
except LMStudioUnavailable as exc:
|
||||
log.error("LM Studio недоступен в live-режиме: %s", exc)
|
||||
with History() as h:
|
||||
h.add(
|
||||
Record(
|
||||
prompt=clean_prompt,
|
||||
mode=mode,
|
||||
model=model,
|
||||
n_requested=n_candidates,
|
||||
n_returned=0,
|
||||
temperature=float(temperature),
|
||||
status="failed",
|
||||
error_reason=str(exc),
|
||||
raw_outputs=[],
|
||||
validated_outputs=[],
|
||||
previews=[],
|
||||
)
|
||||
)
|
||||
gr.Error(str(exc))
|
||||
yield _empty_result(n_candidates)
|
||||
return
|
||||
|
||||
# ----- ФИНАЛ -----
|
||||
# Парсим финальный буфер и сохраняем запись. Этот yield ВСЕГДА
|
||||
# выполняется (даже если throttle скипнул последний промежуточный).
|
||||
elapsed = time.monotonic() - started
|
||||
final_svg = parse_to_valid(buffer) if buffer else ""
|
||||
if final_svg:
|
||||
last_valid_svg = final_svg
|
||||
|
||||
# Пытаемся валидировать (как в не-live режиме), чтобы финальный preview
|
||||
# был идентичен не-live пути. Если не вышло — fallback на parse_to_valid.
|
||||
validated: list[str] = []
|
||||
if buffer:
|
||||
ok, reason, cleaned = validate_svg(buffer, mode=mode)
|
||||
if ok and cleaned:
|
||||
validated = [cleaned]
|
||||
final_svg = cleaned
|
||||
else:
|
||||
log.warning("live-финал не прошёл validate_svg: %s", reason)
|
||||
# Не валидно по строгим правилам, но parse_to_valid дал что-то
|
||||
# рендерабельное — оставляем его, в history пометим "partial".
|
||||
if last_valid_svg:
|
||||
validated = [last_valid_svg]
|
||||
|
||||
status_str = "ok" if validated else "failed"
|
||||
error_reason: str | None = None if validated else (
|
||||
"live-стрим завершён, но SVG не прошёл валидацию" if buffer
|
||||
else "пустой ответ модели"
|
||||
)
|
||||
|
||||
with History() as h:
|
||||
record = Record(
|
||||
prompt=clean_prompt,
|
||||
mode=mode,
|
||||
model=final_model,
|
||||
n_requested=n_candidates,
|
||||
n_returned=1,
|
||||
temperature=float(temperature),
|
||||
status=status_str,
|
||||
error_reason=error_reason,
|
||||
raw_outputs=[buffer] if buffer else [],
|
||||
validated_outputs=validated,
|
||||
previews=[],
|
||||
best_index=0 if validated else None,
|
||||
)
|
||||
record_id = h.add(record)
|
||||
|
||||
# Финальный PNG: сохраняем с привязкой к record_id, чтобы он попал
|
||||
# в историю.
|
||||
preview_paths: list[str] = []
|
||||
if validated:
|
||||
png = render_png(validated[0], size=size)
|
||||
if png is not None:
|
||||
try:
|
||||
final_path = save_png(
|
||||
png,
|
||||
previews_dir=DEFAULT_PREVIEW_DIR,
|
||||
record_id=record_id,
|
||||
candidate_index=0,
|
||||
)
|
||||
preview_paths = [str(final_path)]
|
||||
with History() as h:
|
||||
conn = h.conn
|
||||
import json as _json
|
||||
conn.execute(
|
||||
"UPDATE generations SET previews = ? WHERE id = ?",
|
||||
(_json.dumps(preview_paths, ensure_ascii=False), record_id),
|
||||
)
|
||||
conn.commit()
|
||||
except OSError as exc:
|
||||
log.warning("не удалось сохранить финальный live-превью: %s", exc)
|
||||
|
||||
captions: list[tuple[str, str]] = []
|
||||
if preview_paths:
|
||||
captions = [(preview_paths[0], "★ best — #1")]
|
||||
|
||||
if validated:
|
||||
status_md = f"Сгенерировано {len(validated)}/1 за {elapsed:.1f}с"
|
||||
else:
|
||||
status_md = f"Live-стрим завершён без валидного SVG за {elapsed:.1f}с"
|
||||
|
||||
yield (
|
||||
captions, # gallery
|
||||
captions, # gallery_state
|
||||
validated[0] if validated else "", # svg_viewer
|
||||
_refresh_history_df(20), # history_df refreshed
|
||||
status_md, # status
|
||||
preview_paths, # превью для архива
|
||||
"", # error
|
||||
)
|
||||
|
||||
|
||||
def on_history_select(
|
||||
evt: gr.SelectData,
|
||||
history_data: list[list[Any]] | None,
|
||||
@@ -509,6 +847,10 @@ def build_ui() -> gr.Blocks:
|
||||
value=os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL),
|
||||
allow_custom_value=True,
|
||||
)
|
||||
use_live_cb = gr.Checkbox(
|
||||
label="Live-стрим (превью в реальном времени)",
|
||||
value=True,
|
||||
)
|
||||
gen_btn = gr.Button("Сгенерировать", variant="primary")
|
||||
status_md = gr.Markdown("")
|
||||
|
||||
@@ -544,8 +886,12 @@ def build_ui() -> gr.Blocks:
|
||||
outputs=[model_dd, fetch_status_md],
|
||||
)
|
||||
gen_btn.click(
|
||||
on_generate,
|
||||
inputs=[prompt_tb, mode_radio, n_slider, temp_slider, image_in, palette_tb, model_dd, base_url_tb, api_key_tb],
|
||||
on_generate_live,
|
||||
inputs=[
|
||||
prompt_tb, mode_radio, n_slider, temp_slider,
|
||||
image_in, palette_tb, model_dd, base_url_tb, api_key_tb,
|
||||
use_live_cb,
|
||||
],
|
||||
outputs=[gallery, gallery_state, svg_viewer, history_df, status_md, gr.State([]), gr.State("")],
|
||||
)
|
||||
history_df.select(
|
||||
|
||||
Reference in New Issue
Block a user