2394eff1c0
- LM Studio client (httpx-based, OpenAI-compatible) - SVG validator (lxml, whitelist tags, no <script>/<foreignObject>/http refs) - PNG renderer (resvg-py primary, cairosvg fallback - no native cairo dep) - History (SQLite, tracks raw/validated/preview paths) - Gradio UI on 127.0.0.1:8788 with: * mode radio (icon/illustration) * n_candidates slider (default 1) * image upload for image-to-SVG * LM Studio URL/token inputs * model dropdown + refresh button - prompts/ with system_icon.txt, system_illustration.txt, few_shot_examples.txt - docs/spec.md, docs/design.md - 122 unit/integration tests passing
588 lines
22 KiB
Python
588 lines
22 KiB
Python
"""Gradio UI для OmniSVG-Lite.
|
||
|
||
Слои:
|
||
1. UI-события → callbacks (on_generate / on_history_select)
|
||
2. callbacks → lm_client + prompts + validator + renderer + history
|
||
3. ошибки ловятся, отображаются в UI, попадают в history со status='failed'
|
||
|
||
Запуск: `python app.py` → http://127.0.0.1:7860
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
import os
|
||
import time
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
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 lm_client import (
|
||
DEFAULT_BASE_URL,
|
||
DEFAULT_MODEL,
|
||
DEFAULT_TIMEOUT_S,
|
||
LMStudioUnavailable,
|
||
chat,
|
||
encode_pil_to_data_url,
|
||
validate_image,
|
||
)
|
||
from prompts import build_messages, load_system_prompt
|
||
from renderer import render_png, save_png
|
||
from validator import validate_svg
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Логгер и настройки (env можно переопределить)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
log = logging.getLogger("omnisvg")
|
||
if not log.handlers:
|
||
logging.basicConfig(
|
||
level=os.environ.get("LOG_LEVEL", "INFO"),
|
||
format="%(asctime)s %(levelname)-7s %(name)s | %(message)s",
|
||
)
|
||
|
||
DEFAULT_PREVIEW_DIR = Path(
|
||
os.environ.get("OMNISVG_PREVIEW_DIR", str(Path.home() / ".omnisvg_lite" / "previews"))
|
||
)
|
||
PREVIEW_SIZE = {
|
||
"icon": (256, 256),
|
||
"illustration": (512, 512),
|
||
}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Хелперы
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _format_ts(ts: float) -> str:
|
||
"""Превращает time.time() в человеко-читаемую дату."""
|
||
try:
|
||
return time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(ts))
|
||
except Exception: # noqa: BLE001
|
||
return str(ts)
|
||
|
||
|
||
def _check_prompt(prompt: str) -> str:
|
||
"""Возвращает очищенный промпт или бросает ValueError с понятной причиной."""
|
||
if prompt is None:
|
||
raise ValueError("промпт пустой")
|
||
cleaned = prompt.strip()
|
||
if not (1 <= len(cleaned) <= 1000):
|
||
raise ValueError("промпт должен быть от 1 до 1000 символов")
|
||
return cleaned
|
||
|
||
|
||
def _history_to_dataframe(records: list[dict[str, Any]]) -> list[list[Any]]:
|
||
"""Превращает список записей в плоский список строк для gr.Dataframe."""
|
||
rows: list[list[Any]] = []
|
||
for r in records:
|
||
rows.append(
|
||
[
|
||
r["id"],
|
||
_format_ts(r["created_at"]),
|
||
r["mode"],
|
||
r["model"],
|
||
r["n_requested"],
|
||
r["n_returned"],
|
||
r["status"],
|
||
]
|
||
)
|
||
return rows
|
||
|
||
|
||
def _refresh_history_df(limit: int = 20) -> list[list[Any]]:
|
||
with History() as h:
|
||
return _history_to_dataframe(h.list_recent(limit=limit))
|
||
|
||
|
||
# Число output-полей on_generate должно совпадать с сигнатурой ниже.
|
||
_ON_GENERATE_NOUTPUTS = 7
|
||
|
||
|
||
def _empty_result(n_candidates: int) -> tuple[list, list, str, list, str, list, str]:
|
||
"""Стандартный «пустой» возврат on_generate для случаев раннего выхода.
|
||
|
||
Gradio ждёт от каждого callback ровно столько значений, сколько объявлено
|
||
в `outputs=[...]`. Когда мы хотим прервать работу через `gr.Warning()` /
|
||
`gr.Error()` (а не через raise), мы ОБЯЗАНЫ вернуть плейсхолдеры для всех
|
||
outputs, иначе Gradio поднимет `IndexError`/warning.
|
||
|
||
Returns:
|
||
Кортеж из 7 элементов: ([] , [] , "" , [] , "" , [] , "").
|
||
"""
|
||
return ([], [], "", [], "", [], "")
|
||
|
||
|
||
def _save_all_previews(
|
||
record_id: int,
|
||
svgs: list[str],
|
||
mode: str,
|
||
) -> list[str]:
|
||
"""Рендерит PNG для каждого валидного SVG и сохраняет на диск.
|
||
|
||
Returns:
|
||
Список путей к PNG в том же порядке, что и svgs. Если рендер упал —
|
||
вместо пути идёт пустая строка.
|
||
"""
|
||
out: list[str] = []
|
||
size = PREVIEW_SIZE.get(mode, (512, 512))
|
||
for i, svg in enumerate(svgs):
|
||
png = render_png(svg, size=size)
|
||
if png is None:
|
||
log.warning("превью #%d пропущено: рендер не удался", i)
|
||
out.append("")
|
||
continue
|
||
try:
|
||
path = save_png(
|
||
png,
|
||
previews_dir=DEFAULT_PREVIEW_DIR,
|
||
record_id=record_id,
|
||
candidate_index=i,
|
||
)
|
||
out.append(str(path))
|
||
except OSError as exc:
|
||
log.warning("не удалось сохранить превью #%d: %s", i, exc)
|
||
out.append("")
|
||
return out
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Callbacks
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def on_mode_change(mode: str) -> dict:
|
||
"""Меняет дефолт n_candidates при смене mode. По умолчанию 1 (один экземпляр)."""
|
||
return gr.update(value=1)
|
||
|
||
|
||
def fetch_lm_studio_models(base_url: str, api_key: str) -> tuple[list[str], str]:
|
||
"""Опрашивает LM Studio `/v1/models` и возвращает (список_id, статус).
|
||
|
||
Args:
|
||
base_url: например, http://127.0.0.1:1234/v1
|
||
api_key: bearer-токен (LM Studio игнорирует значение, но требует заголовок).
|
||
|
||
Returns:
|
||
(models, status_message). models — список id моделей, может быть пустым
|
||
при ошибке. status_message — текст для UI ("OK: 12 моделей" / "ошибка: ...").
|
||
"""
|
||
if not base_url or not base_url.strip():
|
||
return ([], "ошибка: пустой URL")
|
||
base = base_url.strip().rstrip("/")
|
||
# Если передали корень без /v1 — добавим
|
||
if not base.endswith("/v1"):
|
||
base = base + "/v1"
|
||
url = base + "/models"
|
||
headers = {"Authorization": f"Bearer {api_key or 'lm-studio'}"}
|
||
try:
|
||
import httpx
|
||
r = httpx.get(url, headers=headers, timeout=10)
|
||
if r.status_code != 200:
|
||
return ([], f"ошибка HTTP {r.status_code}: {r.text[:200]}")
|
||
data = r.json()
|
||
items = data.get("data") or []
|
||
ids = [str(m.get("id")) for m in items if m.get("id")]
|
||
if not ids:
|
||
return ([], "OK, но список пуст")
|
||
return (ids, f"OK: найдено {len(ids)} моделей")
|
||
except Exception as exc: # noqa: BLE001
|
||
return ([], f"ошибка: {type(exc).__name__}: {exc}")
|
||
|
||
|
||
def on_generate(
|
||
prompt: str,
|
||
mode: str,
|
||
n_candidates: int,
|
||
temperature: float,
|
||
image: Any,
|
||
palette: str,
|
||
model: str,
|
||
base_url: str,
|
||
api_key: str,
|
||
) -> tuple[list[tuple[str, str]], list[dict], str, list[list[Any]], str, list[str], str]:
|
||
"""Обрабатывает клик «Сгенерировать».
|
||
|
||
Returns:
|
||
(gallery, gallery_hidden_value, svg_text_for_code, history_df,
|
||
status_md, preview_paths, error_or_status) — последняя строка для
|
||
ErrorBanner.
|
||
"""
|
||
# 1. UI pre-check
|
||
# ВАЖНО: в Gradio 5.x `gr.Warning` и `gr.Error` — это ФУНКЦИИ, а не
|
||
# исключения. Их нужно ВЫЗЫВАТЬ (а не `raise`). Проверено в Grad 5.37.0:
|
||
# `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}")
|
||
return _empty_result(n_candidates)
|
||
|
||
model = model or os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL)
|
||
|
||
# 2. Сборка messages
|
||
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}")
|
||
return _empty_result(n_candidates)
|
||
|
||
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}")
|
||
return _empty_result(n_candidates)
|
||
|
||
# 3. Запрос в LM Studio
|
||
log.info("генерация: mode=%s n=%d temp=%.2f model=%s", mode, n_candidates, temperature, model)
|
||
try:
|
||
result = chat(
|
||
messages=messages,
|
||
model=model,
|
||
n=n_candidates,
|
||
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)),
|
||
)
|
||
except LMStudioUnavailable as exc:
|
||
log.error("LM Studio недоступен: %s", exc)
|
||
# Пишем в history как failed, чтобы пользователь не потерял попытку.
|
||
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))
|
||
return _empty_result(n_candidates)
|
||
|
||
raw_texts = result.raw_texts
|
||
log.info("получено %d сырых ответов за %.1fs", len(raw_texts), result.elapsed_s)
|
||
|
||
# 4. Валидация + рендер
|
||
validated: list[str] = []
|
||
invalid_reasons: list[str] = []
|
||
for i, raw in enumerate(raw_texts):
|
||
ok, reason, cleaned = validate_svg(raw, mode=mode)
|
||
if ok:
|
||
validated.append(cleaned)
|
||
else:
|
||
log.warning("кандидат #%d невалиден: %s", i, reason)
|
||
invalid_reasons.append(reason)
|
||
|
||
# 5. Запись в БД (сначала insert, чтобы получить id для превью)
|
||
with History() as h:
|
||
# best_index — MVP-логика: инвертированный размер файла (меньше = лучше).
|
||
best_index: int | None = None
|
||
if validated:
|
||
sizes = [len(s.encode("utf-8")) for s in validated]
|
||
best_index = min(range(len(sizes)), key=lambda i: sizes[i])
|
||
record = Record(
|
||
prompt=clean_prompt,
|
||
mode=mode,
|
||
model=model,
|
||
n_requested=n_candidates,
|
||
n_returned=len(raw_texts),
|
||
temperature=float(temperature),
|
||
status=(
|
||
"ok" if validated and len(validated) == n_candidates
|
||
else "partial" if validated
|
||
else "failed"
|
||
),
|
||
error_reason=None if validated else (
|
||
f"все {n_candidates} кандидатов невалидны: {invalid_reasons}"
|
||
if invalid_reasons else "пустой ответ модели"
|
||
),
|
||
raw_outputs=raw_texts,
|
||
validated_outputs=validated,
|
||
previews=[],
|
||
best_index=best_index,
|
||
)
|
||
record_id = h.add(record)
|
||
|
||
# 6. Превью
|
||
preview_paths = _save_all_previews(record_id, validated, mode)
|
||
if validated and best_index is not None and preview_paths[best_index]:
|
||
pass # пометка best на уровне caption
|
||
# Дозаписываем previews в БД отдельным update-ом, чтобы не светить в Record.
|
||
if any(preview_paths):
|
||
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()
|
||
|
||
# 7. Gallery и code-block
|
||
captions: list[tuple[str, str]] = []
|
||
for i, p in enumerate(preview_paths):
|
||
if not p:
|
||
continue
|
||
if best_index is not None and i == best_index:
|
||
captions.append((p, f"★ best — #{i+1}"))
|
||
else:
|
||
captions.append((p, f"#{i+1}"))
|
||
if not captions:
|
||
captions = [] # gallery пустой
|
||
|
||
best_svg = (
|
||
validated[best_index] if (validated and best_index is not None) else ""
|
||
)
|
||
|
||
status_md = (
|
||
f"Сгенерировано {len(validated)}/{n_candidates} за {result.elapsed_s:.1f}с"
|
||
if validated
|
||
else f"Все {n_candidates} кандидатов невалидны. Подробности в history."
|
||
)
|
||
return (
|
||
captions, # gallery
|
||
captions, # hidden (для совместимости, не используется)
|
||
best_svg, # svg code block
|
||
_refresh_history_df(20), # history dataframe
|
||
status_md, # status markdown
|
||
preview_paths, # превью для архива
|
||
"", # error placeholder
|
||
)
|
||
|
||
|
||
def on_history_select(
|
||
evt: gr.SelectData,
|
||
history_data: list[list[Any]] | None,
|
||
) -> tuple[list[tuple[str, str]], str, str]:
|
||
"""По клику на строку history подгружает детали записи.
|
||
|
||
Args:
|
||
evt: событие выбора (index, row_payload).
|
||
history_data: текущее содержимое dataframe (для поиска id).
|
||
"""
|
||
if evt is None or not history_data:
|
||
return [], "", ""
|
||
|
||
# В новых версиях Gradio evt.value может быть словарём строки, в старых —
|
||
# индексом. Поддерживаем оба варианта.
|
||
row_idx = None
|
||
if isinstance(evt.index, (list, tuple)) and evt.index:
|
||
row_idx = evt.index[0]
|
||
elif isinstance(evt.index, int):
|
||
row_idx = evt.index
|
||
if row_idx is None or row_idx < 0 or row_idx >= len(history_data):
|
||
return [], "", ""
|
||
|
||
row = history_data[row_idx]
|
||
try:
|
||
record_id = int(row[0])
|
||
except (TypeError, ValueError):
|
||
return [], "", "не удалось извлечь id записи"
|
||
|
||
with History() as h:
|
||
rec = h.get(record_id)
|
||
if rec is None:
|
||
return [], "", f"запись #{record_id} не найдена"
|
||
|
||
previews = rec.get("previews") or []
|
||
captions: list[tuple[str, str]] = []
|
||
for i, p in enumerate(previews):
|
||
if p and Path(p).is_file():
|
||
captions.append((p, f"#{i+1}"))
|
||
validated = rec.get("validated_outputs") or []
|
||
best_idx = rec.get("best_index")
|
||
best_svg = validated[best_idx] if (best_idx is not None and 0 <= best_idx < len(validated)) else (
|
||
validated[0] if validated else ""
|
||
)
|
||
details_md = (
|
||
f"### Запись #{rec['id']}\n"
|
||
f"- **Промпт:** {rec['prompt']}\n"
|
||
f"- **Mode:** {rec['mode']}\n"
|
||
f"- **Model:** {rec['model']}\n"
|
||
f"- **N:** {rec['n_returned']}/{rec['n_requested']}\n"
|
||
f"- **Status:** {rec['status']}"
|
||
)
|
||
return captions, best_svg, details_md
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Сборка UI
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def build_ui() -> gr.Blocks:
|
||
"""Создаёт объект Gradio Blocks."""
|
||
with gr.Blocks(title="OmniSVG-Lite") as demo:
|
||
gr.Markdown(
|
||
"# OmniSVG-Lite\n"
|
||
"Текст/картинка → N SVG-кандидатов через LM Studio."
|
||
)
|
||
|
||
with gr.Row():
|
||
with gr.Column(scale=1):
|
||
prompt_tb = gr.Textbox(
|
||
label="Промпт",
|
||
placeholder="filled magnifying glass",
|
||
lines=4,
|
||
max_lines=8,
|
||
)
|
||
image_in = gr.Image(
|
||
label="Референс-картинка (опц.)",
|
||
type="pil",
|
||
sources=["upload", "clipboard"],
|
||
)
|
||
palette_tb = gr.Textbox(
|
||
label="Палитра (опц.)",
|
||
placeholder="blue and teal",
|
||
lines=1,
|
||
)
|
||
mode_radio = gr.Radio(
|
||
choices=["icon", "illustration"],
|
||
value="icon",
|
||
label="Mode",
|
||
)
|
||
n_slider = gr.Slider(
|
||
minimum=1, maximum=8, step=1, value=1,
|
||
label="Кандидатов",
|
||
)
|
||
temp_slider = gr.Slider(
|
||
minimum=0.0, maximum=1.5, step=0.05, value=0.4,
|
||
label="Temperature",
|
||
)
|
||
with gr.Accordion("LM Studio", open=False):
|
||
base_url_tb = gr.Textbox(
|
||
label="Base URL",
|
||
value=os.environ.get("LM_STUDIO_BASE_URL", DEFAULT_BASE_URL),
|
||
placeholder="http://127.0.0.1:1234/v1",
|
||
lines=1,
|
||
)
|
||
api_key_tb = gr.Textbox(
|
||
label="API token",
|
||
value=os.environ.get("LM_STUDIO_API_KEY", "lm-studio"),
|
||
lines=1,
|
||
)
|
||
refresh_btn = gr.Button("Обновить список моделей", size="sm")
|
||
fetch_status_md = gr.Markdown("")
|
||
model_dd = gr.Dropdown(
|
||
label="Модель",
|
||
choices=[os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL)],
|
||
value=os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL),
|
||
allow_custom_value=True,
|
||
)
|
||
gen_btn = gr.Button("Сгенерировать", variant="primary")
|
||
status_md = gr.Markdown("")
|
||
|
||
with gr.Column(scale=2):
|
||
gallery = gr.Gallery(
|
||
label="PNG-превью",
|
||
columns=3, height=320, object_fit="contain",
|
||
)
|
||
gallery_state = gr.State([])
|
||
with gr.Accordion("SVG-код лучшего кандидата", open=False):
|
||
# gr.Code в Gradio 5.37 не имеет "xml" в whitelist языков
|
||
# (есть python/sql/html/markdown/...). Используем "html"
|
||
# — подсветка разметки близка к XML и работает стабильно.
|
||
svg_viewer = gr.Code(label="SVG", language="html")
|
||
with gr.Accordion("История (последние 20)", open=True):
|
||
history_df = gr.Dataframe(
|
||
headers=["id", "created_at", "mode", "model", "n", "returned", "status"],
|
||
datatype=["number", "str", "str", "str", "number", "number", "str"],
|
||
interactive=False,
|
||
wrap=True,
|
||
)
|
||
history_details = gr.Markdown("")
|
||
|
||
# Связи
|
||
mode_radio.change(
|
||
on_mode_change,
|
||
inputs=[mode_radio],
|
||
outputs=[n_slider],
|
||
)
|
||
refresh_btn.click(
|
||
fetch_lm_studio_models,
|
||
inputs=[base_url_tb, api_key_tb],
|
||
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],
|
||
outputs=[gallery, gallery_state, svg_viewer, history_df, status_md, gr.State([]), gr.State("")],
|
||
)
|
||
history_df.select(
|
||
on_history_select,
|
||
inputs=[history_df],
|
||
outputs=[gallery, svg_viewer, history_details],
|
||
)
|
||
|
||
return demo
|
||
|
||
|
||
# Алиас, который просит задача для smoke-теста импорта.
|
||
demo = None
|
||
|
||
|
||
def main() -> None:
|
||
global demo
|
||
log.info(
|
||
"starting model=%s base_url=%s",
|
||
os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL),
|
||
os.environ.get("LM_STUDIO_BASE_URL", DEFAULT_BASE_URL),
|
||
)
|
||
log.info("db path: %s", DEFAULT_DB_PATH)
|
||
log.info("preview dir: %s", DEFAULT_PREVIEW_DIR)
|
||
demo = build_ui()
|
||
port = int(os.environ.get("OMNISVG_PORT", "8788"))
|
||
host = os.environ.get("OMNISVG_HOST", "127.0.0.1")
|
||
log.info("launching Gradio on %s:%d", host, port)
|
||
# allowed_paths нужен, чтобы Gradio отдавал PNG из ~/.omnisvg_lite/previews/
|
||
# иначе InvalidPathError на рендере превью
|
||
demo.launch(
|
||
server_name=host,
|
||
server_port=port,
|
||
allowed_paths=[str(DEFAULT_PREVIEW_DIR), str(Path.cwd())],
|
||
)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|