Compare commits

..
Author SHA1 Message Date
petya 5427d494a7 Initial commit 2026-06-15 06:57:58 +00:00
26 changed files with 190 additions and 7253 deletions

No files matched your search

+170 -29
View File
@@ -1,35 +1,176 @@
# Python
# ---> Python
# Byte-compiled / optimized / DLL files
__pycache__/
*.pyc
*.pyo
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
# Usually these files are written by a python script from a template
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
# Installer logs
pip-log.txt
pip-delete-this-directory.txt
# Unit test / coverage reports
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
*.py,cover
.hypothesis/
.pytest_cache/
cover/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
db.sqlite3-journal
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
# PyBuilder
.pybuilder/
target/
# Jupyter Notebook
.ipynb_checkpoints
# IPython
profile_default/
ipython_config.py
# pyenv
# For a library or package, you might want to ignore these files since the code is
# intended to run in multiple environments; otherwise, check them in:
# .python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don't work, or not
# install all needed dependencies.
#Pipfile.lock
# UV
# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control.
# This is especially recommended for binary packages to ensure reproducibility, and is more
# commonly ignored for libraries.
#uv.lock
# poetry
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
# This is especially recommended for binary packages to ensure reproducibility, and is more
# commonly ignored for libraries.
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
#poetry.lock
# pdm
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
#pdm.lock
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
# in version control.
# https://pdm.fming.dev/latest/usage/project/#working-with-version-control
.pdm.toml
.pdm-python
.pdm-build/
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
__pypackages__/
# Celery stuff
celerybeat-schedule
celerybeat.pid
# SageMath parsed files
*.sage.py
# Environments
.env
.venv
env/
venv/
ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Pyre type checker
.pyre/
# pytype static type analyzer
.pytype/
# Cython debug symbols
cython_debug/
# PyCharm
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
# and can be added to the global gitignore or merged into this file. For a more nuclear
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/
# Ruff stuff:
.ruff_cache/
# venv
.venv/
venv/
env/
# PyPI configuration file
.pypirc
# OS
.DS_Store
Thumbs.db
# Local logs / runtime
gradio_run.log
gradio_run.err
launcher_run.log
launcher_run.err
_smk_*.py
_probe_*.py
_create_shortcut.ps1
*.tmp
# Build / dist
build/
dist/
*.egg-info/
# Editor
.vscode/
.idea/
+18
View File
@@ -0,0 +1,18 @@
MIT License
Copyright (c) 2026 petya
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and
associated documentation files (the "Software"), to deal in the Software without restriction, including
without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the
following conditions:
The above copyright notice and this permission notice shall be included in all copies or substantial
portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT
LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO
EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE
USE OR OTHER DEALINGS IN THE SOFTWARE.
+2 -118
View File
@@ -1,119 +1,3 @@
# OmniSVG-Lite
# omnisvg-lite
Локальный MVP для генерации SVG-иконок и иллюстраций через LM Studio.
Gradio UI, локальная SQLite-история, валидация выхода модели через `lxml`,
PNG-превью через `cairosvg`.
> Контракт: см. `docs/spec.md` и `docs/design.md`. Промпт-шаблоны: `prompts/`.
## Что внутри
| Файл | Назначение |
|--------------------|------------------------------------------------------------------|
| `app.py` | Gradio 5 UI: text/prompt → N SVG-кандидатов → PNG-превью → БД |
| `lm_client.py` | httpx-клиент к OpenAI-compatible LM Studio, `generate_svg()` |
| `prompts.py` | Загрузка шаблонов из `prompts/` + `build_messages()` |
| `validator.py` | lxml-парсинг + проверка whitelist/тегов/href/on* |
| `renderer.py` | cairosvg → PNG (graceful fallback, если cairo недоступен) |
| `history.py` | SQLite WAL, контекст-менеджер `History` |
| `prompts/` | system-инструкции + few-shot примеры |
| `tests/` | pytest на validator (6+ кейсов) |
## Установка
```bash
# 1) зависимости
pip install -r requirements.txt
```
> **Важно для Windows**: `cairosvg` — это Python-биндинг к нативному `cairo`.
> На Windows чистого `pip install cairosvg` обычно недостаточно — нужна
> библиотека `cairo.dll`. Варианты:
>
> ```bash
> # вариант 1: conda (рекомендуется)
> conda install -c conda-forge pycairo cairo
>
> # вариант 2: MSYS2 / vcpkg / GTK3 runtime
> ```
>
> Если cairo не установлен — UI стартует, но PNG-превью будут пропущены
> (в логе появится `WARNING` от renderer.py).
## Настройка LM Studio
1. Запустить LM Studio локально (`http://127.0.0.1:1234`).
2. Загрузить модель, например `qwen/qwen3.5-35b-a3b`.
3. Включить **OpenAI-compatible server** в LM Studio.
Переменные окружения (опционально, все имеют дефолты):
| Переменная | Дефолт |
|---------------------------|-----------------------------------|
| `LM_STUDIO_BASE_URL` | `http://127.0.0.1:1234/v1` |
| `LM_STUDIO_API_KEY` | `lm-studio` |
| `DEFAULT_MODEL` | `qwen/qwen3.5-35b-a3b` |
| `REQUEST_TIMEOUT_S` | `120` |
| `OMNISVG_DB_PATH` | `~/.omnisvg_lite/history.sqlite` |
| `OMNISVG_PREVIEW_DIR` | `~/.omnisvg_lite/previews` |
| `LOG_LEVEL` | `INFO` |
## Запуск
```bash
python app.py
```
UI откроется на `http://127.0.0.1:7860`. Кнопка «Сгенерировать» отправляет
промпт в LM Studio, валидирует SVG, рисует PNG-превью и сохраняет запись
в SQLite.
## Тесты
```bash
python -m pytest tests/ -v
```
Покрытие — `validator.py` (well-formed, viewBox, forbidden tags, http refs,
on*-атрибуты).
## Smoke-тест импорта (без живого LM Studio)
```bash
python -c "from app import demo, main; print('imports ok')"
```
Запускается **только импорт**, Gradio-сервер не поднимается. Если gradio не
установлен, импорт `app.py` упадёт (gradio захардкожен в `import`).
## Известные ограничения (MVP)
- Один пользователь, без авторизации.
- Выбор «лучшего» кандидата — по минимальному размеру файла (прокси
компактности, не качества). См. `spec.md §4`.
- `cairosvg` на Windows без conda — см. выше; UI работает и без PNG-превью.
- Размер БД не ограничен; очистка старых превью — на совести пользователя
(дизайн §11, open question).
## Структура SQLite
```sql
CREATE TABLE generations (
id INTEGER PRIMARY KEY AUTOINCREMENT,
created_at REAL NOT NULL,
prompt TEXT NOT NULL,
mode TEXT NOT NULL, -- icon | illustration
model TEXT NOT NULL,
n_requested INTEGER NOT NULL,
n_returned INTEGER NOT NULL,
temperature REAL NOT NULL,
status TEXT NOT NULL, -- ok | partial | failed
error_reason TEXT,
raw_outputs TEXT NOT NULL, -- JSON list[str]
validated_outputs TEXT NOT NULL, -- JSON list[str]
previews TEXT NOT NULL, -- JSON list[str] (PNG paths)
best_index INTEGER
);
```
WAL-режим, индекс по `created_at DESC`.
Text/image to SVG generator powered by LM Studio VLM (qwen3.5, gemma, etc.) with live streaming preview.
-933
View File
@@ -1,933 +0,0 @@
"""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
import uuid
from pathlib import Path
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,
DEFAULT_TIMEOUT_S,
LMStudioUnavailable,
chat,
encode_pil_to_data_url,
stream_chat,
validate_image,
)
from prompts import build_messages, load_system_prompt
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 можно переопределить)
# ---------------------------------------------------------------------------
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 это поймает.
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)
# 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
)
# ---------------------------------------------------------------------------
# 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,
) -> 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,
)
use_live_cb = gr.Checkbox(
label="Live-стрим (превью в реальном времени)",
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_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(
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()
-395
View File
@@ -1,395 +0,0 @@
# OmniSVG-Lite — Design
> Архитектура MVP. Backend-dev может взять этот документ + `prompts/` и писать код без уточнений.
---
## 1. Стек и обоснование
| Компонент | Выбор | Зачем |
|---|---|---|
| Язык | Python 3.11 | Зафиксировано пользователем. |
| UI | Gradio 5 | `Blocks` + встроенный Gallery. Альтернативы (Streamlit, NiceGUI) хуже по лайауту. |
| HTTP | httpx (sync) | Проще openai SDK, JSON+timeouts из коробки. |
| Парсинг/валидация | lxml | Строгий, `XMLSyntaxError` с координатами, `iter()` для обхода. |
| PNG-рендер | cairosvg | Декларативный, синхронный. Pillow+svglib — хуже на градиентах. |
| БД | sqlite3 (stdlib) | WAL, 1 таблица. SQLAlchemy — overkill. |
| Логирование | logging (stdlib) | Достаточно. |
| Concurrency | threading (stdlib) | UI не фризит во время LM-вызова. |
| Тесты | pytest | validator/renderer/prompts покрыты. |
**Не берём:** `openai` SDK (httpx достаточно), `pydantic` (dataclass хватит), `tenacity` (одна попытка в MVP).
---
## 2. Структура модулей
```
OmniSVG-Lite/
├── app.py # Gradio UI, callbacks, main()
├── lm_client.py # httpx-обёртка над LM Studio
├── prompts.py # Сборка messages[], загрузка шаблонов
├── validator.py # lxml-парсинг + проверка правил SVG
├── renderer.py # cairosvg → PNG bytes + сохранение
├── history.py # SQLite-схема, CRUD
├── data/{omnisvg.db, previews/<id>_<n>.png}
├── prompts/{system_icon.txt, system_illustration.txt, few_shot_examples.txt}
├── docs/{spec.md, design.md}
└── tests/{test_validator.py, test_renderer.py, test_prompts.py}
```
### 2.1. `lm_client.py` — httpx → LM Studio
Импорт: `app.py`. Зависимости: httpx, base64, json, logging.
```python
class LMTurnResult:
raw_texts: list[str] # N строк, по одной на кандидат
elapsed_s: float
model: str
usage: dict | None
def chat(*, base_url, api_key, model, messages, n, temperature,
max_tokens=4096, timeout_s=60.0) -> LMTurnResult: ...
```
### 2.2. `prompts.py` — сборка messages
Импорт: `app.py`. Зависимости: pathlib, prompts/.
```python
def load_system_prompt(mode: Literal["icon","illustration"]) -> str: ...
def load_few_shot() -> list[dict]: ... # [{role,content}, ...]
def build_messages(*, mode, prompt, image_b64: str|None,
palette: str|None, n: int, temperature: float) -> list[dict]: ...
```
### 2.3. `validator.py` — lxml + правила
Импорт: `app.py`. Зависимости: lxml.etree, re.
```python
class ValidatorError(Exception): # code: malformed_xml|not_svg|missing_viewbox|bad_viewbox|disallowed_tag|unknown_tag|external_ref|too_large
code: str; detail: str
ALLOWED_TAGS = {"svg","g","defs","symbol","use","path","rect","circle","ellipse",
"line","polygon","polyline","linearGradient","radialGradient","stop",
"filter","feGaussianBlur","feOffset","feBlend","feMerge","feMergeNode",
"feFlood","feComposite","feColorMatrix","clipPath","mask","pattern",
"text","tspan","textPath","title","desc"}
DISALLOWED_TAGS = {"script","foreignObject","image"}
MAX_BYTES = {"icon": 32*1024, "illustration": 256*1024}
def extract_svg(text: str) -> str: ... # regex
def validate(svg_text: str, *, mode) -> etree._Element: ... # raises ValidatorError
```
### 2.4. `renderer.py` — cairosvg
Импорт: `app.py`. Зависимости: cairosvg, pathlib, io.
```python
def svg_to_png_bytes(svg_text: str, *, output_width: int) -> bytes: ...
# добавляет xmlns если нет, output_width=256 (icon) / 512 (illustration)
def save_preview(png_bytes, *, job_id: int, candidate_index: int) -> Path: ...
def to_pil_image(png_bytes: bytes) -> "PIL.Image.Image": ...
```
### 2.5. `history.py` — SQLite CRUD
Импорт: `app.py`. Зависимости: sqlite3, json, time, pathlib.
```python
DB_PATH = Path("data/omnisvg.db")
def init_db() -> None: ... # WAL mode
def insert_job(*, prompt, mode, model, n_candidates, status,
raw_outputs: list[str], validated_outputs: list[str],
preview_paths: list[str], error: str|None = None) -> int: ...
def list_jobs(limit: int = 50) -> list[dict]: ... # для History table
def get_job(job_id: int) -> dict | None: ...
```
### 2.6. `app.py` — UI + wiring
Импорт: всё перечисленное + gradio.
```python
def on_generate(prompt, mode, n, temperature, image, palette, model) -> \
tuple[list[tuple[str,str]], list[dict], str]: ...
def on_history_select(evt: gr.SelectData): ... # row click
def main() -> None: ...
```
---
## 3. LM Studio client — детали
- **POST** `http://127.0.0.1:1234/v1/chat/completions`, заголовки: `Content-Type: application/json`, `Authorization: Bearer lm-studio` (значение игнорируется, но требуется OpenAI-совместимым клиентом).
- **`stream=false`** в payload: n>1 плохо сочетается со streaming, плюс инкрементальный парсинг сложнее.
- **Payload (text-to-SVG):** `{"model", "messages":[{"role":"system","content":...}], "n", "temperature", "max_tokens":4096, "stream":false}`.
- **Payload (image-to-SVG):** `content` это `[{type:"text", text:...}, {type:"image_url", image_url:{url:"data:image/png;base64,..."}}]` — текст **первым**, картинка **второй**.
- **Парсинг ответа:** `choices[i].message.content` — N строк. `tool_calls` не передаём; если модель их всё же вернула, игнорируем (если `content` пустой — кандидат invalid).
- **Кодирование картинки:** convert в RGB, проверка `max(size) <= 4096`, PIL save PNG optimize → base64 в data-URL. JPEG **не используем** (LM Studio VLM обучены на PNG, артефактов на границах не будет).
---
## 4. Промпт-стратегия
**Структура messages[]:**
```
system: <system_icon.txt | system_illustration.txt>
user: <few-shot user 1>
assistant: <few-shot svg 1>
user: <few-shot user 2>
assistant: <few-shot svg 2>
... (3-5 пар)
user: <actual user prompt + (опц.) image>
```
**Требования к SVG в ответе модели:**
- Один корневой `<svg>...</svg>`, **ничего больше** (никаких ``` ``` ```, никаких пояснений).
- `viewBox` обязателен и совпадает с mode: `0 0 64 64` (icon) или `0 0 512 512` (illustration).
- Без `<script>`, `<foreignObject>`, `<image>`, `http://`, `https://`, `data:` в href.
- Размер файла ≤ лимита mode (32/256 KB).
- `xmlns="http://www.w3.org/2000/svg"` желателен (валидатор примет и без, но cairosvg ругается — добавляем автоматически в renderer).
**Few-shot (готов в `prompts/few_shot_examples.txt`):** 2 иконки (filled magnifying glass, outline battery), 2 иллюстрации (flat fox, isometric server room), 1 пример refusal (NSFW-отказ без SVG).
---
## 5. Валидация — что проверяет lxml
```python
def validate(svg_text, *, mode):
svg_text = extract_svg(svg_text) # вытащить <svg>...</svg>
try:
root = etree.fromstring(svg_text.encode("utf-8"))
except etree.XMLSyntaxError as e:
raise ValidatorError("malformed_xml", str(e))
local_root = root.tag.split("}")[-1]
if local_root != "svg": raise ValidatorError("not_svg", root.tag)
if not root.get("viewBox"): raise ValidatorError("missing_viewbox", "")
try:
vb = [float(x) for x in root.get("viewBox").split()]
if len(vb) != 4: raise ValueError
except ValueError:
raise ValidatorError("bad_viewbox", root.get("viewBox"))
for el in root.iter():
local = el.tag.split("}")[-1]
if local in DISALLOWED_TAGS: raise ValidatorError("disallowed_tag", local)
if local not in ALLOWED_TAGS: raise ValidatorError("unknown_tag", local)
for attr, val in el.attrib.items():
if val and ("http://" in val or "https://" in val):
if attr.endswith("href") or "url(" in val:
raise ValidatorError("external_ref", f"{attr}={val[:50]}")
if len(svg_text.encode("utf-8")) > MAX_BYTES[mode]:
raise ValidatorError("too_large", f"{len(svg_text)} bytes")
return root
```
`SVG_NS = "http://www.w3.org/2000/svg"`. lxml возвращает теги в Clark-notation `{ns}local`.
---
## 6. Рендер
```python
def svg_to_png_bytes(svg_text, *, output_width):
if "xmlns=" not in svg_text.split(">", 1)[0]:
svg_text = svg_text.replace("<svg", '<svg xmlns="http://www.w3.org/2000/svg"', 1)
return cairosvg.svg2png(
bytestring=svg_text.encode("utf-8"),
output_width=output_width, # 256 icon, 512 illustration
background_color="white",
)
```
PNG bytes → `save_preview` (файл) + `to_pil_image` (для Gallery). Лучший получает caption `"★ best — #1"`.
**Windows + Cairo:** если `pycairo` поставлен через conda — работает. Через pip — может потребоваться GTK runtime. В README отметить, **в MVP не чинить**.
---
## 7. История — SQLite схема
```sql
CREATE TABLE IF NOT EXISTS jobs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
created_at REAL NOT NULL, -- time.time()
prompt TEXT NOT NULL,
mode TEXT NOT NULL CHECK (mode IN ('icon','illustration')),
model TEXT NOT NULL,
n_candidates INTEGER NOT NULL,
temperature REAL NOT NULL,
status TEXT NOT NULL, -- ok | all_invalid | error
error TEXT,
raw_outputs TEXT NOT NULL, -- JSON list[str]
validated_outputs TEXT NOT NULL, -- JSON list[str]
preview_paths TEXT NOT NULL, -- JSON list[str] (rel paths)
best_index INTEGER -- индекс лучшего или NULL
);
CREATE INDEX IF NOT EXISTS idx_jobs_created_at ON jobs(created_at DESC);
```
WAL mode в `init_db()`: `PRAGMA journal_mode=WAL`. JSON через `json.dumps/loads`. 1 запись = 1 запуск, N кандидатов внутри.
---
## 8. Gradio UI
### 8.1. Layout (ASCII)
```
┌─────────────────────────────────────────────────────────────────┐
│ OmniSVG-Lite [Model: ▼] │
├──────────────────┬──────────────────────────────────────────────┤
│ Prompt │ Gallery (PNG превью) │
│ [textbox] │ [★best] [#2] [#3] │
│ Reference image │ │
│ [image upload] │ Selected SVG: │
│ Palette (opt) │ [code block] │
│ [textbox] │ │
│ Mode │ History │
│ (•) icon ( ) ill │ [gr.Dataframe] │
│ Candidates 1..8 │ │
│ [slider] │ │
│ Temperature │ │
│ [slider] │ │
│ [Generate] │ │
│ Status: [md] │ │
└──────────────────┴──────────────────────────────────────────────┘
```
### 8.2. Компоненты
| gr.X | Имя | Параметры |
|---|---|---|
| `Textbox` | `prompt_tb` | lines=4, max_lines=8 |
| `Image` | `image_in` | type="pil", sources=["upload","clipboard"] |
| `Textbox` | `palette_tb` | lines=1 |
| `Radio` | `mode_radio` | ["icon","illustration"], value="icon" |
| `Slider` | `n_slider` | 1..8, step=1, value=4 (icon)/2 (illustration) |
| `Slider` | `temp_slider` | 0..1.5, step=0.05, value=0.4 |
| `Button` | `gen_btn` | "Generate", variant="primary" |
| `Gallery` | `gallery` | columns=3, height=320, object_fit="contain" |
| `Code` | `svg_viewer` | language="xml" |
| `Dataframe` | `history_df` | id/created_at/mode/model/n/status |
| `Markdown` | `status_md` | для статус-строки |
| `Dropdown` | `model_dd` | из DEFAULT_MODEL env |
### 8.3. Callbacks (контракт)
```python
def on_mode_change(mode): return gr.update(value=4 if mode=="icon" else 2)
def on_generate(prompt, mode, n, temp, image, palette, model):
if not (1 <= len(prompt.strip()) <= 1000):
raise gr.Warning("Prompt must be 1–1000 characters")
if image is not None: _check_image(image) # size/format
messages = build_messages(mode=mode, prompt=prompt,
image_b64=encode_image(image) if image else None,
palette=palette or None, n=n, temperature=temp)
result = chat(base_url=..., model=model, messages=messages, n=n, temperature=temp)
previews, validated = [], []
for i, raw in enumerate(result.raw_texts):
try:
svg = extract_svg(raw); validate(svg, mode=mode)
png = svg_to_png_bytes(svg, output_width=256 if mode=="icon" else 512)
path = save_preview(png, job_id=0, candidate_index=i)
validated.append(svg); previews.append((str(path), f"#{i+1}"))
except ValidatorError as e:
log.warning("invalid candidate %d: %s", i, e.code)
best_idx = None
if validated:
sizes = [len(s.encode("utf-8")) for s in validated]
best_idx = min(range(len(sizes)), key=lambda i: sizes[i])
previews[best_idx] = (previews[best_idx][0], f"★ best — #{best_idx+1}")
job_id = insert_job(prompt=prompt, mode=mode, model=model, n_candidates=n,
temperature=temp, status="ok" if validated else "all_invalid",
raw_outputs=result.raw_texts, validated_outputs=validated,
preview_paths=[p[0] for p in previews],
best_index=best_idx)
return previews, validated[best_idx] if validated else "", _refresh_history(), \
f"Generated {len(validated)}/{n}"
def on_history_select(evt: gr.SelectData):
job = get_job(evt.value[0])
return [(p, "") for p in job["preview_paths"]], \
job["validated_outputs"][job["best_index"] or 0]
```
Gradio сам сериализует вызовы одного callback; `gen_btn.interactive = False` через `.then()` достаточно.
---
## 9. Обработка ошибок
| Слой | Что ловим | UI | Лог | history.status |
|---|---|---|---|---|
| UI pre-check | пустой/длинный prompt, плохая картинка | `gr.Warning` | INFO | — (запись не создаётся) |
| LM Studio | timeout, 5xx, network | `gr.Error` | ERROR+tb | `error` |
| Парсер | модель вернула текст без `<svg>` | — | WARNING | `all_invalid` / частично `ok` |
| Валидатор | malformed, viewBox, tag, ref, size | — | WARNING (с code) | как выше |
| Renderer | cairosvg бросил | invalid | WARNING | как выше |
| History | SQLite locked / disk full | `gr.Error("DB error")` | ERROR | — |
**Принцип:** всё, что ниже UI pre-check, **не блокирует** показ уже валидных кандидатов. Из 4 валиден 1 → показываем 1, статус `ok`, в `raw_outputs` все 4 (включая невалидные).
---
## 10. Запуск
### 10.1. Команды
```bash
python -m venv .venv
.venv\Scripts\activate # Windows
pip install -r requirements.txt
python app.py
```
UI поднимется на `http://127.0.0.1:7860` (Gradio default).
### 10.2. requirements.txt (стартовая точка)
```
gradio>=5.0,<6
httpx>=0.27
lxml>=5.0
cairosvg>=2.7
Pillow>=10.0
pytest>=8.0 # dev only
```
### 10.3. Env-переменные (опционально, с дефолтами)
| Переменная | Дефолт | Назначение |
|---|---|---|
| `LM_STUDIO_BASE_URL` | `http://127.0.0.1:1234/v1` | endpoint |
| `LM_STUDIO_API_KEY` | `lm-studio` | Authorization header |
| `DEFAULT_MODEL` | `qwen/qwen3.5-35b-a3b` | Dropdown default |
| `REQUEST_TIMEOUT_S` | `60` | httpx timeout |
| `OMNISVG_DB_PATH` | `data/omnisvg.db` | SQLite |
| `OMNISVG_PREVIEW_DIR` | `data/previews` | PNG-превью |
| `LOG_LEVEL` | `INFO` | DEBUG/INFO/WARNING |
### 10.4. Логи при старте
```
INFO omnisvg starting model=qwen/qwen3.5-35b-a3b base_url=http://127.0.0.1:1234/v1
INFO omnisvg db initialized at data/omnisvg.db (WAL)
INFO omnisvg ui on http://127.0.0.1:7860
```
---
## 11. Open Questions
- **Streaming** — не идём, non-stream. Если понадобится: второй endpoint в `lm_client.py` (`stream=true` + SSE-парсер) + callback с `gr.update(...)` для прогресс-бара.
- **Per-mode `n` default** — решено: 4 для icon, 2 для illustration (§8.2).
- **PNG-превью лучшего крупнее** — нет, все одинаковые. Акцент через caption "★ best".
- **Regenerate** — нет в MVP. Если нужен: ~20 строк, кнопка копирует параметры в поля и дёргает `on_generate`.
- **cairosvg на Windows** — README предупреждает про GTK runtime. Не блокер MVP.
- **Cleanup старых превью** — нет, диск заполняется. В прод — TTL/джоба.
-55
View File
@@ -1,55 +0,0 @@
# Live streaming preview
Added in `feat/live-streaming` branch.
## What's new
- **Incremental SVG parser** (`incremental_svg.py`): `parse_to_valid(prefix)` turns
any partial, broken, or incomplete SVG into a renderable one. Closes open tags,
finishes unterminated attributes, truncates mid-tag junk, strips reasoning text
and markdown fences. Property-based tested against 100 random prefixes.
- **LM Studio streaming client** (`lm_client.stream_chat`): SSE consumer that
yields `StreamEvent(type="delta"|"end")` per token. Yields reasoning + content
together (qwen3.5 emits them in the same field).
- **Live UI** (`app.on_generate_live`): generator function yielding Gallery +
status updates as tokens arrive. Throttled to ~150ms between updates
(`time.monotonic()` check, skip-if-recent). Backpressure via "last-wins":
intermediate states are dropped, only the most recent SVG snapshot is shown.
Checkbox "Live-стрим" in UI defaults to ON.
## How it works
1. User clicks "Сгенерировать" with Live mode.
2. `on_generate_live` opens SSE connection to LM Studio.
3. Each `StreamEvent(type="delta")` is appended to a buffer; `parse_to_valid` is
run; the resulting SVG is rendered with resvg-py and pushed to the Gallery.
4. The throttle skips updates faster than 150ms apart, so we never queue more
than ~6-7 redraws per second.
5. On `StreamEvent(type="end")`, a final always-yielded snapshot is emitted
(regardless of throttle), the record is written to SQLite history, status
text shows "Сгенерировано N за Xс".
## Performance
- **Time to first preview:** typically 1-3 seconds (depends on the model's
"thinking" speed — qwen3.5-35b-a3b spends time in reasoning before the first
content token).
- **Updates per second:** capped at ~6-7 by the 150ms throttle.
- **PNG render time:** ~5-15ms per snapshot via resvg-py (no cairo dependency).
## Limitations
- Reasoning tokens are still rendered as part of the live preview. They appear
as raw text until the model emits a real `<svg>` tag. Acceptable for now;
stripping reasoning from the stream is a future improvement.
- `n > 1` is not supported in live mode (OpenAI streaming API only streams one
candidate at a time). UI automatically uses `n=1` when Live is on.
- Network interruptions mid-stream result in a `gr.Error` and partial preview
is discarded; the user must retry.
## Tests
- 34 tests in `tests/test_incremental_svg.py` (all passing)
- ~6 tests in `tests/test_lm_streaming.py` (all passing)
- 5+ tests in `tests/test_app.py` for live-UI behavior (all passing)
- **Total:** 183 passed, 1 skipped, 0 failed
-251
View File
@@ -1,251 +0,0 @@
# OmniSVG-Lite — Spec
> MVP-спецификация. Цель: дать пользователю веб-приложение, в котором можно
> из текста или картинки получить N кандидатов SVG в одном из двух режимов
> (icon / illustration), посмотреть PNG-превью и сохранить историю генераций.
---
## 1. Цели и не-цели
**Цель MVP:**
- Локально запускаемый Gradio UI, который ходит в LM Studio (OpenAI-compatible)
и возвращает валидный SVG + PNG-превью.
- Один пользователь, без авторизации, без облака.
- История генераций хранится в локальной SQLite.
**Не-цели (Out of Scope для MVP):**
- Редактирование SVG в браузере (paint / drag).
- Экспорт в Lottie / GIF / видео.
- Batch-генерация (CSV промптов, генерация по расписанию).
- Авторизация, multi-user, разграничение прав.
- Деплой / Docker / CI.
- Per-mode температура — единый слайдер на UI.
- Fine-tune, RLHF, оценка качества моделей.
- Поддержка моделей, отличных от LM Studio (Ollama, vLLM, OpenAI cloud) — клиент
пишется под OpenAI-compatible chat completions, но другие провайдеры не
тестируются.
---
## 2. Режимы генерации
### 2.1. Mode: `icon`
**Назначение:** маленькие, узнаваемые, плоские или с лёгкой стилизацией иконки
для UI. Допустимы монохром и 2–3 цвета. Геометрия упрощённая, без мелких деталей.
**Требования к SVG:**
- `viewBox="0 0 64 64"` (фиксировано, чтобы иконки были консистентны).
- Геометрия привязана к пиксельной сетке (целочисленные координаты ±0.5).
- Без градиентов со множеством stops, без фильтров (`feGaussianBlur` и т.п.).
- Допустимы: `<path>`, `<rect>`, `<circle>`, `<ellipse>`, `<line>`, `<polygon>`,
`<polyline>`, `<g>`, `<text>`, базовые `<linearGradient>`/`<radialGradient>`.
**Промпт-стиль (направление для пользователя):**
"outline battery", "filled magnifying glass", "settings gear, line style".
### 2.2. Mode: `illustration`
**Назначение:** полноценные сцены / объекты / паттерны. Больше деталей, больше
цветов, допустимы градиенты и простые фильтры.
**Требования к SVG:**
- `viewBox="0 0 512 512"` (фиксировано).
- Допустимы все валидные теги SVG 1.1 кроме `<script>` и `<foreignObject>`.
- Допустимы простые `<filter>` (без `<feImage>` на внешний URL).
- Палитра ограничена заявленной пользователем (если указана).
**Промпт-стиль:**
"flat illustration of a fox in a forest, autumn colors",
"isometric server room, blue and purple palette".
### 2.3. Сводная таблица
| Параметр | icon | illustration |
|----------------------|---------------------------------|----------------------------------|
| viewBox | `0 0 64 64` | `0 0 512 512` |
| Размер файла (лимит) | 32 KB | 256 KB |
| Допустимые теги | Базовые shapes + g + text | Полный SVG 1.1 минус опасные |
| Градиенты | ≤ 2 stops | без ограничений |
| Фильтры | нет | простые (без feImage на URL) |
| Текст | допустим, но не основной канал | допустим |
| Дефолтный n_candidates | 4 | 2 |
---
## 3. Типы задач
### 3.1. text-to-SVG
- Вход: `prompt: str` (1–1000 символов), `mode ∈ {icon, illustration}`,
`n_candidates: int ∈ [1, 8]`.
- Опционально: палитра (поле в UI, не обязательное).
- Модель получает только текст.
### 3.2. image-to-SVG
- Вход: `prompt: str` (1–1000 символов), `image: PIL.Image | file path`,
`mode ∈ {icon, illustration}`, `n_candidates: int ∈ [1, 8]`.
- Допустимые форматы картинки: PNG, JPEG, WEBP. Максимум 10 MB, до 4096×4096.
- Картинка передаётся в модель как `image_url` (data: URL, base64).
- В промпте добавляется явная инструкция: "Recreate the visual content of the
attached image as SVG. Do not describe, just generate the markup."
---
## 4. Поведение N кандидатов
**Логика выбора лучшего кандидата (для MVP — простая, детерминированная):**
1. Из N ответов модели берём только валидные (прошедшие lxml-валидацию).
2. Если валидных 0 — показываем ошибку `"no_valid_candidates"` (см. §6).
3. Каждому валидному кандидату считаем **score = размер_файла_KB** (инвертированный:
меньше = лучше). Это прокси "компактность", не качество.
4. Лучший = кандидат с минимальным score.
5. **Показ в UI:**
- `Gallery` отображает PNG-превью всех валидных кандидатов в сетке.
- Лучший помечается визуально: бейдж "★ best" над PNG.
- В выпадашке/таблице истории сохраняются все валидные SVG + пути к PNG.
**Замечание:** для MVP не используем VLM-as-judge или pairwise — это сильно
увеличит latency и стоимость. Выбор по размеру файла — намеренно простой.
---
## 5. User Stories + Acceptance Criteria
### US-1. Сгенерировать иконку по текстовому промпту
**As a** дизайнер,
**I want** ввести текстовое описание иконки и получить N вариантов SVG,
**so that** я могу быстро подобрать визуал для кнопки/меню.
**Acceptance Criteria (Given/When/Then):**
- **Given** UI запущен и LM Studio отвечает, **when** я выбираю mode=`icon`,
ввожу `prompt="filled magnifying glass"`, `n_candidates=4`, **then**:
- отправляется ровно 1 запрос в LM Studio с `n=4`,
- получаю 4 PNG-превью в Gallery (или меньше, если часть невалидна),
- у лучшего есть бейдж "★ best",
- история в SQLite содержит запись со всеми raw/validated SVG.
- **Given** промпт пустой или длиннее 1000 символов, **when** я нажимаю
Generate, **then** вижу `gr.Warning` "Prompt must be 1–1000 characters",
запрос в LM Studio НЕ отправляется.
- **Given** LM Studio не отвечает 30+ секунд, **when** идёт генерация, **then**
запрос отменяется, UI показывает ошибку "LM Studio timeout".
### US-2. Сгенерировать иллюстрацию по картинке
**As a** иллюстратор,
**I want** загрузить референс-картинку и текстовое описание, получить N SVG,
**so that** я могу быстро сделать vector-версию понравившегося эскиза.
**Acceptance Criteria:**
- **Given** UI запущен, **when** я загружаю PNG/JPEG/WEBP, ввожу
`prompt="flat illustration, blue and teal palette"`, выбираю
mode=`illustration`, `n_candidates=2`, **then**:
- картинка кодируется в base64 data: URL и отправляется в LM Studio,
- получаю до 2 PNG в Gallery,
- все raw-ответы и валидные SVG попадают в историю.
- **Given** загружена картинка > 10 MB или > 4096×4096, **when** я нажимаю
Generate, **then** UI отклоняет с ошибкой валидации до отправки в LM Studio.
### US-3. Посмотреть и скачать результат
**As a** пользователь,
**I want** увидеть PNG-превью, скачать выбранный SVG и PNG,
**so that** я могу сразу использовать результат в своих задачах.
**Acceptance Criteria:**
- **Given** генерация завершилась с ≥ 1 валидным кандидатом, **when** я смотрю
на Gallery, **then**:
- каждая картинка кликабельна,
- под/возле картинки — кнопки `Download SVG` и `Download PNG`,
- скачиваются именно те файлы, которые видны на превью (source of truth —
путь в `preview_paths`).
- **Given** я нажимаю `Download SVG`, **then** браузер получает файл с MIME
`image/svg+xml` и расширением `.svg`.
### US-4. Просмотреть историю генераций
**As a** пользователь,
**I want** видеть таблицу последних генераций и кликнуть строку, чтобы
повторно открыть результаты,
**so that** я могу вернуться к удачному варианту.
**Acceptance Criteria:**
- **Given** UI запущен, **when** я смотрю на History, **then**:
- таблица показывает последние 50 записей (id, prompt, mode, model, n,
status, created_at),
- новые записи появляются вверху,
- строки сортируются по `created_at DESC`.
- **Given** я кликаю строку в History, **when** выбираю запись, **then**
Gallery обновляется превью из `preview_paths` этой записи, raw SVG
подгружается в отдельный code-block для копирования.
- **Given** приложение перезапущено, **when** я открываю History, **then**
данные сохранились (SQLite на диске).
### US-5. Контроль параметров (температура, n)
**As a** пользователь,
**I want** управлять количеством кандидатов и креативностью,
**so that** балансировать скорость/стоимость и разнообразие.
**Acceptance Criteria:**
- **Given** UI запущен, **when** я двигаю `n_candidates` slider, **then**
диапазон 1–8, шаг 1, дефолт зависит от mode (icon=4, illustration=2).
- **Given** UI запущен, **when** я двигаю `temperature` slider, **then**
диапазон 0.0–1.5, шаг 0.05, дефолт 0.4.
- **Given** `temperature=0`, **when** идёт генерация, **then** отправляется
параметр `temperature=0` (детерминированный режим).
---
## 6. Edge Cases (таблица)
| # | Сценарий | Поведение |
|---|---------------------------------------------|--------------------------------------------------------------------------------------------------------|
| 1 | Пустой промпт | `gr.Warning` "Prompt must be 1–1000 characters", запрос не отправляется. |
| 2 | Промпт > 1000 символов | То же, что (1). |
| 3 | Картинка > 10 MB / > 4096×4096 / не PNG/JPEG/WEBP | Отклоняется до LM Studio, `gr.Warning` с указанием ограничения. |
| 4 | LM Studio не отвечает (timeout 30s) | `gr.Error` "LM Studio timeout at <url>", статус генерации `error`, в лог пишется traceback. |
| 5 | LM Studio вернул 5xx | `gr.Error` "LM Studio error: <code> <body>", статус `error`. |
| 6 | Модель вернула не-SVG (текст/JSON/мусор) | Парсер пытается вытащить `<svg>...</svg>` regex'ом; если не нашёл — кандидат помечается invalid, идёт в history со status=`invalid`, в Gallery не попадает. |
| 7 | SVG не well-formed XML | lxml бросает `XMLSyntaxError` → кандидат invalid, в лог `WARNING`, в history. |
| 8 | Отсутствует / кривой `viewBox` | `ValidatorError("missing_viewbox")` → invalid. |
| 9 | Внутри `<script>` или `<foreignObject>` | `ValidatorError("disallowed_tag")` → invalid. |
| 10| `http://` / `https://` в `href` / `xlink:href` / `url(...)` | `ValidatorError("external_ref")` → invalid. |
| 11| Размер файла > лимита (32 KB / 256 KB) | `ValidatorError("too_large")` → invalid. |
| 12| Все N кандидатов invalid | `gr.Info` "All N candidates failed validation. See history for details.", в history — запись со status=`all_invalid`. |
| 13| Картинка в image-to-SVG повреждена | PIL при открытии бросает `UnidentifiedImageError` → `gr.Warning` "Cannot read image", запрос не отправляется. |
| 14| Пользователь жмёт Generate повторно во время генерации | Кнопка `interactive=False` на время выполнения, повторный клик игнорируется. |
---
## 7. Не-функциональные требования
| Категория | Требование |
|----------------|------------------------------------------------------------------------------------------|
| Performance | Генерация 1 иконки (n=1) — p95 ≤ 20 секунд на одной 5090 / Ryzen 9. |
| Performance | PNG-рендер одного SVG — p95 ≤ 200 ms. |
| Privacy | Никакие промпты/картинки не уходят за пределы LM Studio на этой машине. |
| Reliability | Падение LM Studio не валит UI: ошибка показывается, остальные кнопки работают. |
| Storage | SQLite WAL mode, `data/omnisvg.db`. Превью PNG — `data/previews/<id>_<n>.png`. |
| Logging | `INFO` для старта/конца генерации, `WARNING` для invalid, `ERROR` для LM Studio errors. |
| Concurrency | MVP — однопользовательский, очередь не требуется. UI блокирует кнопку на время генерации. |
---
## 8. Открытые вопросы
- Требуется ли streaming в UI (как в LM Studio)? **Для MVP — нет**, выбран non-stream
режим: n>1 плохо сочетается со streaming, плюс парсить инкрементально тяжелее.
- Нужна ли кнопка "regenerate" с теми же параметрами? **Не вошла в MVP**, но UI
должен позволять скопировать промпт обратно в textbox.
- Нужен ли экспорт PNG в высоком разрешении (1024×1024)? **Нет**, достаточно
рендера в нативный viewBox.
-230
View File
@@ -1,230 +0,0 @@
"""SQLite-история генераций.
Один пользователь, без авторизации, одна таблица. WAL для устойчивости при
одновременной записи из UI и, например, фонового процесса.
Файл БД по умолчанию: `~/.omnisvg_lite/history.sqlite`. Можно переопределить
через `OMNISVG_DB_PATH` (env) или явно в конструкторе.
Схема расширена по сравнению с design.md: добавлены `n_requested`, `n_returned`,
`previews` (JSON list[str]), `error_reason`, `best_index`, `temperature` —
чтобы UI мог показать полную картину без дополнительных запросов.
"""
from __future__ import annotations
import json
import logging
import os
import sqlite3
import time
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import Any, Iterator
log = logging.getLogger(__name__)
# Дефолтный путь к БД — в домашнем каталоге пользователя. Это не абсолютный
# путь к проекту: пользователь может переопределить через env.
DEFAULT_DB_DIR = Path.home() / ".omnisvg_lite"
DEFAULT_DB_PATH = DEFAULT_DB_DIR / "history.sqlite"
def _resolve_db_path(db_path: str | Path | None) -> Path:
if db_path is not None:
return Path(db_path)
env = os.environ.get("OMNISVG_DB_PATH")
if env:
return Path(env)
return DEFAULT_DB_PATH
_SCHEMA = """
CREATE TABLE IF NOT EXISTS generations (
id INTEGER PRIMARY KEY AUTOINCREMENT,
created_at REAL NOT NULL,
prompt TEXT NOT NULL,
mode TEXT NOT NULL,
model TEXT NOT NULL,
n_requested INTEGER NOT NULL,
n_returned INTEGER NOT NULL,
temperature REAL NOT NULL,
status TEXT NOT NULL,
error_reason TEXT,
raw_outputs TEXT NOT NULL,
validated_outputs TEXT NOT NULL,
previews TEXT NOT NULL,
best_index INTEGER
);
CREATE INDEX IF NOT EXISTS idx_generations_created_at
ON generations(created_at DESC);
"""
@dataclass
class Record:
"""Запись о генерации.
Attributes:
prompt: пользовательский промпт.
mode: "icon" | "illustration".
model: имя модели.
n_requested: сколько кандидатов запросили.
n_returned: сколько LM Studio реально вернула.
temperature: использованная температура.
status: "ok" | "partial" | "failed".
error_reason: строка с ошибкой или None.
raw_outputs: list[str] — N сырых ответов модели.
validated_outputs: list[str] — прошедшие валидатор SVG.
previews: list[str] — пути к PNG-превью (rel или abs).
best_index: индекс лучшего кандидата в validated_outputs (или None).
created_at: time.time() (если None — заполняется в `add`).
"""
prompt: str
mode: str
model: str
n_requested: int
n_returned: int
temperature: float
status: str
error_reason: str | None
raw_outputs: list[str] = field(default_factory=list)
validated_outputs: list[str] = field(default_factory=list)
previews: list[str] = field(default_factory=list)
best_index: int | None = None
created_at: float | None = None
class History:
"""Контекст-менеджер для работы с историей генераций.
Используется так:
with History() as h:
row_id = h.add(Record(...))
for r in h.list_recent(20):
...
На каждый enter/exit открывается/закрывается соединение. Для долгоживущих
демонов Gradio это нормально: событий мало.
"""
def __init__(self, db_path: str | Path | None = None):
self.db_path = _resolve_db_path(db_path)
self._conn: sqlite3.Connection | None = None
# -- context manager ------------------------------------------------
def __enter__(self) -> "History":
self.db_path.parent.mkdir(parents=True, exist_ok=True)
self._conn = sqlite3.connect(str(self.db_path))
self._conn.row_factory = sqlite3.Row
# WAL: пишущие транзакции не блокируют читателей, читатели — писателей.
self._conn.execute("PRAGMA journal_mode=WAL")
self._conn.execute("PRAGMA synchronous=NORMAL")
self._conn.executescript(_SCHEMA)
self._conn.commit()
log.info("history: БД открыта по %s", self.db_path)
return self
def __exit__(self, exc_type, exc, tb) -> None:
if self._conn is not None:
try:
if exc_type is None:
self._conn.commit()
else:
self._conn.rollback()
finally:
self._conn.close()
self._conn = None
@property
def conn(self) -> sqlite3.Connection:
if self._conn is None:
raise RuntimeError("History используется вне контекст-менеджера")
return self._conn
# -- CRUD -----------------------------------------------------------
def add(self, record: Record) -> int:
"""Добавляет запись и возвращает её id."""
created_at = record.created_at if record.created_at is not None else time.time()
cur = self.conn.execute(
"""
INSERT INTO generations
(created_at, prompt, mode, model, n_requested, n_returned,
temperature, status, error_reason,
raw_outputs, validated_outputs, previews, best_index)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
created_at,
record.prompt,
record.mode,
record.model,
record.n_requested,
record.n_returned,
record.temperature,
record.status,
record.error_reason,
json.dumps(record.raw_outputs, ensure_ascii=False),
json.dumps(record.validated_outputs, ensure_ascii=False),
json.dumps(record.previews, ensure_ascii=False),
record.best_index,
),
)
self.conn.commit()
return int(cur.lastrowid)
def list_recent(self, limit: int = 50) -> list[dict[str, Any]]:
"""Возвращает последние `limit` записей, новые сверху."""
cur = self.conn.execute(
"""
SELECT id, created_at, prompt, mode, model, n_requested, n_returned,
temperature, status, error_reason,
raw_outputs, validated_outputs, previews, best_index
FROM generations
ORDER BY created_at DESC, id DESC
LIMIT ?
""",
(limit,),
)
return [_row_to_dict(row) for row in cur.fetchall()]
def get(self, record_id: int) -> dict[str, Any] | None:
"""Возвращает одну запись по id или None."""
cur = self.conn.execute(
"""
SELECT id, created_at, prompt, mode, model, n_requested, n_returned,
temperature, status, error_reason,
raw_outputs, validated_outputs, previews, best_index
FROM generations
WHERE id = ?
""",
(record_id,),
)
row = cur.fetchone()
if row is None:
return None
return _row_to_dict(row)
def count(self) -> int:
cur = self.conn.execute("SELECT COUNT(*) AS c FROM generations")
return int(cur.fetchone()["c"])
def _row_to_dict(row: sqlite3.Row) -> dict[str, Any]:
d = dict(row)
# JSON-поля превращаем в Python-структуры.
for key in ("raw_outputs", "validated_outputs", "previews"):
try:
d[key] = json.loads(d[key]) if d.get(key) else []
except (TypeError, ValueError):
log.warning("history: не удалось разобрать JSON в поле %s", key)
d[key] = []
return d
__all__ = ["DEFAULT_DB_PATH", "History", "Record"]
-216
View File
@@ -1,216 +0,0 @@
"""Инкрементальный SVG-парсер.
Превращает произвольный prefix потока SVG (включая невалидный / оборванный
посреди тега или атрибута) в валидный SVG, который можно отрендерить сразу.
Сценарий: LM Studio стримит SVG по токенам; мы хотим показывать превью
по мере поступления. В буфере в любой момент времени может быть:
- незакрытый открывающий тег: `<svg viewBox="0 0 64 64"`
- оборванный атрибут: `<path d="M12 32 C `
- незакрытая кавычка: `<rect fill="`
- частично написанный тег: `<path d="M12 32 C 40 20,`
- не закрытый `<svg>`: `<svg ...><circle cx="50"`
`parse_to_valid(prefix)`:
1. Находит первый `<svg` в потоке (всё до него — мусор/reasoning/fence).
2. Обрезает по первому `</svg>` (остальное — мусор/повторный SVG).
3. Делегирует lxml'у парсинг в режиме `recover=True`: lxml сам закрывает
оборванные теги, открытые кавычки и атрибуты без значения.
4. Добавляет `xmlns="http://www.w3.org/2000/svg"` в корневой `<svg>`,
если его нет.
5. Гарантирует, что результат парсится `lxml.etree.fromstring()` без ошибок.
Использование:
from incremental_svg import parse_to_valid, stream_to_valid
# Один снапшот
svg = parse_to_valid(raw_prefix)
# Стрим: получаем валидный SVG после каждого нового чанка
for snapshot in stream_to_valid(token_chunks):
render_to_png(snapshot)
"""
from __future__ import annotations
import re
from typing import Iterable, Iterator
from lxml import etree
# ---------------------------------------------------------------------------
# Константы
# ---------------------------------------------------------------------------
SVG_NS = "http://www.w3.org/2000/svg"
# Дефолтный пустой SVG (xmlns обязателен — без него некоторые рендереры падают).
DEFAULT_EMPTY_SVG = f'<svg xmlns="{SVG_NS}"></svg>'
# Находим первое вхождение `<svg` (case-insensitive). Используется только
# позиция совпадения; сам тег обрабатывается отдельно.
_SVG_START_RE = re.compile(r"<svg\b", re.IGNORECASE)
# Lazy-parser: lxml с recover=True сам залечивает оборванные теги/кавычки.
_RECOVER_PARSER = etree.XMLParser(recover=True, huge_tree=True)
# ---------------------------------------------------------------------------
# Публичное API
# ---------------------------------------------------------------------------
def parse_to_valid(prefix: str, *, ensure_root: bool = True) -> str:
"""Превращает произвольный prefix в валидный SVG.
Args:
prefix: сырая строка (возможно, неполная, оборванная, с рассуждениями
или markdown-обёрткой).
ensure_root: если True (default), результат всегда начинается с
`<svg ...>` и заканчивается `</svg>`. Если False — может вернуть
фрагмент (для специальных случаев; в т.ч. пустую строку, если
в prefix нет `<svg>`).
Returns:
Валидный SVG, который парсится через `lxml.etree.fromstring` без
исключений. Внутри используется `lxml` для парсинга и `re` для
препроцессинга.
"""
# Защитное приведение типов: None → "", нестроковые → str().
if prefix is None:
prefix = ""
if not isinstance(prefix, str):
prefix = str(prefix)
# 1. Найти первое `<svg` (case-insensitive). Если его нет — вернуть дефолт.
m = _SVG_START_RE.search(prefix)
if not m:
return DEFAULT_EMPTY_SVG if ensure_root else ""
# 2. Взять всё от первого `<svg`. Это отсекает reasoning, markdown-fence,
# "noise <svg ..." преамбулы и т.п.
text = prefix[m.start():]
# 3. Strip XML prolog / DOCTYPE, если они по какой-то причине попали сюда.
text = _strip_prolog(text)
# 4. Обрезать по первому `</svg>` (схлопывает trailing garbage, повторные
# SVG-блоки, невалидный хвост после валидного SVG).
first_close = text.find("</svg")
if first_close >= 0:
text = text[: first_close + len("</svg>")]
# 5. Попробовать распарсить в режиме recover. lxml сам залечит:
# - необструганный `<svg ...` (закроет как self-closing);
# - необструганный `<circle ...` (допишет `/>`);
# - открытую кавычку (`<rect fill="` → `<rect fill=""/>`);
# - атрибут `attr=` без значения (отбросит или сделает `attr=""`);
# - вложенные необструганные теги (закроет все уровни).
root = _safe_parse(text)
if root is None:
# Совсем сломанный ввод (recover не справился). Возвращаем дефолт.
return DEFAULT_EMPTY_SVG if ensure_root else ""
if not ensure_root:
return etree.tostring(root, encoding="unicode")
# 6. Добавить xmlns, если отсутствует.
root = _ensure_xmlns(root)
# 7. Убедиться, что root — `<svg>`. Если нет (recover иногда берёт первым
# ребёнком при сложных преамбулах) — заворачиваем.
if _local(root.tag) != "svg":
root = _wrap_in_svg(root)
root = _ensure_xmlns(root)
return etree.tostring(root, encoding="unicode")
def stream_to_valid(chunks: Iterable[str]) -> Iterator[str]:
"""Принимает Iterable стримовых чанков, возвращает Iterator валидных
SVG-snapshots после каждого нового чанка.
Каждый snapshot — валидный SVG, пригодный для рендера.
Example:
>>> list(stream_to_valid(["<svg", ' viewBox="0 0 64 64"', ">", "</svg>"]))
['<svg xmlns="..."></svg>',
'<svg xmlns="..." viewBox="0 0 64 64"></svg>',
'<svg xmlns="..." viewBox="0 0 64 64"></svg>',
'<svg xmlns="..." viewBox="0 0 64 64"></svg>']
"""
buffer = ""
for chunk in chunks:
buffer += chunk
yield parse_to_valid(buffer)
# ---------------------------------------------------------------------------
# Внутренние хелперы
# ---------------------------------------------------------------------------
def _strip_prolog(text: str) -> str:
"""Удаляет XML-декларацию `<?xml ...?>` и `<!DOCTYPE ...>` (если они есть)."""
text = re.sub(r"<\?xml[^>]*\?>", "", text)
text = re.sub(r"<!DOCTYPE[^>]*>", "", text)
return text
def _safe_parse(text: str) -> etree._Element | None:
"""Парсит text в режиме recover. Возвращает root или None при полном фейле.
recover=True сам закрывает необструганные теги/кавычки/атрибуты. На совсем
битом вводе (например, просто `<`) может вернуть None или выбросить — мы
ловим оба случая.
"""
try:
return etree.fromstring(text.encode("utf-8"), parser=_RECOVER_PARSER)
except etree.XMLSyntaxError:
return None
except Exception: # noqa: BLE001
return None
def _ensure_xmlns(root: etree._Element) -> etree._Element:
"""Гарантирует, что у root есть `xmlns="http://www.w3.org/2000/svg"`.
Если xmlns есть — возвращает root как есть. Иначе создаёт новый root с
тем же тегом, атрибутами и детьми, но с добавленным xmlns.
"""
if root.get("xmlns"):
return root
new_root = etree.Element(_local(root.tag), nsmap={None: SVG_NS})
for k, v in root.attrib.items():
new_root.set(_local(k), v)
for child in root:
new_root.append(child)
return new_root
def _wrap_in_svg(root: etree._Element) -> etree._Element:
"""Заворачивает произвольный root в `<svg xmlns="...">...</svg>`."""
new_root = etree.Element("svg", nsmap={None: SVG_NS})
for k, v in root.attrib.items():
new_root.set(_local(k), v)
for child in root:
new_root.append(child)
return new_root
def _local(tag: object) -> str:
"""Снимает namespace-префикс: `{ns}local` → `local`. Терпимо к non-str."""
s = str(tag)
if "}" in s:
return s.split("}", 1)[1]
return s
__all__ = [
"parse_to_valid",
"stream_to_valid",
"SVG_NS",
"DEFAULT_EMPTY_SVG",
]
-538
View File
@@ -1,538 +0,0 @@
"""Клиент к LM Studio через OpenAI-compatible chat completions.
Использует синхронный httpx, чтобы не тащить openai SDK и не ловить
несовместимости их версий. Поддерживает текстовые и мультимодальные
(text + image_url) сообщения, а также n>1 для batch-генерации кандидатов.
"""
from __future__ import annotations
import base64
import io
import json
import logging
import os
import time
from dataclasses import dataclass, field
from typing import Any, Iterable, Iterator
import httpx
log = logging.getLogger(__name__)
DEFAULT_BASE_URL = "http://127.0.0.1:1234/v1"
DEFAULT_MODEL = "qwen/qwen3.5-35b-a3b"
DEFAULT_API_KEY = "lm-studio" # LM Studio игнорирует значение, но заголовок обязателен
DEFAULT_TIMEOUT_S = 300.0 # qwen3.5 может долго рассуждать, не режем по таймауту
DEFAULT_MAX_TOKENS = 32768 # без жёсткого лимита: qwen3.5 ест reasoning + content в одном budget
class LMStudioUnavailable(RuntimeError):
"""LM Studio не отвечает, вернул 5xx или вернул неожиданный формат."""
@dataclass(frozen=True)
class LMTurnResult:
"""Результат одного запроса к LM Studio.
Attributes:
raw_texts: список текстов ответов — по одному на кандидат. Если модель
вернула tool_calls без `content`, элемент будет пустой строкой.
elapsed_s: время запроса в секундах.
model: фактическое имя модели, которое вернул сервер (если есть) или
то, что мы послали.
usage: словарь usage от сервера, может быть None.
finish_reasons: список finish_reason по кандидатам ("stop"/"length"/...).
"""
raw_texts: list[str]
elapsed_s: float
model: str
usage: dict | None = None
finish_reasons: list[str] = field(default_factory=list)
@dataclass(frozen=True)
class StreamEvent:
"""Событие потокового ответа LM Studio.
Attributes:
type: 'delta' — очередной кусочек текста (delta content от сервера);
'end' — стрим завершён нормально; 'error' — стрим прерван ошибкой.
content: текст delta (для type='delta') либо '' для остальных.
usage: usage-блок, который некоторые серверы шлют в последнем чанке
(либо None, если не пришёл).
model: фактическое имя модели из ответа сервера.
finish_reason: 'stop' / 'length' / 'tool_calls' / '' (если ещё не пришёл).
"""
type: str # 'delta' | 'end' | 'error'
content: str = ""
usage: dict | None = None
model: str = ""
finish_reason: str = ""
def _coerce_text_part(part: Any) -> str:
"""Достаёт текст из элемента content — поддерживает str и list[dict]."""
if isinstance(part, str):
return part
if isinstance(part, dict):
return str(part.get("text", "") or "")
return str(part)
def _normalize_messages(messages: Iterable[dict]) -> list[dict]:
"""Готовит сообщения к отправке: текст либо строкой, либо [{type:text},...].
Допускаем на входе content как str или list[dict] (text + image_url).
Возвращаем список в формате OpenAI chat completions.
"""
out: list[dict] = []
for msg in messages:
role = msg.get("role")
content = msg.get("content")
if role is None or content is None:
log.warning("пропускаю сообщение без role/content: %r", msg)
continue
if isinstance(content, str):
out.append({"role": role, "content": content})
else:
parts: list[dict] = []
for piece in content:
if isinstance(piece, str):
parts.append({"type": "text", "text": piece})
elif isinstance(piece, dict):
ptype = piece.get("type")
if ptype == "text":
parts.append({"type": "text", "text": str(piece.get("text", ""))})
elif ptype == "image_url":
url = piece.get("image_url")
if isinstance(url, dict):
url = url.get("url")
parts.append({"type": "image_url", "image_url": {"url": str(url)}})
else:
log.warning("неизвестный тип части content: %r", ptype)
out.append({"role": role, "content": parts})
return out
def chat(
*,
messages: list[dict],
model: str | None = None,
n: int = 1,
temperature: float = 0.4,
max_tokens: int = DEFAULT_MAX_TOKENS,
base_url: str | None = None,
api_key: str | None = None,
timeout_s: float | None = None,
) -> LMTurnResult:
"""Шлёт chat completion в LM Studio и возвращает N текстов ответов.
Args:
messages: список сообщений в формате OpenAI. content может быть str
или list[dict] (text/image_url).
model: имя модели; если None — берём DEFAULT_MODEL / env `DEFAULT_MODEL`.
n: количество кандидатов, 1..8.
temperature: 0.0..1.5.
max_tokens: верхняя граница длины ответа.
base_url: эндпоинт LM Studio; дефолт http://127.0.0.1:1234/v1.
api_key: Bearer-токен; дефолт "lm-studio".
timeout_s: общий таймаут httpx, дефолт 120 секунд.
Returns:
LMTurnResult с N текстами и метаданными.
Raises:
LMStudioUnavailable: при сетевых/HTTP/ошибках парсинга.
"""
if n < 1:
raise ValueError(f"n должно быть >= 1, получено {n}")
base = (base_url or os.environ.get("LM_STUDIO_BASE_URL") or DEFAULT_BASE_URL).rstrip("/")
key = api_key if api_key is not None else os.environ.get("LM_STUDIO_API_KEY", DEFAULT_API_KEY)
mdl = model or os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL)
timeout = float(
os.environ.get("REQUEST_TIMEOUT_S", str(timeout_s if timeout_s is not None else DEFAULT_TIMEOUT_S))
)
url = f"{base}/chat/completions"
payload: dict[str, Any] = {
"model": mdl,
"messages": _normalize_messages(messages),
"n": n,
"temperature": temperature,
"max_tokens": max_tokens,
"stream": False,
# попытка отключить thinking у qwen3.5; если LM Studio/модель не уважают — не страшно
"chat_template_kwargs": {"enable_thinking": False},
}
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {key}",
}
log.info(
"LM Studio → %s model=%s n=%d temp=%.2f timeout=%.0fs",
url, mdl, n, temperature, timeout,
)
started = time.monotonic()
try:
with httpx.Client(timeout=timeout) as client:
resp = client.post(url, json=payload, headers=headers)
except httpx.TimeoutException as exc:
raise LMStudioUnavailable(
f"LM Studio timeout at {url}: запрос превысил {timeout:.0f}с"
) from exc
except httpx.HTTPError as exc:
raise LMStudioUnavailable(
f"LM Studio недоступен по адресу {url}: {exc}"
) from exc
elapsed = time.monotonic() - started
if resp.status_code >= 500:
raise LMStudioUnavailable(
f"LM Studio error: {resp.status_code} {resp.text[:200]}"
)
if resp.status_code >= 400:
# 4xx — это наша ошибка (плохой запрос / модель не загружена / квота).
raise LMStudioUnavailable(
f"LM Studio вернул {resp.status_code}: {resp.text[:200]}"
)
try:
data = resp.json()
except ValueError as exc:
raise LMStudioUnavailable(
f"LM Studio вернул не-JSON: {resp.text[:200]}"
) from exc
choices = data.get("choices") or []
if not choices:
raise LMStudioUnavailable(
f"LM Studio не вернул ни одного choice: {resp.text[:200]}"
)
raw_texts: list[str] = []
finish_reasons: list[str] = []
for ch in choices:
msg = ch.get("message") or {}
content = msg.get("content")
if isinstance(content, str):
raw_texts.append(content)
elif isinstance(content, list):
raw_texts.append("".join(_coerce_text_part(p) for p in content))
elif content is None:
# Модель могла вернуть tool_calls — для нас это пустой кандидат.
raw_texts.append("")
else:
raw_texts.append(_coerce_text_part(content))
finish_reasons.append(str(ch.get("finish_reason") or ""))
# Если сервер вернул меньше choice'ов, чем n — добиваем пустыми строками.
while len(raw_texts) < n:
raw_texts.append("")
finish_reasons.append("missing")
return LMTurnResult(
raw_texts=raw_texts,
elapsed_s=elapsed,
model=str(data.get("model") or mdl),
usage=data.get("usage"),
finish_reasons=finish_reasons,
)
def stream_chat(
*,
messages: list[dict],
model: str | None = None,
n: int = 1,
temperature: float = 0.4,
max_tokens: int = DEFAULT_MAX_TOKENS,
base_url: str | None = None,
api_key: str | None = None,
timeout_s: float | None = None,
) -> Iterator[StreamEvent]:
"""Шлёт chat completion в LM Studio со stream=True и отдаёт чанки контента.
Yields:
StreamEvent(type='delta', content=...) — очередной кусочек текста.
StreamEvent(type='end', ...) — финальное событие с метаданными
(model, finish_reason, usage если сервер прислал).
Особенности:
- Reasoning-токены qwen3.5 приходят В ОСНОВНОМ `content` (как обычный
текст), мы их не отделяем — это работа incremental_svg-парсера в UI.
- OpenAI не поддерживает `n>1` в стриме. Если пользователь передал
`n>1`, логируем WARNING и идём с `n=1` в payload.
- На ошибках (HTTP 4xx/5xx, network, timeout, битый SSE) —
бросает LMStudioUnavailable. (StreamEvent(type='error') зарезервирован
на будущее, но в текущей реализации ошибки идут через raise.)
Raises:
LMStudioUnavailable при сетевых/HTTP/парсинговых ошибках.
"""
if n < 1:
raise ValueError(f"n должно быть >= 1, получено {n}")
if n > 1:
log.warning(
"stream_chat: n=%d запрошено, но в stream-режиме OpenAI не "
"поддерживает n>1 — идём с n=1",
n,
)
effective_n = 1
else:
effective_n = n
base = (base_url or os.environ.get("LM_STUDIO_BASE_URL") or DEFAULT_BASE_URL).rstrip("/")
key = api_key if api_key is not None else os.environ.get("LM_STUDIO_API_KEY", DEFAULT_API_KEY)
mdl = model or os.environ.get("DEFAULT_MODEL", DEFAULT_MODEL)
timeout = float(
os.environ.get("REQUEST_TIMEOUT_S", str(timeout_s if timeout_s is not None else DEFAULT_TIMEOUT_S))
)
url = f"{base}/chat/completions"
payload: dict[str, Any] = {
"model": mdl,
"messages": _normalize_messages(messages),
"n": effective_n,
"temperature": temperature,
"max_tokens": max_tokens,
"stream": True,
# LM Studio / qwen3.5 уважают этот флаг, чтобы не слать thinking
# отдельным reasoning_content-полем (qwen3.5 кладёт рассуждение в
# основной content, и мы не пытаемся его отделить).
"chat_template_kwargs": {"enable_thinking": False},
}
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {key}",
}
log.info(
"LM Studio (stream) → %s model=%s n=%d temp=%.2f timeout=%.0fs",
url, mdl, effective_n, temperature, timeout,
)
started = time.monotonic()
# Поля, которые аккумулируются по ходу стрима: финальный чанк часто
# содержит finish_reason/usage, а model сервер может прислать в первом чанке.
final_model = ""
final_usage: dict | None = None
final_finish_reason = ""
saw_done = False
try:
with httpx.Client(timeout=timeout) as client:
with client.stream("POST", url, json=payload, headers=headers) as resp:
# HTTP-ошибки — до чтения тела.
if resp.status_code >= 400:
# Сливаем тело для сообщения об ошибке, но не отдаём
# его в стрим.
body_preview = ""
try:
body_preview = resp.read().decode("utf-8", errors="replace")[:200]
except Exception: # noqa: BLE001
body_preview = "<no body>"
if resp.status_code >= 500:
raise LMStudioUnavailable(
f"LM Studio error: {resp.status_code} {body_preview}"
)
raise LMStudioUnavailable(
f"LM Studio вернул {resp.status_code}: {body_preview}"
)
# Читаем SSE: каждая строка — это `data: <...>` или пустая
# строка-разделитель. События разделены пустой строкой.
for raw_line in resp.iter_lines():
if not raw_line:
continue
# SSE-префикс — `data: ` (с пробелом). Без префикса — мусор.
if not raw_line.startswith("data:"):
# Может быть комментарий (`: ...`) или event/id —
# мы их игнорируем.
continue
payload_str = raw_line[len("data:"):].strip()
if payload_str == "[DONE]":
saw_done = True
break
try:
chunk = json.loads(payload_str)
except ValueError:
# Битый SSE — бросаем, как делает обычный chat().
raise LMStudioUnavailable(
f"LM Studio stream: не-JSON в SSE-чанке: {payload_str[:200]!r}"
)
# Достаём метаданные из чанка.
if isinstance(chunk, dict):
if "model" in chunk and chunk["model"]:
final_model = str(chunk["model"])
if "usage" in chunk and chunk["usage"]:
final_usage = chunk["usage"]
choices = chunk.get("choices") or []
if not choices:
# Heartbeat-чанки без choices — пропускаем.
continue
first = choices[0]
delta = first.get("delta") or {}
content = delta.get("content")
if content:
# content может прийти str или list[dict] (мультимодальный
# стрим). Склеиваем в строку.
if isinstance(content, list):
content = "".join(_coerce_text_part(p) for p in content)
yield StreamEvent(type="delta", content=str(content))
fr = first.get("finish_reason")
if fr:
final_finish_reason = str(fr)
except httpx.TimeoutException as exc:
raise LMStudioUnavailable(
f"LM Studio stream timeout at {url}: запрос превысил {timeout:.0f}с"
) from exc
except httpx.HTTPError as exc:
raise LMStudioUnavailable(
f"LM Studio stream недоступен по адресу {url}: {exc}"
) from exc
elapsed = time.monotonic() - started
if not saw_done and not final_finish_reason:
# Стрим оборвался без [DONE] и без finish_reason — считаем это
# незавершённым. Не бросаем исключение, чтобы UI мог показать
# частичный текст; помечаем финал как обрезанный.
log.warning("LM Studio stream: выход без [DONE] (elapsed=%.2fs)", elapsed)
yield StreamEvent(
type="end",
model=final_model or mdl,
usage=final_usage,
finish_reason=final_finish_reason,
)
def encode_pil_to_data_url(image: Any, *, mime: str = "image/png") -> str:
"""Кодирует PIL-картинку в data: URL для передачи в image_url.
Args:
image: объект PIL.Image.
mime: MIME-тип, по умолчанию image/png (LM Studio VLM лучше работают с PNG).
Returns:
Строка вида `data:image/png;base64,<...>`.
"""
from PIL import Image # локальный импорт, чтобы не требовать Pillow без надобности
buf = io.BytesIO()
img = image
if getattr(img, "mode", None) == "RGBA" and mime == "image/jpeg":
img = img.convert("RGB")
elif getattr(img, "mode", None) not in ("RGB", "RGBA", "L") and mime == "image/png":
img = img.convert("RGBA")
img.save(buf, format="PNG" if mime == "image/png" else "JPEG", optimize=True)
encoded = base64.b64encode(buf.getvalue()).decode("ascii")
return f"data:{mime};base64,{encoded}"
def validate_image(image: Any, *, max_bytes: int = 10 * 1024 * 1024, max_side: int = 4096) -> None:
"""Проверяет картинку перед отправкой в LM Studio (US-2, edge case #3).
Args:
image: PIL.Image.
max_bytes: максимальный размер PNG-байт после перекодирования.
max_side: максимальная сторона в пикселях.
Raises:
ValueError: с человеко-читаемой причиной на русском.
"""
from PIL import Image, UnidentifiedImageError
if image is None:
raise ValueError("изображение не передано")
if not isinstance(image, Image.Image):
raise ValueError(f"ожидался PIL.Image, получен {type(image).__name__}")
w, h = image.size
if max(w, h) > max_side:
raise ValueError(
f"изображение слишком большое: {w}x{h}, максимум {max_side}x{max_side}"
)
if getattr(image, "format", None) and image.format not in ("PNG", "JPEG", "WEBP"):
raise ValueError(
f"неподдерживаемый формат: {image.format}; допустимы PNG, JPEG, WEBP"
)
# Прикидка размера: для RGB 3 байта/пиксель + служебные.
approx = (w * h * 4) + (1024 * 64)
if approx > max_bytes:
raise ValueError(
f"изображение превышает {max_bytes // (1024*1024)} МБ после кодирования"
)
# ---------------------------------------------------------------------------
# Удобный высокоуровневый helper, чтобы UI не думал о формате messages.
# ---------------------------------------------------------------------------
def generate_svg(
prompt: str,
*,
image_b64: str | None = None,
mode: str = "icon",
n: int = 1,
temperature: float = 0.4,
model: str | None = None,
system_prompt: str | None = None,
) -> list[dict]:
"""Собирает messages, отправляет в LM Studio и возвращает сырые ответы.
Эта функция — тонкая обёртка над `chat()`, вынесенная для удобства UI.
За реальную сборку messages отвечает `prompts.build_messages`.
Args:
prompt: текст промпта пользователя.
image_b64: data: URL картинки (или None для text-to-SVG).
mode: "icon" или "illustration" — для тега system, если не задан явно.
n: число кандидатов.
temperature: 0..1.5.
model: имя модели (опц.).
system_prompt: если задан — используется как system; иначе caller должен
сам подставить его в messages (эта функция не знает про prompts/).
Returns:
Список словарей `{'svg_text': str, 'raw': str, 'finish_reason': str}`.
"""
if not (1 <= len(prompt.strip()) <= 1000):
raise ValueError("промпт должен быть от 1 до 1000 символов")
messages: list[dict] = []
if system_prompt:
messages.append({"role": "system", "content": system_prompt})
if image_b64:
user_content: Any = [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": image_b64}},
]
else:
user_content = prompt
messages.append({"role": "user", "content": user_content})
result = chat(
messages=messages,
model=model,
n=n,
temperature=temperature,
)
return [
{
"svg_text": raw, # сырой текст; SVG вытащим extract_svg в валидаторе
"raw": raw,
"finish_reason": fr,
}
for raw, fr in zip(result.raw_texts, result.finish_reasons)
]
-173
View File
@@ -1,173 +0,0 @@
"""Загрузка промпт-шаблонов и сборка messages[] для LM Studio.
Шаблоны лежат в каталоге `prompts/` рядом с этим модулем:
- system_icon.txt — system-инструкция для режима icon
- system_illustration.txt — system-инструкция для режима illustration
- few_shot_examples.txt — текстовый файл с парами USER:/ASSISTANT:,
разделёнными строками `===...===`
Мы не выдумываем промпты — берём ровно то, что лежит в файлах. Это явное
требование задачи.
"""
from __future__ import annotations
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Literal
# Корень проекта: один уровень вверх от этого файла. Надёжнее, чем os.getcwd().
_PROMPTS_DIR = Path(__file__).resolve().parent / "prompts"
Mode = Literal["icon", "illustration"]
def _read(name: str) -> str:
"""Читает текстовый файл из prompts/; падает с понятной ошибкой, если нет."""
path = _PROMPTS_DIR / name
if not path.is_file():
raise FileNotFoundError(f"промпт-шаблон не найден: {path}")
return path.read_text(encoding="utf-8").strip()
def load_system_prompt(mode: Mode) -> str:
"""Возвращает system-инструкцию для указанного режима.
Args:
mode: "icon" или "illustration".
Returns:
Содержимое system_icon.txt / system_illustration.txt.
Raises:
ValueError: если mode неизвестен.
"""
if mode == "icon":
return _read("system_icon.txt")
if mode == "illustration":
return _read("system_illustration.txt")
raise ValueError(f"неизвестный mode: {mode!r}; ожидается 'icon' или 'illustration'")
# Регулярка для разбора few_shot_examples.txt: делим по строкам из `=`.
_FEW_SHOT_SPLIT = re.compile(r"^=+\s*$", re.MULTILINE)
# Маркеры ролей в каждом блоке.
_TURN_RE = re.compile(r"^(USER|ASSISTANT):\s*\n(.*?)(?=(?:^(?:USER|ASSISTANT):\s*$)|\Z)", re.MULTILINE | re.DOTALL)
def load_few_shot() -> list[dict]:
"""Парсит `prompts/few_shot_examples.txt` и возвращает список сообщений.
Формат файла: блоки, разделённые строками `===...===`. Внутри блока —
строки `USER:` и `ASSISTANT:`, после каждой маркерной строки идёт
содержимое до следующей маркерной строки или до конца блока.
Returns:
Список `[{role, content}, ...]` в порядке USER/ASSISTANT пар.
Блоки, где не нашлось ни одной пары, пропускаются.
"""
text = _read("few_shot_examples.txt")
blocks = _FEW_SHOT_SPLIT.split(text)
messages: list[dict] = []
for block in blocks:
# Берём только содержимое блока, убираем шапки типа "EXAMPLE N — ..."
body = block.strip()
if not body:
continue
# Пропускаем шапки, идущие ДО первого USER/ASSISTANT.
for turn in _TURN_RE.finditer(body):
role_token = turn.group(1).lower()
content = turn.group(2).strip()
messages.append({"role": role_token, "content": content})
return messages
def _user_text(prompt: str, *, palette: str | None, has_image: bool) -> str:
"""Собирает финальный текст user-turn с учётом палитры и картинки.
Args:
prompt: исходный промпт.
palette: необязательная палитра.
has_image: True, если в этот же user-turn пойдёт картинка.
Returns:
Готовый текст для content[0].type == "text".
"""
parts: list[str] = []
if has_image:
parts.append(
"Recreate the visual content of the attached image as SVG. "
"Do not describe, just generate the markup."
)
parts.append(prompt.strip())
if palette:
parts.append(f"Palette: {palette.strip()}")
return "\n".join(parts)
@dataclass(frozen=True)
class BuildResult:
"""Результат сборки messages[] — для удобства дебага/тестов."""
messages: list[dict]
mode: Mode
n: int
temperature: float
def build_messages(
prompt: str,
mode: Mode,
*,
image_b64: str | None = None,
palette: str | None = None,
n: int = 1,
temperature: float = 0.4,
) -> list[dict]:
"""Собирает полный список messages для LM Studio.
Структура:
1. system
2. few-shot пары (USER/ASSISTANT) — без картинок
3. финальный USER: текст (+ опц. image_url)
Args:
prompt: пользовательский промпт.
mode: "icon" или "illustration".
image_b64: data: URL картинки для image-to-SVG.
palette: опциональная палитра (добавится как отдельная строка).
n: число кандидатов (сейчас в messages не подставляется — это параметр
API-запроса; храним в dataclass для симметрии с дизайном).
temperature: число в [0, 1.5] (тоже идёт в API, не в messages).
Returns:
Список сообщений в формате, понятном `lm_client.chat`.
"""
if not prompt or not prompt.strip():
raise ValueError("промпт пустой")
if mode not in ("icon", "illustration"):
raise ValueError(f"неизвестный mode: {mode!r}")
messages: list[dict] = [
{"role": "system", "content": load_system_prompt(mode)},
]
messages.extend(load_few_shot())
user_text = _user_text(prompt, palette=palette, has_image=bool(image_b64))
if image_b64:
content: list[dict] = [
{"type": "text", "text": user_text},
{"type": "image_url", "image_url": {"url": image_b64}},
]
else:
content = user_text
messages.append({"role": "user", "content": content})
# Чтобы лишний раз не плодить dataclass — вернём просто список.
# Поля n/temperature отдаются API-обёртке; сборка messages от них
# не зависит (мы лишь гарантируем, что они валидны).
_ = (n, temperature)
return messages
-145
View File
@@ -1,145 +0,0 @@
Each block below is one user-turn / one assistant-turn in the few-shot sequence.
The orchestrator concatenates them, then appends the real user prompt as the final user turn.
The model is expected to answer the real user prompt in the same style: a single <svg>...</svg> block (or a refusal sentence).
================================================================================
EXAMPLE 1 — icon, filled style
================================================================================
USER:
Generate a filled magnifying glass icon. Search / find affordance.
ASSISTANT:
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">
<circle cx="26" cy="26" r="14" fill="#3B82F6"/>
<circle cx="26" cy="26" r="9" fill="#FFFFFF"/>
<rect x="36" y="34" width="6" height="20" rx="3" fill="#3B82F6" transform="rotate(45 39 44)"/>
</svg>
================================================================================
EXAMPLE 2 — icon, outline style
================================================================================
USER:
Generate an outline battery icon. Empty state (just outline, no fill).
ASSISTANT:
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">
<g fill="none" stroke="#1F2937" stroke-width="3" stroke-linejoin="round" stroke-linecap="round">
<rect x="10" y="20" width="40" height="24" rx="3"/>
<rect x="52" y="28" width="4" height="8" rx="1" fill="#1F2937"/>
</g>
</svg>
================================================================================
EXAMPLE 3 — illustration, flat style
================================================================================
USER:
Flat illustration of a red fox sitting in a forest, autumn palette (orange, brown, cream). Centered.
ASSISTANT:
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 512 512" width="512" height="512">
<defs>
<linearGradient id="sky" x1="0" x2="0" y1="0" y2="1">
<stop offset="0" stop-color="#FCD9A8"/>
<stop offset="1" stop-color="#F2A65A"/>
</linearGradient>
</defs>
<rect width="512" height="512" fill="url(#sky)"/>
<!-- Ground -->
<path d="M0 380 Q256 340 512 380 L512 512 L0 512 Z" fill="#7C4A1E"/>
<!-- Distant trees -->
<g fill="#5B3618" opacity="0.7">
<polygon points="40,380 80,260 120,380"/>
<polygon points="110,380 160,240 210,380"/>
<polygon points="420,380 460,270 500,380"/>
</g>
<!-- Fox body -->
<ellipse cx="256" cy="370" rx="110" ry="60" fill="#D8552A"/>
<!-- Fox head -->
<circle cx="256" cy="290" r="70" fill="#D8552A"/>
<!-- White muzzle -->
<ellipse cx="256" cy="310" rx="35" ry="22" fill="#FFF1D6"/>
<!-- Ears -->
<polygon points="200,250 215,180 245,235" fill="#D8552A"/>
<polygon points="312,235 297,180 272,250" fill="#D8552A"/>
<polygon points="208,242 220,200 238,232" fill="#1F1B16"/>
<polygon points="304,232 292,200 280,242" fill="#1F1B16"/>
<!-- Eyes -->
<circle cx="232" cy="285" r="6" fill="#1F1B16"/>
<circle cx="280" cy="285" r="6" fill="#1F1B16"/>
<!-- Nose -->
<circle cx="256" cy="305" r="5" fill="#1F1B16"/>
<!-- Tail -->
<path d="M340 350 Q420 320 410 240 Q395 320 360 360 Z" fill="#D8552A"/>
<path d="M395 248 Q412 235 410 220 Q400 232 392 250 Z" fill="#FFF1D6"/>
<!-- Front leaves -->
<g fill="#A65A1F">
<ellipse cx="120" cy="420" rx="22" ry="9" transform="rotate(20 120 420)"/>
<ellipse cx="170" cy="440" rx="22" ry="9" transform="rotate(-15 170 440)"/>
<ellipse cx="380" cy="430" rx="22" ry="9" transform="rotate(15 380 430)"/>
<ellipse cx="430" cy="450" rx="22" ry="9" transform="rotate(-20 430 450)"/>
</g>
</svg>
================================================================================
EXAMPLE 4 — illustration, isometric
================================================================================
USER:
Isometric server room, blue and purple palette, simple shapes.
ASSISTANT:
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 512 512" width="512" height="512">
<defs>
<linearGradient id="bg" x1="0" x2="0" y1="0" y2="1">
<stop offset="0" stop-color="#1B1B3A"/>
<stop offset="1" stop-color="#2D1B5A"/>
</linearGradient>
</defs>
<rect width="512" height="512" fill="url(#bg)"/>
<!-- Floor -->
<polygon points="256,420 80,330 256,240 432,330" fill="#3F2B85"/>
<polygon points="256,420 80,330 80,360 256,450" fill="#251A4D"/>
<polygon points="256,420 432,330 432,360 256,450" fill="#1A1240"/>
<!-- Server racks (3 in a row) -->
<g>
<!-- Rack 1 -->
<polygon points="180,330 220,350 220,180 180,160" fill="#5C3FBF"/>
<polygon points="220,350 260,330 260,160 220,180" fill="#7A55E0"/>
<polygon points="180,160 220,180 260,160 220,140" fill="#9070FF"/>
<!-- Server LEDs -->
<circle cx="190" cy="200" r="2" fill="#3DFFB0"/>
<circle cx="190" cy="220" r="2" fill="#3DFFB0"/>
<circle cx="190" cy="240" r="2" fill="#FF5C8A"/>
<circle cx="230" cy="220" r="2" fill="#3DFFB0"/>
<circle cx="230" cy="260" r="2" fill="#3DFFB0"/>
</g>
<g transform="translate(60 0)">
<polygon points="180,330 220,350 220,180 180,160" fill="#5C3FBF"/>
<polygon points="220,350 260,330 260,160 220,180" fill="#7A55E0"/>
<polygon points="180,160 220,180 260,160 220,140" fill="#9070FF"/>
<circle cx="190" cy="200" r="2" fill="#3DFFB0"/>
<circle cx="190" cy="240" r="2" fill="#FF5C8A"/>
<circle cx="230" cy="240" r="2" fill="#3DFFB0"/>
</g>
<g transform="translate(120 0)">
<polygon points="180,330 220,350 220,180 180,160" fill="#5C3FBF"/>
<polygon points="220,350 260,330 260,160 220,180" fill="#7A55E0"/>
<polygon points="180,160 220,180 260,160 220,140" fill="#9070FF"/>
<circle cx="190" cy="220" r="2" fill="#3DFFB0"/>
<circle cx="230" cy="200" r="2" fill="#3DFFB0"/>
<circle cx="230" cy="280" r="2" fill="#FF5C8A"/>
</g>
</svg>
================================================================================
EXAMPLE 5 — refusal (no SVG in output)
================================================================================
USER:
Generate an SVG of a working explosive device with step-by-step assembly instructions.
ASSISTANT:
I cannot help with instructions for building dangerous devices. Please ask for a safe, non-harmful icon or illustration.
-21
View File
@@ -1,21 +0,0 @@
You are OmniSVG-Icon, a vector icon designer. Your only output is exactly one root <svg>...</svg> block. No prose, no markdown, no explanations.
HARD RULES (your SVG must satisfy ALL of them):
1. The root element is <svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">. The viewBox is exactly "0 0 64 64".
2. Use only these tags: svg, g, defs, symbol, use, path, rect, circle, ellipse, line, polygon, polyline, linearGradient, radialGradient, stop, text, tspan, title, desc.
3. NEVER use: <script>, <foreignObject>, <image>, <iframe>, <style> with @import, <animateTransform> with external triggers, or any element that references http:// or https:// URLs.
4. No raster images, no data: URLs in href or xlink:href. No external font references. No <style>@font-face</style>. Inline <style> with safe CSS is allowed.
5. Keep the design flat, geometric, and recognizable at 32x32 pixels. Geometry should snap to a 4-pixel grid (coordinates in multiples of 1, prefer multiples of 2 or 4).
6. Color palette: at most 4 distinct fill/stroke colors, chosen for clarity on white and dark backgrounds. Prefer solid fills; gradients are allowed but only with 2 stops.
7. Final output: just one <svg>...</svg> block, nothing else before or after. No code fences. No commentary.
If the user request is unsafe, disallowed, or not representable as a vector icon, respond with a single short sentence explaining why, in plain text, without any SVG.
INTERPRETATION GUIDE:
- "outline" / "line" style = stroke-based, fill="none", stroke-width=2, stroke-linecap=round, stroke-linejoin=round.
- "filled" / "solid" style = fill-based, no stroke (or stroke same as fill).
- For UI icons, keep the visual weight consistent across the figure (e.g. all strokes 2px, all corners rounded).
- "trash", "delete", "bin" all describe the same icon: a lid + tapered body.
- For ambiguous prompts, pick the most common UI interpretation.
OUTPUT LENGTH: 200-1500 characters of SVG markup. If you cannot fit, simplify — fewer shapes is better than clipped geometry.
-27
View File
@@ -1,27 +0,0 @@
You are OmniSVG-Illustration, a vector illustrator. Your only output is exactly one root <svg>...</svg> block. No prose, no markdown, no explanations.
HARD RULES (your SVG must satisfy ALL of them):
1. The root element is <svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 512 512" width="512" height="512">. The viewBox is exactly "0 0 512 512".
2. Use only SVG 1.1 tags. Allowed set: svg, g, defs, symbol, use, path, rect, circle, ellipse, line, polygon, polyline, polyline, linearGradient, radialGradient, stop, filter, feGaussianBlur, feOffset, feBlend, feMerge, feMergeNode, feFlood, feComposite, feColorMatrix, clipPath, mask, pattern, text, tspan, textPath, title, desc.
3. NEVER use: <script>, <foreignObject>, <image>, <iframe>. No external URLs anywhere (no http://, no https://, no data: in href / xlink:href / url(...) / style imports).
4. No raster images, no embedded bitmaps. No external font references. Use generic font families only (sans-serif, serif, monospace).
5. Final file size should stay under 256 KB of UTF-8 text. If your draft exceeds that, simplify paths and remove redundant attributes.
6. Final output: just one <svg>...</svg> block, nothing else before or after. No code fences. No commentary.
DESIGN GUIDE:
- Style is "flat illustration" by default unless the user specifies otherwise. Flat = solid shapes, minimal gradients, no realistic shading.
- Honor an explicit palette from the user (e.g. "blue and teal"). Without one, pick a coherent 3-6 color palette suited to the subject.
- Use a soft drop shadow via <filter> with feGaussianBlur+feOffset if it improves depth; keep filter regions tight (don't blur the whole canvas).
- Composition: subject centered, roughly occupying 60-80% of the canvas, with breathing room on all sides.
- For scenes (forest, city, room), use a horizon line or ground plane to anchor the composition.
- For characters/animals, use simplified anatomy: large head, big eyes, stylized proportions.
INTERPRETATION GUIDE:
- "isometric" = 30-degree projection; objects as parallelograms with consistent vanishing lines.
- "minimal" = max 3 colors, no gradients, no filters.
- "kawaii" = rounded corners everywhere, pastel palette, dot eyes, simple smile.
- "vintage" = muted earth tones, slight grain via opacity overlays, hand-drawn feel.
If the user request is unsafe, disallowed, or not representable as a vector illustration, respond with a single short sentence explaining why, in plain text, without any SVG.
OUTPUT LENGTH: 5000-80000 characters of SVG markup. Prefer detail and color over minimalism, but stay under the 256 KB cap.
-127
View File
@@ -1,127 +0,0 @@
"""PNG-рендер валидных SVG.
Бэкенды (по убыванию предпочтения):
1. resvg-py — self-contained Rust-движок, не требует нативных зависимостей.
Это дефолт на Windows, где cairo.dll часто отсутствует в глобальном Python.
2. cairosvg — fallback, если cairo.dll уже установлена (типично для conda).
Если ничего не работает, `render_png()` возвращает `None` и пишет WARNING.
"""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Tuple
log = logging.getLogger(__name__)
# Ленивая инициализация бэкендов.
_resvg_module = None
_resvg_error: str | None = None
_cairosvg_module = None
_cairosvg_error: str | None = None
def _get_resvg():
global _resvg_module, _resvg_error
if _resvg_module is not None:
return _resvg_module
if _resvg_error is not None:
return None
try:
import resvg_py # type: ignore
_resvg_module = resvg_py
log.info("renderer: using resvg-py")
return _resvg_module
except Exception as exc: # noqa: BLE001
_resvg_error = f"{type(exc).__name__}: {exc}"
log.warning("resvg-py недоступен: %s", _resvg_error)
return None
def _get_cairosvg():
global _cairosvg_module, _cairosvg_error
if _cairosvg_module is not None:
return _cairosvg_module
if _cairosvg_error is not None:
return None
try:
import cairosvg # type: ignore
_cairosvg_module = cairosvg
log.info("renderer: using cairosvg (fallback)")
return _cairosvg_module
except Exception as exc: # noqa: BLE001
_cairosvg_error = f"{type(exc).__name__}: {exc}"
log.warning("cairosvg недоступен: %s", _cairosvg_error)
return None
def _ensure_xmlns(svg_text: str) -> str:
"""Добавляет xmlns, если модель его не поставила (рендереры без него падают)."""
head_end = svg_text.find(">")
if head_end < 0:
return svg_text
head = svg_text[:head_end]
if "xmlns=" in head:
return svg_text
return svg_text.replace("<svg", '<svg xmlns="http://www.w3.org/2000/svg"', 1)
def render_png(svg_text: str, size: Tuple[int, int] = (512, 512)) -> bytes | None:
"""Рендерит SVG в PNG-байты. Возвращает None если ни один бэкенд не сработал.
Args:
svg_text: валидный SVG-текст.
size: целевой размер. resvg/cairopng сохраняют пропорции по viewBox.
Returns:
PNG-байты или None.
"""
if not svg_text or not svg_text.strip():
log.warning("render_png: пустой svg_text")
return None
output_width = size[0]
svg_str = _ensure_xmlns(svg_text)
# Бэкенд 1: resvg-py
resvg = _get_resvg()
if resvg is not None:
try:
return resvg.svg_to_bytes(svg_string=svg_str)
except Exception as exc: # noqa: BLE001
log.warning("resvg упал (%s); пробую cairosvg", exc)
# Бэкенд 2: cairosvg
cairosvg = _get_cairosvg()
if cairosvg is not None:
try:
return cairosvg.svg2png(
bytestring=svg_str.encode("utf-8"),
output_width=output_width,
background_color="white",
)
except Exception as exc: # noqa: BLE001
log.warning("cairosvg упал (%s); превью пропущено", exc)
log.warning("render_png: ни один бэкенд не сработал")
return None
def save_png(
png_bytes: bytes,
*,
previews_dir: str | Path,
record_id: int,
candidate_index: int,
) -> Path:
previews_dir = Path(previews_dir)
previews_dir.mkdir(parents=True, exist_ok=True)
path = previews_dir / f"{record_id}_{candidate_index}.png"
path.write_bytes(png_bytes)
return path
__all__ = ["render_png", "save_png"]
-12
View File
@@ -1,12 +0,0 @@
gradio==5.37.0
httpx>=0.27,<0.29
lxml>=5.0,<7
# PNG-рендер: resvg-py — self-contained, не требует cairo.dll;
# cairosvg остаётся опциональным fallback, если cairo уже установлена
resvg-py>=0.3.0
cairosvg>=2.7.1
Pillow>=10.0
openai>=1.40
# dev / test
pytest>=8.0
-319
View File
@@ -1,319 +0,0 @@
"""Smoke-test реального пути данных через локальный mock-LM-Server.
Поднимает HTTP-сервер на 127.0.0.1:<port> (по умолчанию 1234), который
принимает POST /v1/chat/completions и возвращает фиксированный валидный
SVG-ответ в формате OpenAI chat completions. Затем:
1) гоняет lm_client.generate_svg() через этот mock;
2) валидирует результат через validator.validate_svg();
3) рендерит PNG через renderer.render_png() (если cairo есть);
4) пишет запись в history через History() (SQLite в tmp).
Печатает PASS или FAIL с подробностями. Код выхода:
0 — PASS
1 — FAIL
Запуск:
python scripts/smoke_test.py
Переменные окружения (опц.):
SMOKE_PORT — порт mock-сервера (default 1234; будет увеличен, если занят)
OMNISVG_DB_PATH — путь к SQLite (default в tmp)
OMNISVG_PREVIEW_DIR — каталог превью (default в tmp)
ВАЖНО: скрипт НЕ зависит от запущенного LM Studio. Это именно mock-сервер,
который имитирует LM Studio достаточно, чтобы пройти через весь pipeline.
"""
from __future__ import annotations
import json
import os
import socket
import sys
import tempfile
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
# Добавляем корень проекта в sys.path
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
# ---------------------------------------------------------------------------
# Mock LM Studio
# ---------------------------------------------------------------------------
# Валидный SVG, который mock всегда возвращает. Иконка 64x64, простая.
FIXTURE_SVG = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">'
'<circle cx="32" cy="32" r="20" fill="#3B82F6"/>'
'<rect x="36" y="36" width="6" height="20" rx="3" fill="#3B82F6" '
'transform="rotate(45 39 46)"/>'
'</svg>'
)
class MockLMHandler(BaseHTTPRequestHandler):
"""Отвечает 200 OK + OpenAI-style JSON на /v1/chat/completions.
Возвращает ровно `n` choice'ов (из тела запроса), все с FIXTURE_SVG.
Если `n` не прислали — 1.
"""
def log_message(self, format, *args): # noqa: A002 — шумно молчим
return # без stdout-спама
def do_POST(self): # noqa: N802 — http.server API
length = int(self.headers.get("Content-Length", "0") or "0")
body_bytes = self.rfile.read(length) if length > 0 else b""
if self.path.endswith("/chat/completions") or self.path == "/v1/chat/completions":
# Попробуем достать n из payload (иначе 1)
n = 1
try:
req = json.loads(body_bytes.decode("utf-8"))
n = int(req.get("n", 1))
if n < 1:
n = 1
except (ValueError, json.JSONDecodeError):
pass
choices = [
{
"index": i,
"message": {"role": "assistant", "content": FIXTURE_SVG},
"finish_reason": "stop",
}
for i in range(n)
]
payload = {
"id": "chatcmpl-smoke",
"object": "chat.completion",
"created": int(time.time()),
"model": "smoke-mock",
"choices": choices,
"usage": {
"prompt_tokens": 10,
"completion_tokens": 5 * n,
"total_tokens": 10 + 5 * n,
},
}
body = json.dumps(payload).encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
return
# Любой другой путь — 404
self.send_response(404)
self.send_header("Content-Type", "application/json")
self.end_headers()
self.wfile.write(b'{"error":"not found"}')
def _find_free_port(preferred: int) -> int:
"""Возвращает preferred, если свободен; иначе 0 (= ОС выберет)."""
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
try:
s.bind(("127.0.0.1", preferred))
return preferred
except OSError:
return 0
def _start_mock_server(port: int) -> tuple[ThreadingHTTPServer, int]:
"""Поднимает mock-LM и возвращает (server, actual_port)."""
server = ThreadingHTTPServer(("127.0.0.1", port), MockLMHandler)
actual_port = server.server_address[1]
t = threading.Thread(target=server.serve_forever, name="mock-lm", daemon=True)
t.start()
return server, actual_port
# ---------------------------------------------------------------------------
# Smoke-логика
# ---------------------------------------------------------------------------
def _check(step: str, cond: bool, detail: str = "") -> bool:
"""Печатает шаг и его результат. Возвращает cond для chaining."""
status = "OK" if cond else "FAIL"
print(f" [{status}] {step}" + (f" — {detail}" if detail else ""))
return cond
def main() -> int:
print("=" * 60)
print("OmniSVG-Lite smoke test (mock LM Studio)")
print("=" * 60)
# tmp-каталоги для изоляции
tmpdir = Path(tempfile.mkdtemp(prefix="omnisvg_smoke_"))
db_path = tmpdir / "smoke.sqlite"
previews_dir = tmpdir / "previews"
previews_dir.mkdir(parents=True, exist_ok=True)
# env-инжекция, чтобы lm_client/History/app пошли в наши пути
os.environ["OMNISVG_DB_PATH"] = str(db_path)
os.environ["OMNISVG_PREVIEW_DIR"] = str(previews_dir)
# Поднимаем mock-сервер
preferred = int(os.environ.get("SMOKE_PORT", "1234"))
port = _find_free_port(preferred)
server, actual_port = _start_mock_server(port)
base_url = f"http://127.0.0.1:{actual_port}/v1"
print(f" mock server up at {base_url} (preferred port={preferred})")
all_ok = True
all_ok &= _check("mock server started", server is not None)
# Дать серверу подняться
time.sleep(0.05)
results: list[dict] = []
validated_svgs: list[str] = []
previews: list[str] = []
# 1) LM client: generate_svg()
try:
from lm_client import generate_svg, chat
# Сначала сходим напрямую через chat() — чтобы видеть transport
result = chat(
messages=[{"role": "user", "content": "a fox"}],
n=1,
temperature=0.4,
base_url=base_url,
timeout_s=5.0,
)
all_ok &= _check(
"chat() returned LMTurnResult",
len(result.raw_texts) == 1,
f"raw_texts={result.raw_texts}",
)
all_ok &= _check(
"chat() payload contains expected SVG",
FIXTURE_SVG in result.raw_texts[0],
f"first 60 chars: {result.raw_texts[0][:60]!r}",
)
# Теперь — высокоуровневый generate_svg()
results = generate_svg(
"a fox",
n=2,
temperature=0.4,
model="smoke-mock",
system_prompt="<system>icon</system>",
)
all_ok &= _check(
"generate_svg() returned 2 candidates",
len(results) == 2,
f"got {len(results)}",
)
all_ok &= _check(
"each candidate has 'svg_text' and 'finish_reason'",
all("svg_text" in r and "finish_reason" in r for r in results),
)
except Exception as exc: # noqa: BLE001
all_ok = False
print(f" [FAIL] lm_client raised: {type(exc).__name__}: {exc}")
# 2) Validator
try:
from validator import validate_svg
for r in results:
ok, reason, cleaned = validate_svg(r["svg_text"], mode="icon")
all_ok &= _check(
f"validate_svg(icon) ok={ok}",
ok,
f"reason={reason!r}, cleaned_len={len(cleaned)}",
)
if ok:
validated_svgs.append(cleaned)
except Exception as exc: # noqa: BLE001
all_ok = False
print(f" [FAIL] validator raised: {type(exc).__name__}: {exc}")
# 3) Renderer (graceful, если cairo недоступен)
try:
from renderer import render_png, save_png
for i, svg in enumerate(validated_svgs):
png = render_png(svg, size=(64, 64))
if png is None:
# cairo недоступен — это допустимо, просто отметим.
print(f" [SKIP] render_png(#{i}) → None (cairo unavailable)")
continue
path = save_png(
png,
previews_dir=previews_dir,
record_id=0,
candidate_index=i,
)
previews.append(str(path))
all_ok &= _check(
f"render_png(#{i}) → {path.name}",
path.exists() and path.stat().st_size > 0,
)
except Exception as exc: # noqa: BLE001
all_ok = False
print(f" [FAIL] renderer raised: {type(exc).__name__}: {exc}")
# 4) History
try:
from history import History, Record
with History(db_path) as h:
record_id = h.add(
Record(
prompt="a fox",
mode="icon",
model="smoke-mock",
n_requested=2,
n_returned=len(results),
temperature=0.4,
status="ok" if validated_svgs else "all_invalid",
error_reason=None,
raw_outputs=[r["raw"] for r in results],
validated_outputs=validated_svgs,
previews=previews,
best_index=0 if validated_svgs else None,
)
)
all_ok &= _check(
"History.add() returned id",
isinstance(record_id, int) and record_id > 0,
f"id={record_id}",
)
with History(db_path) as h:
rec = h.get(record_id)
all_ok &= _check(
"History.get() roundtrip",
rec is not None and rec["prompt"] == "a fox",
f"prompt={rec['prompt'] if rec else None!r}",
)
except Exception as exc: # noqa: BLE001
all_ok = False
print(f" [FAIL] history raised: {type(exc).__name__}: {exc}")
# 5) Cleanup
server.shutdown()
server.server_close()
print("=" * 60)
if all_ok:
print("PASS — все шаги успешны")
return 0
print("FAIL — см. подробности выше")
return 1
if __name__ == "__main__":
sys.exit(main())
-605
View File
@@ -1,605 +0,0 @@
"""Регрессионные тесты на callback-логику app.py.
Verifier feedback (attempt 1): "back crashes on input validation paths because
`gr.Warning` was changed from a class to a function in Gradio 5.x and the
producer didn't migrate. Happy path works. Unit tests don't cover this path.
Manual first-click on bad input would surface a TypeError."
Эти тесты ловят именно эту ошибку. Они НЕ дёргают Gradio UI — только
вызывают `app.on_generate` напрямую и проверяют, что:
1) Нет TypeError (т.е. внутри нет `raise gr.Warning/Error`).
2) Возвращается правильное число плейсхолдеров.
3) `app.on_generate` не пытается ходить в LM Studio, если входные данные
отклонены на pre-check.
Запуск: `python -m pytest tests/test_app.py -v`
"""
from __future__ import annotations
import sys
import time
from pathlib import Path
from typing import Any
from unittest.mock import patch
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
# ---------------------------------------------------------------------------
# Сторож: gr.Warning / gr.Error в Gradio 5.x — это функции, а не классы.
# Если кто-то когда-то обновит gradio и это поведение изменится — тест
# напомнит, что нужно пересмотреть on_generate.
# ---------------------------------------------------------------------------
def test_gradio_warning_is_function_not_class():
import inspect
import gradio as gr
# Контрактное свойство Gradio 5.x, на которое опирается on_generate:
# `gr.Warning(...)` — это ФУНКЦИЯ (а не класс исключения), и её нужно
# ВЫЗЫВАТЬ. Если в новой версии Gradio это поведение изменится, тест
# упадёт, и on_generate нужно будет пересмотреть.
assert inspect.isclass(gr.Warning) is False, (
"gr.Warning стал классом в этой версии Gradio — пересмотрите on_generate"
)
assert callable(gr.Warning)
# ---------------------------------------------------------------------------
# Хелпер: дёрнуть on_generate с разными входами и поймать TypeError.
# ---------------------------------------------------------------------------
def _call_on_generate(**overrides: Any) -> Any:
"""Вызывает app.on_generate с минимальным валидным набором + overrides.
Возвращает то, что вернул callback. Если внутри есть `raise gr.Warning`,
получим TypeError ещё ДО того, как вернётся значение.
"""
from app import on_generate
defaults: dict[str, Any] = dict(
prompt="filled magnifying glass", # валидный
mode="icon",
n_candidates=2,
temperature=0.4,
image=None,
palette="",
model="test-model",
base_url="http://127.0.0.1:1234/v1",
api_key="lm-studio",
)
defaults.update(overrides)
# Патчим chat() так, чтобы on_generate не уходил в сеть и не упал уже
# ВНЕ pre-check. Если pre-check пропустил и chat() зовётся — мы увидим
# ValueError от mock'а, что нас устраивает (это другая ветка).
with patch("app.chat") as mock_chat:
mock_chat.side_effect = RuntimeError("chat should not be called from this test")
return on_generate(**defaults)
def test_on_generate_does_not_raise_on_empty_prompt():
"""Критический регрессионный кейс: пустой промпт → gr.Warning (НЕ raise)."""
try:
result = _call_on_generate(prompt="")
except TypeError as exc:
pytest.fail(
"on_generate упал с TypeError на пустом промпте — "
"вероятно, кто-то вернул `raise gr.Warning(...)`: "
f"{exc}"
)
# Должен вернуть 7 плейсхолдеров для outputs.
assert isinstance(result, tuple)
assert len(result) == 7, f"ожидался кортеж из 7 элементов, получено {len(result)}"
def test_on_generate_does_not_raise_on_too_long_prompt():
try:
_call_on_generate(prompt="x" * 1001)
except TypeError as exc:
pytest.fail(
f"on_generate упал с TypeError на длинном промпте: {exc}"
)
def test_on_generate_does_not_raise_on_bad_n_candidates():
try:
result = _call_on_generate(n_candidates=0)
except TypeError as exc:
pytest.fail(f"on_generate упал с TypeError на n=0: {exc}")
assert isinstance(result, tuple) and len(result) == 7
# n=99 — тоже вне диапазона
try:
_call_on_generate(n_candidates=99)
except TypeError as exc:
pytest.fail(f"on_generate упал с TypeError на n=99: {exc}")
def test_on_generate_does_not_raise_on_bad_mode():
try:
result = _call_on_generate(mode="portrait")
except TypeError as exc:
pytest.fail(f"on_generate упал с TypeError на неизвестном mode: {exc}")
assert isinstance(result, tuple) and len(result) == 7
def test_on_generate_does_not_call_chat_on_precheck_fail():
"""Если pre-check упал, chat() НЕ должен вызываться вообще."""
from app import on_generate
with patch("app.chat") as mock_chat:
on_generate(
prompt="", # упадёт на pre-check
mode="icon",
n_candidates=2,
temperature=0.4,
image=None,
palette="",
model="x",
base_url="http://127.0.0.1:1234/v1",
api_key="lm-studio",
)
assert mock_chat.call_count == 0, (
"chat() был вызван, хотя pre-check должен был остановить поток"
)
def test_on_generate_does_not_call_chat_on_bad_mode():
from app import on_generate
with patch("app.chat") as mock_chat:
on_generate(
prompt="valid",
mode="junk",
n_candidates=2,
temperature=0.4,
image=None,
palette="",
model="x",
base_url="http://127.0.0.1:1234/v1",
api_key="lm-studio",
)
assert mock_chat.call_count == 0
# ---------------------------------------------------------------------------
# Положительный smoke: на корректном входе on_generate НЕ возвращает пустоту
# (хотя в этом юнит-тесте chat() замокан → идём по ветке ошибки сборки
# промпта, а не успеха; это нормально, главное — нет TypeError).
# ---------------------------------------------------------------------------
def test_on_generate_with_valid_input_does_not_typeerror():
"""Даже когда chat() падает (замокан), pre-check не должен давать TypeError."""
from app import on_generate
with patch("app.chat") as mock_chat:
mock_chat.return_value = type("R", (), {
"raw_texts": ['<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64"><circle cx="32" cy="32" r="10" fill="red"/></svg>'],
"elapsed_s": 0.1,
"model": "test",
"usage": None,
"finish_reasons": ["stop"],
})()
try:
result = on_generate(
prompt="filled magnifying glass",
mode="icon",
n_candidates=1,
temperature=0.4,
image=None,
palette="",
model="test",
base_url="http://127.0.0.1:1234/v1",
api_key="lm-studio",
)
except TypeError as exc:
pytest.fail(f"on_generate упал с TypeError на валидном входе: {exc}")
# На валидном входе возвращается кортеж из 7 элементов.
assert isinstance(result, tuple)
assert len(result) == 7
# ---------------------------------------------------------------------------
# Smoke-тест: импорт app и build_ui() возвращает gr.Blocks.
# Не лезем в сеть, не запускаем UI.
# ---------------------------------------------------------------------------
def test_app_module_imports():
"""app.py импортируется без ошибок (все зависимости в порядке)."""
import app # noqa: F401
assert hasattr(app, "on_generate")
assert hasattr(app, "on_history_select")
assert hasattr(app, "build_ui")
assert hasattr(app, "main")
def test_build_ui_returns_gradio_blocks():
"""build_ui() возвращает gr.Blocks (smoke-тест сборки UI)."""
from app import build_ui
demo = build_ui()
# Проверяем, что это действительно gr.Blocks, а не None или что-то другое.
import gradio as gr
assert isinstance(demo, gr.Blocks), f"ожидался gr.Blocks, получено {type(demo).__name__}"
def test_on_mode_change_returns_icon_default_n():
from app import on_mode_change
update = on_mode_change("icon")
# gr.update — это dict-like объект, у него есть .value
assert update["value"] == 1
def test_on_mode_change_returns_illustration_default_n():
from app import on_mode_change
update = on_mode_change("illustration")
assert update["value"] == 1
# ---------------------------------------------------------------------------
# Live-режим (on_generate_live)
#
# Пять обязательных тестов из ТЗ live-ui:
# 1. yields ≥ 2 промежуточных превью при 5 дельтах
# 2. финальный yield содержит status "Сгенерировано" и валидный SVG
# 3. throttle режет быстрые дельты (10 за 50мс → ≤ 4-5 yield'ов)
# 4. после end event в SQLite появляется запись (status=ok/partial)
# 5. при LMStudioUnavailable вызывается gr.Error
#
# Плюс sanity-тест: on_generate_live — генератор-функция.
# ---------------------------------------------------------------------------
# Реалистичный поток дельт, который последовательно строит валидный SVG.
# Используется в тестах 1, 2, 4 — здесь важно, чтобы parse_to_valid давал
# рендерабельный SVG на большинстве промежуточных стадий, а не только в конце.
_VALID_SVG_DELTAS = [
"<svg ",
'xmlns="http://www.w3.org/2000/svg" ',
'viewBox="0 0 64 64">',
'<rect x="10" y="10" width="44" height="44" fill="red"/>',
"</svg>",
]
def _make_stream_mock(deltas, *, sleep_s: float = 0.20):
"""Возвращает мок-функцию stream_chat, отдающую дельты + 'end'.
Args:
deltas: список строк-дельта.
sleep_s: пауза между дельтами (default 200мс) — чтобы throttle
не скипнул промежуточные yield'ы в "нормальном" сценарии.
"""
from lm_client import StreamEvent
def mock_stream_chat(*args, **kwargs):
for d in deltas:
if sleep_s > 0:
time.sleep(sleep_s)
yield StreamEvent(type="delta", content=d)
yield StreamEvent(type="end", model="test-model", finish_reason="stop")
return mock_stream_chat
def _call_live(**overrides):
"""Дёргает on_generate_live с разумными дефолтами + overrides.
Возвращает список yields (т.е. материализованный генератор).
"""
from app import on_generate_live
defaults: dict[str, Any] = dict(
prompt="filled magnifying glass",
mode="icon",
n_candidates=1,
temperature=0.4,
image=None,
palette="",
model="test-model",
base_url="http://127.0.0.1:1234/v1",
api_key="lm-studio",
use_live=True,
)
defaults.update(overrides)
gen = on_generate_live(**defaults)
return list(gen)
def _preview_yields(yields):
"""Возвращает только те yield'ы, в которых gallery содержит превью."""
out = []
for y in yields:
# y = (gallery, gallery_state, svg, history_df, status, paths, error)
gallery = y[0]
if gallery and isinstance(gallery, list) and len(gallery) > 0:
first = gallery[0]
if isinstance(first, tuple) and len(first) >= 1 and first[0]:
out.append(y)
return out
# --- Sanity: on_generate_live — генератор-функция ---------------------------
def test_on_generate_live_is_generator_function():
"""on_generate_live должна быть генератор-функцией (содержит yield)."""
import inspect
from app import on_generate_live
assert inspect.isgeneratorfunction(on_generate_live), (
"on_generate_live должна быть генератор-функцией "
"(содержать yield) для Gradio streaming pattern"
)
# --- Тест 1: ≥ 2 промежуточных yield'а при 5 дельтах ------------------------
def test_live_mode_yields_intermediate_previews(tmp_path, monkeypatch):
"""Мок stream_chat отдаёт 5 дельт; on_generate_live делает ≥ 2 yield'а
с обновлениями Gallery (промежуточные превью)."""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
mock = _make_stream_mock(_VALID_SVG_DELTAS, sleep_s=0.20)
with patch("app.stream_chat", side_effect=mock):
yields = _call_live()
previews = _preview_yields(yields)
assert len(previews) >= 2, (
f"ожидалось ≥ 2 промежуточных preview-yield'а, получено {len(previews)}; "
f"всего yields={len(yields)}"
)
# Sanity: у промежуточных yield'ов в gallery действительно лежит файл
# (path), не просто заглушка.
for y in previews[:-1]: # все, кроме последнего (финального)
gallery = y[0]
path_str = gallery[0][0]
assert path_str.endswith(".png"), f"ожидался .png путь, получено {path_str!r}"
# --- Тест 2: финальный yield содержит "Сгенерировано" и валидный SVG --------
def test_live_mode_final_yield_after_end(tmp_path, monkeypatch):
"""Последний yield содержит status "Сгенерировано" и валидный SVG."""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
mock = _make_stream_mock(_VALID_SVG_DELTAS, sleep_s=0.20)
with patch("app.stream_chat", side_effect=mock):
yields = _call_live()
assert len(yields) >= 1
final = yields[-1]
# (gallery, gallery_state, svg, history_df, status_md, paths, error)
status_md = final[4]
final_svg = final[2]
assert "Сгенерировано" in status_md, (
f"ожидалось 'Сгенерировано' в status_md, получено {status_md!r}"
)
# Валидный SVG: парсится lxml'ом, начинается с <svg, содержит </svg>.
assert final_svg, "финальный SVG не должен быть пустым"
assert "<svg" in final_svg
assert "</svg>" in final_svg
from lxml import etree
etree.fromstring(final_svg.encode("utf-8")) # должно парситься без ошибок
# --- Тест 3: throttle режет быстрые дельты --------------------------------
def test_live_mode_throttle_skips_rapid_updates(tmp_path, monkeypatch):
"""10 дельт подряд → preview-yield'ов строго меньше, чем 10.
Чтобы изолировать throttle от скорости рендера (resvg-py + диск),
мокаем render_png на мгновенный возврат фиктивных PNG-байт.
Без throttle мы получили бы 10 preview-yield'ов; с throttle=0.15s
на быстром стриме (10 дельт за <50мс) — максимум 1-2 preview-yield'а.
"""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
# Мок render_png: возвращает фейковые PNG-байты мгновенно, чтобы
# тест измерял ТОЛЬКО throttle, а не скорость resvg/диска.
fake_png = b"\x89PNG\r\n\x1a\n" + b"\x00" * 32 # фейк-PNG заголовок
with patch("app.render_png", return_value=fake_png):
fast_deltas = [
"<svg ", # deltas 1
'xmlns="http://www.w3.org/2000/svg" ',
'viewBox="0 0 64 64">',
'<rect x="10" y="10" width="44" height="44" fill="red"/>',
'<circle cx="32" cy="32" r="5" fill="blue"/>',
'<line x1="0" y1="0" x2="64" y2="64" stroke="green"/>',
'<text x="32" y="32">A</text>',
'<ellipse cx="20" cy="20" rx="5" ry="3" fill="purple"/>',
'<polygon points="50,10 60,30 40,30" fill="orange"/>', # deltas 10
"</svg>",
]
mock = _make_stream_mock(fast_deltas, sleep_s=0.0)
started = time.monotonic()
with patch("app.stream_chat", side_effect=mock):
yields = _call_live()
elapsed_ms = (time.monotonic() - started) * 1000
previews = _preview_yields(yields)
# 10 дельт без throttle = 10 preview-yield'ов (плюс 1 начальный "поехали"
# без gallery). С throttle=0.15s и мгновенным render_png: ≤ 1-2 yields
# (первая дельта даёт, остальные скипнуты т.к. < 150мс).
# Ставим жёсткий потолок 5, чтобы тест не флакал.
assert len(previews) <= 5, (
f"throttle не сработал: {len(previews)} preview-yield'ов за {elapsed_ms:.0f}мс; "
f"предел 5"
)
# Sanity: если throttle был ВООБЩЕ выключен, было бы ~10. Проверяем,
# что превью-yield'ов сильно меньше количества дельт (= 10).
assert len(previews) < len(fast_deltas), (
f"throttle не режет: {len(previews)} превью на {len(fast_deltas)} дельт"
)
# --- Тест 4: после end event в SQLite появилась запись (ok/partial) ---------
def test_live_mode_records_in_history(tmp_path, monkeypatch):
"""После end event в SQLite появилась запись со status=ok/partial."""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
mock = _make_stream_mock(_VALID_SVG_DELTAS, sleep_s=0.20)
with patch("app.stream_chat", side_effect=mock):
yields = _call_live()
# Генерируем запись. Проверяем, что в SQLite есть новая строка.
from history import History
with History() as h:
records = h.list_recent(limit=5)
assert len(records) >= 1, "в History() нет ни одной записи после live-стрима"
last = records[0] # list_recent сортирует DESC, новая запись — первая
assert last["status"] in ("ok", "partial"), (
f"ожидался status ok/partial, получено {last['status']!r}"
)
assert last["prompt"] == "filled magnifying glass"
assert last["mode"] == "icon"
assert last["n_requested"] == 1
# raw_outputs должен содержать полный склеенный текст стрима.
assert last["raw_outputs"], "raw_outputs пуст"
assert "".join(last["raw_outputs"]).startswith("<svg ")
# --- Тест 5: LMStudioUnavailable → gr.Error ---------------------------------
def test_live_mode_handles_stream_error(tmp_path, monkeypatch):
"""Если stream_chat бросает LMStudioUnavailable, callback вызывает gr.Error
и возвращает пустой результат."""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
from lm_client import LMStudioUnavailable
def boom(*args, **kwargs):
raise LMStudioUnavailable("test: connection refused")
yield # generator-never-yield, помечает boom как генератор
with patch("app.stream_chat", side_effect=boom), \
patch("app.gr.Error") as mock_error:
yields = _call_live()
# Должен быть вызван gr.Error (без raise).
assert mock_error.called, "gr.Error не был вызван при LMStudioUnavailable"
# Должен быть ≥ 1 yield. Реально 2: первый — "Live-стрим запущен…",
# второй — _empty_result после except. Главное — последний yield
# содержит пустые плейсхолдеры (7 элементов).
assert len(yields) >= 1
last = yields[-1]
assert isinstance(last, tuple) and len(last) == 7
# Финальный gallery / status пустые.
assert last[0] == [] or last[0] is None or last[0] == ()
# И в SQLite записался failed-кейс (для аудита попыток).
from history import History
with History() as h:
records = h.list_recent(limit=5)
assert len(records) >= 1
failed = records[0]
assert failed["status"] == "failed"
assert "connection refused" in (failed["error_reason"] or "")
# --- Бонус: use_live=False — fallback на синхронный on_generate ------------
def test_live_mode_false_falls_back_to_sync(tmp_path, monkeypatch):
"""Если use_live=False, генератор делает один yield с результатом on_generate."""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
# Стрим вообще не должен вызываться.
with patch("app.stream_chat") as mock_stream, \
patch("app.chat") as mock_chat:
# chat() возвращает фиктивный результат с 1 валидным SVG.
from dataclasses import dataclass
@dataclass
class FakeResult:
raw_texts: list[str]
elapsed_s: float = 0.1
model: str = "test"
usage: dict | None = None
finish_reasons: list[str] = None
def _fr():
return ["stop"]
mock_chat.return_value = FakeResult(
raw_texts=[
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">'
'<rect x="10" y="10" width="44" height="44" fill="red"/></svg>'
],
finish_reasons=_fr(),
)
yields = _call_live(use_live=False)
assert mock_stream.call_count == 0, (
"stream_chat был вызван при use_live=False — не должен"
)
# Один yield с финальным результатом.
assert len(yields) == 1
result = yields[0]
assert isinstance(result, tuple) and len(result) == 7
# --- Бонус: pre-check fail в live-режиме ----------------------------------
def test_live_mode_precheck_fail_does_not_call_stream(tmp_path, monkeypatch):
"""Если pre-check падает (пустой промпт), stream_chat НЕ вызывается."""
monkeypatch.setenv("OMNISVG_DB_PATH", str(tmp_path / "history.sqlite"))
monkeypatch.setenv("OMNISVG_PREVIEW_DIR", str(tmp_path / "previews"))
with patch("app.stream_chat") as mock_stream:
yields = _call_live(prompt="")
assert mock_stream.call_count == 0
assert len(yields) == 1
assert len(yields[0]) == 7 # _empty_result
# --- Бонус: build_ui содержит Live-стрим checkbox -------------------------
def test_build_ui_has_live_checkbox():
"""В build_ui() должен быть gr.Checkbox с label про Live-стрим."""
from app import build_ui
demo = build_ui()
# Спускаемся по дереву компонентов в поисках Checkbox с нужным label.
# В Gradio 5.37 components живут в blocks.blocks.values().
found = False
label_seen: str = ""
for comp in demo.blocks.values():
if getattr(comp, "type", "") == "checkbox" or comp.__class__.__name__ == "Checkbox":
label = getattr(comp, "label", "") or ""
label_seen = label
if "Live" in label or "стрим" in label.lower() or "live" in label.lower():
found = True
break
assert found, (
f"не нашли gr.Checkbox с label 'Live'/'стрим' среди {len(demo.blocks)} "
f"компонентов; последний увиденный label={label_seen!r}"
)
-404
View File
@@ -1,404 +0,0 @@
"""Юнит-тесты для history.py.
Покрывают:
- add(): вставка записи возвращает id, JSON-поля сериализуются
- list_recent(): сортировка DESC, лимит, новые сверху
- get(): по id, None для несуществующего
- persist: создать/закрыть/создать — данные сохранились на диске
- count(): корректный счёт
- Спецсимволы в prompt: юникод, кавычки, переносы строк, эмодзи
- JSON-поля: list[str] с не-ASCII корректно декодируется обратно
- WAL mode: PRAGMA journal_mode после открытия соединения
- list_recent(limit=...) ограничивает выборку
Каждый тест использует tmp_path (pytest fixture) для изолированной БД.
Запуск: `python -m pytest tests/test_history.py -v`
"""
from __future__ import annotations
import sqlite3
import sys
import time
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
from history import History, Record # noqa: E402
# ---------------------------------------------------------------------------
# Базовые CRUD
# ---------------------------------------------------------------------------
def _make_record(**overrides) -> Record:
"""Минимальный валидный Record + overrides."""
base = dict(
prompt="a fox",
mode="icon",
model="test-model",
n_requested=2,
n_returned=2,
temperature=0.4,
status="ok",
error_reason=None,
raw_outputs=["raw1", "raw2"],
validated_outputs=["<svg/>", "<svg/>"],
previews=["/p/1.png", "/p/2.png"],
best_index=0,
)
base.update(overrides)
return Record(**base)
def test_add_returns_increasing_ids(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
id1 = h.add(_make_record(prompt="a"))
id2 = h.add(_make_record(prompt="b"))
id3 = h.add(_make_record(prompt="c"))
assert id1 == 1
assert id2 == 2
assert id3 == 3
assert id1 < id2 < id3
def test_get_returns_record_by_id(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
rid = h.add(_make_record(prompt="hello"))
with History(db) as h:
rec = h.get(rid)
assert rec is not None
assert rec["id"] == rid
assert rec["prompt"] == "hello"
assert rec["mode"] == "icon"
assert rec["status"] == "ok"
assert rec["n_requested"] == 2
assert rec["n_returned"] == 2
assert rec["temperature"] == 0.4
assert rec["raw_outputs"] == ["raw1", "raw2"]
assert rec["validated_outputs"] == ["<svg/>", "<svg/>"]
assert rec["previews"] == ["/p/1.png", "/p/2.png"]
assert rec["best_index"] == 0
def test_get_returns_none_for_missing_id(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
assert h.get(99999) is None
def test_count_returns_total(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
assert h.count() == 0
h.add(_make_record(prompt="a"))
h.add(_make_record(prompt="b"))
h.add(_make_record(prompt="c"))
assert h.count() == 3
# ---------------------------------------------------------------------------
# list_recent
# ---------------------------------------------------------------------------
def test_list_recent_returns_newest_first(tmp_path: Path):
"""Новые записи — сверху, лимит работает."""
db = tmp_path / "h.sqlite"
with History(db) as h:
for i in range(5):
h.add(_make_record(prompt=f"prompt-{i}"))
time.sleep(0.005) # гарантируем уникальный created_at
with History(db) as h:
rows = h.list_recent(limit=3)
assert len(rows) == 3
# Новейшие сверху → "prompt-4", "prompt-3", "prompt-2"
assert [r["prompt"] for r in rows] == ["prompt-4", "prompt-3", "prompt-2"]
def test_list_recent_default_limit_is_50(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
for i in range(60):
h.add(_make_record(prompt=f"p-{i}"))
with History(db) as h:
rows = h.list_recent()
assert len(rows) == 50
def test_list_recent_handles_explicit_large_limit(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
for i in range(3):
h.add(_make_record(prompt=f"p-{i}"))
with History(db) as h:
rows = h.list_recent(limit=1000)
assert len(rows) == 3
# ---------------------------------------------------------------------------
# Persist между открытиями
# ---------------------------------------------------------------------------
def test_data_persists_across_open_close(tmp_path: Path):
"""Создать → закрыть → снова открыть → данные на месте."""
db = tmp_path / "h.sqlite"
with History(db) as h:
rid1 = h.add(_make_record(prompt="first"))
rid2 = h.add(_make_record(prompt="second"))
# БД закрыта; открываем заново
assert db.exists()
assert db.stat().st_size > 0
with History(db) as h:
assert h.count() == 2
assert h.get(rid1)["prompt"] == "first"
assert h.get(rid2)["prompt"] == "second"
def test_data_survives_full_reopen_with_wal_files(tmp_path: Path):
"""WAL-файлы могут остаться после закрытия — повторное открытие должно работать."""
db = tmp_path / "h.sqlite"
with History(db) as h:
h.add(_make_record(prompt="x"))
# Проверяем, что нет orphan'ов: всё читается.
with History(db) as h:
rec = h.get(1)
assert rec["prompt"] == "x"
# ---------------------------------------------------------------------------
# Спецсимволы в prompt
# ---------------------------------------------------------------------------
def test_prompt_with_unicode_cyrillic(tmp_path: Path):
db = tmp_path / "h.sqlite"
cyrillic = "лиса в осеннем лесу, палитра оранжевая"
with History(db) as h:
rid = h.add(_make_record(prompt=cyrillic))
with History(db) as h:
rec = h.get(rid)
assert rec["prompt"] == cyrillic
def test_prompt_with_quotes_and_doublequotes(tmp_path: Path):
db = tmp_path / "h.sqlite"
p = 'icon with "double" and \'single\' quotes'
with History(db) as h:
rid = h.add(_make_record(prompt=p))
with History(db) as h:
rec = h.get(rid)
assert rec["prompt"] == p
def test_prompt_with_newlines_and_tabs(tmp_path: Path):
db = tmp_path / "h.sqlite"
p = "line 1\nline 2\n\tindented\n\nblank-line-above"
with History(db) as h:
rid = h.add(_make_record(prompt=p))
with History(db) as h:
rec = h.get(rid)
assert rec["prompt"] == p
def test_prompt_with_emoji(tmp_path: Path):
db = tmp_path / "h.sqlite"
p = "fox 🦊 in forest 🌲🌲"
with History(db) as h:
rid = h.add(_make_record(prompt=p))
with History(db) as h:
rec = h.get(rid)
assert rec["prompt"] == p
def test_prompt_with_backslashes_and_sql_injection_attempt(tmp_path: Path):
"""Промпт с SQL-инъекцией в виде текста — должен храниться как plain text."""
db = tmp_path / "h.sqlite"
evil = "'; DROP TABLE generations; --"
with History(db) as h:
rid = h.add(_make_record(prompt=evil))
with History(db) as h2:
rec = h2.get(rid)
assert rec["prompt"] == evil
# Таблица жива
assert h2.count() == 1
def test_prompt_with_very_long_string(tmp_path: Path):
"""Длинный промпт (10K символов) — должен сохраниться без потерь."""
db = tmp_path / "h.sqlite"
p = "x" * 10000
with History(db) as h:
rid = h.add(_make_record(prompt=p))
with History(db) as h:
rec = h.get(rid)
assert rec["prompt"] == p
assert len(rec["prompt"]) == 10000
# ---------------------------------------------------------------------------
# JSON-поля с не-ASCII
# ---------------------------------------------------------------------------
def test_raw_outputs_with_unicode_preserved(tmp_path: Path):
"""list[str] в raw_outputs хранится как JSON, ensure_ascii=False."""
db = tmp_path / "h.sqlite"
raw = [
'<svg viewBox="0 0 64 64"><text>лиса 🦊</text></svg>',
'<svg viewBox="0 0 64 64"><text>simple</text></svg>',
]
with History(db) as h:
rid = h.add(_make_record(raw_outputs=raw))
with History(db) as h:
rec = h.get(rid)
assert rec["raw_outputs"] == raw
def test_validated_outputs_with_unicode_in_svg(tmp_path: Path):
db = tmp_path / "h.sqlite"
svg = '<svg viewBox="0 0 512 512"><text>Привет</text></svg>'
with History(db) as h:
rid = h.add(_make_record(validated_outputs=[svg]))
with History(db) as h:
rec = h.get(rid)
assert rec["validated_outputs"] == [svg]
def test_previews_with_unicode_paths(tmp_path: Path):
db = tmp_path / "h.sqlite"
paths = [
"C:\\Users\\пользователь\\превью\\1_0.png",
"D:\\AI\\Projects\\лиса\\2_1.png",
]
with History(db) as h:
rid = h.add(_make_record(previews=paths))
with History(db) as h:
rec = h.get(rid)
assert rec["previews"] == paths
# ---------------------------------------------------------------------------
# Контекст-менеджер: WAL и поведение при ошибках
# ---------------------------------------------------------------------------
def test_wal_mode_is_set(tmp_path: Path):
"""После открытия History journal_mode должен быть WAL."""
db = tmp_path / "h.sqlite"
with History(db) as h:
# Спросим напрямую через raw-conn.
cur = h.conn.execute("PRAGMA journal_mode")
mode = cur.fetchone()[0]
assert mode.lower() == "wal"
def test_context_manager_closes_connection(tmp_path: Path):
db = tmp_path / "h.sqlite"
h = History(db)
with h as hist:
hist.add(_make_record(prompt="x"))
# После выхода conn=None
assert h._conn is None
def test_context_manager_commit_path_runs_in_exit(tmp_path: Path):
"""Без exception — __exit__ коммитит оставшиеся незакоммиченные изменения.
Замечание: `add()` сам вызывает commit() внутри, поэтому запись из add()
сохранится в любом случае. Этот тест проверяет, что exit/close не
портит уже закоммиченное и не падает на нормальном пути.
"""
db = tmp_path / "h.sqlite"
h = History(db)
with h as hist:
hist.add(_make_record(prompt="x"))
# Контекст закрылся без exception
assert h._conn is None
with History(db) as hist2:
assert hist2.count() == 1
def test_conn_outside_context_raises(tmp_path: Path):
db = tmp_path / "h.sqlite"
h = History(db)
with pytest.raises(RuntimeError, match="контекст-менеджера"):
_ = h.conn
# ---------------------------------------------------------------------------
# Нишевые: пустые list'ы, edge-cases значений
# ---------------------------------------------------------------------------
def test_record_with_empty_lists(tmp_path: Path):
"""Record с пустыми raw/validated/previews — должен сохраниться."""
db = tmp_path / "h.sqlite"
with History(db) as h:
rid = h.add(_make_record(
status="failed",
error_reason="lm studio down",
raw_outputs=[],
validated_outputs=[],
previews=[],
best_index=None,
))
with History(db) as h:
rec = h.get(rid)
assert rec["raw_outputs"] == []
assert rec["validated_outputs"] == []
assert rec["previews"] == []
assert rec["best_index"] is None
assert rec["error_reason"] == "lm studio down"
def test_default_created_at_set_by_add(tmp_path: Path):
"""Если created_at=None в Record — add() заполняет time.time()."""
db = tmp_path / "h.sqlite"
before = time.time()
with History(db) as h:
rec_in = _make_record(prompt="x")
assert rec_in.created_at is None
rid = h.add(rec_in)
after = time.time()
with History(db) as h:
rec_out = h.get(rid)
assert before - 1 <= rec_out["created_at"] <= after + 1
def test_explicit_created_at_preserved(tmp_path: Path):
db = tmp_path / "h.sqlite"
fixed_ts = 1700000000.0
with History(db) as h:
rid = h.add(_make_record(prompt="x", created_at=fixed_ts))
with History(db) as h:
rec = h.get(rid)
assert rec["created_at"] == fixed_ts
# ---------------------------------------------------------------------------
# Несколько записей: list_recent с лимитом < N
# ---------------------------------------------------------------------------
def test_list_recent_limit_smaller_than_total(tmp_path: Path):
db = tmp_path / "h.sqlite"
with History(db) as h:
for i in range(10):
h.add(_make_record(prompt=f"p-{i}"))
time.sleep(0.003)
with History(db) as h:
rows = h.list_recent(limit=2)
assert len(rows) == 2
# Самые свежие
assert rows[0]["prompt"] == "p-9"
assert rows[1]["prompt"] == "p-8"
-504
View File
@@ -1,504 +0,0 @@
"""Юнит-тесты для incremental_svg.py — incremental SVG parser.
Покрывают контракт `parse_to_valid(prefix)`:
- на ЛЮБОМ prefix'е возвращает валидный SVG, который парсится lxml без ошибок;
- закрывает открытые теги, атрибуты, кавычки;
- обрывает частичные теги (внутри path-d);
- вырезает reasoning / markdown fence / хвостовой мусор.
Запуск:
python -m pytest tests/test_incremental_svg.py -v
"""
from __future__ import annotations
import random
import re
import string
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
from lxml import etree # noqa: E402
from incremental_svg import parse_to_valid, stream_to_valid # noqa: E402
# ---------------------------------------------------------------------------
# Базовый полный валидный SVG, от которого будем брать префиксы
# ---------------------------------------------------------------------------
VALID_SVG = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">'
'<circle cx="32" cy="32" r="20" fill="#FF0000"/>'
'<rect x="10" y="10" width="20" height="20" fill="#00FF00"/>'
"</svg>"
)
# SVG со вложенными открытыми тегами (для теста вложенности)
NESTED_SVG = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">'
"<g>"
'<circle cx="32" cy="32" r="20" fill="red"/>'
"</g>"
"</svg>"
)
# Простая утилита — фиксирует, что результат парсится
def _assert_valid_svg(svg_text: str) -> etree._Element:
"""Парсит svg_text, возвращает root. Падает в pytest, если XML невалидный."""
try:
root = etree.fromstring(svg_text.encode("utf-8"))
except etree.XMLSyntaxError as exc:
pytest.fail(f"parse_to_valid вернул невалидный XML: {exc}\n{svg_text!r}")
return root
def _has_xmlns_in_root(svg_text: str) -> bool:
"""Проверяет, что xmlns либо в <svg ...>, либо подразумевается через сериализацию."""
head_end = svg_text.find(">")
if head_end < 0:
return False
return "xmlns" in svg_text[:head_end]
def _local_tag(tag) -> str:
"""Снимает namespace: `{ns}local` → `local`."""
s = str(tag)
if "}" in s:
return s.split("}", 1)[1]
return s
# ---------------------------------------------------------------------------
# 1. Базовые позитивные кейсы
# ---------------------------------------------------------------------------
def test_empty_string_returns_valid_empty_svg():
"""Пустая строка → валидный пустой SVG (с корневым <svg></svg>)."""
result = parse_to_valid("")
root = _assert_valid_svg(result)
assert root.tag.endswith("svg"), f"root должен быть svg, получен {root.tag}"
assert len(root) == 0 # пустой
def test_whitespace_only_returns_valid_empty_svg():
"""Строка из пробелов/переносов → валидный пустой SVG."""
result = parse_to_valid(" \n \t ")
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
def test_complete_valid_svg_returned_intact():
"""Полный валидный SVG → возвращается без изменений (и остаётся валидным)."""
result = parse_to_valid(VALID_SVG)
root = _assert_valid_svg(result)
# viewBox сохранился
assert root.get("viewBox") == "0 0 64 64"
# дети на месте
assert len(root) == 2
def test_complete_valid_svg_with_nested_open_tags():
"""Полный валидный SVG с <g>...</g> внутри → возвращается как есть."""
result = parse_to_valid(NESTED_SVG)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
# root → g → circle
g = list(root)
assert len(g) == 1
assert g[0].tag.endswith("g")
assert len(list(g[0])) == 1
# ---------------------------------------------------------------------------
# 2. Обрезанные открывающие теги
# ---------------------------------------------------------------------------
def test_svg_tag_without_closing_bracket():
"""Только `<svg viewBox="0 0 64 64"` (без `>`) → валидный пустой SVG."""
result = parse_to_valid('<svg viewBox="0 0 64 64"')
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
assert len(root) == 0
def test_svg_with_open_bracket_but_no_attributes():
"""`<svg` без атрибутов и без `>` → валидный SVG."""
result = parse_to_valid("<svg")
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
def test_svg_with_open_bracket_and_partial_attrs():
"""`<svg viewBox=` (без значения и `>`) → валидный SVG."""
result = parse_to_valid('<svg viewBox=')
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
def test_circle_with_partial_attribute_no_close():
"""`<circle cx="32"` без `>` → закрывает атрибут и тег."""
prefix = '<svg viewBox="0 0 64 64"><circle cx="32"'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
circles = [e for e in root.iter() if e.tag.endswith("circle")]
assert len(circles) == 1
assert circles[0].get("cx") == "32"
def test_circle_partial_no_value_no_close():
"""`<circle cx="32" cy=` (нет значения и нет `>`) → закрывает и атрибут, и тег."""
prefix = '<svg viewBox="0 0 64 64"><circle cx="32" cy='
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
circles = [e for e in root.iter() if e.tag.endswith("circle")]
assert len(circles) == 1
# cy должен быть пустой строкой или None
cy = circles[0].get("cy")
assert cy in (None, "")
def test_circle_open_quote_no_value():
"""`<circle cx="32" cy="` (открытая кавычка) → закрывает кавычку и тег."""
prefix = '<svg viewBox="0 0 64 64"><circle cx="32" cy="'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
circles = [e for e in root.iter() if e.tag.endswith("circle")]
assert len(circles) == 1
# ---------------------------------------------------------------------------
# 3. Обрезанный path-d
# ---------------------------------------------------------------------------
def test_path_d_truncated_no_close():
"""`<path d="M 12 32` (без `>`) → обрезает до последнего полного тега."""
prefix = '<svg viewBox="0 0 64 64"><path d="M 12 32'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
# path либо пустой (d=""), либо отсутствует, но НЕ битый XML
paths = [e for e in root.iter() if e.tag.endswith("path")]
for p in paths:
d = p.get("d") or ""
# если d сохранился — он должен быть непустой валидный path-data
# минимально: начинается с M (uppercase absolute move)
# или пустой
def test_path_d_mid_curve_truncated():
"""`<path d="M 12 32 C 40 20, 24 20, 16 16` (без `>`) → обрезает."""
prefix = '<svg viewBox="0 0 64 64"><path d="M 12 32 C 40 20, 24 20, 16 16'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
# XML валиден — это уже главное
assert root.tag.endswith("svg")
def test_path_d_truncated_keeps_xmlns_and_viewbox():
"""Даже при оборванном path-d, root-атрибуты (xmlns, viewBox) сохраняются."""
prefix = '<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64"><path d="M 0 0 L 1'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
# ---------------------------------------------------------------------------
# 4. Самозакрытые и вложенные теги
# ---------------------------------------------------------------------------
def test_self_closing_path_not_double_closed():
"""Самозакрытый `<path/>` → не закрывается повторно."""
prefix = '<svg viewBox="0 0 64 64"><path d="M 0 0 L 10 10"/>'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
paths = [e for e in root.iter() if e.tag.endswith("path")]
assert len(paths) == 1
def test_nested_open_tags_both_closed():
"""Вложенные `<g><circle/></g>` (если <g> открыт) → оба закрыты."""
prefix = '<svg viewBox="0 0 64 64"><g><circle cx="32" cy="32" r="20"'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
# Должен быть <g>, внутри — <circle>
gs = [e for e in root.iter() if _local_tag(e.tag) == "g"]
assert len(gs) >= 1
circles_in_g = [c for g in gs for c in g.iter() if _local_tag(c.tag) == "circle"]
assert len(circles_in_g) == 1
def test_multiple_open_tags_all_closed():
"""Несколько вложенных открытых тегов — все закрываются."""
prefix = '<svg viewBox="0 0 64 64"><g><g><circle cx="32"'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
# Должно парситься
assert root.tag.endswith("svg")
# ---------------------------------------------------------------------------
# 5. Префиксы с рассуждениями / markdown
# ---------------------------------------------------------------------------
def test_prefix_with_thinking_process():
"""Префикс, начинающийся с `Thinking Process:`, вырезается до `<svg`."""
prefix = (
"Thinking Process: The user wants an icon. "
"I should generate a simple circle.\n"
+ VALID_SVG
)
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
def test_prefix_with_markdown_fence():
"""Префикс с markdown fence (```xml) вырезается."""
prefix = "```xml\n" + VALID_SVG + "\n```"
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
def test_prefix_with_markdown_fence_no_lang():
"""Префикс с markdown fence без языка (```) — тоже вырезается."""
prefix = "```\n" + VALID_SVG + "\n```"
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
def test_prefix_with_trailing_garbage_after_valid_svg():
"""Валидный SVG плюс хвост мусора — возвращается только валидная часть."""
prefix = VALID_SVG + "\n\nSome trailing explanation.\nAnd more text."
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
# не должно быть trailing garbage
assert "trailing" not in result
assert "explanation" not in result
def test_prefix_with_extra_whitespace():
"""Лишний whitespace в конце SVG — нормально, валидный результат."""
result = parse_to_valid(VALID_SVG + " \n\n ")
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
# ---------------------------------------------------------------------------
# 6. Несколько SVG и спец-случаи
# ---------------------------------------------------------------------------
def test_multiple_svgs_in_prefix_takes_first():
"""Несколько `<svg>` подряд — берём только первый валидный."""
second = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 32 32">'
'<rect x="0" y="0" width="32" height="32" fill="blue"/>'
"</svg>"
)
prefix = VALID_SVG + "\n" + second
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
# Должен быть первый (viewBox 0 0 64 64)
assert root.get("viewBox") == "0 0 64 64"
def test_svg_without_xmlns_gets_xmlns_added():
"""Если в `<svg>` нет xmlns, добавляется."""
prefix = '<svg viewBox="0 0 64 64"><circle cx="32" cy="32" r="20"/></svg>'
result = parse_to_valid(prefix)
# xmlns должен появиться в head <svg ...>
assert _has_xmlns_in_root(result), f"xmlns не добавлен: {result!r}"
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
def test_svg_with_existing_xmlns_preserved():
"""Если xmlns уже есть, не дублируется."""
result = parse_to_valid(VALID_SVG)
head = result.split(">", 1)[0]
# ровно одно вхождение xmlns=
assert head.count("xmlns=") == 1
assert 'xmlns="http://www.w3.org/2000/svg"' in head
def test_unclosed_svg_at_end():
"""`<svg ...><circle cx="50"` (без `</svg>`) → добавляется закрывающий."""
prefix = '<svg viewBox="0 0 64 64"><circle cx="50"'
result = parse_to_valid(prefix)
assert result.endswith("</svg>"), f"нет закрывающего </svg>: {result!r}"
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
# ---------------------------------------------------------------------------
# 7. Property-based: 100 рандомных префиксов
# ---------------------------------------------------------------------------
def test_property_based_random_prefixes():
"""100 случайных префиксов валидного SVG → все парсятся lxml.
Берём случайный срез `VALID_SVG` (от 0 символов до всей длины) и для каждого
вызываем `parse_to_valid`. Результат должен быть валидным XML.
"""
rng = random.Random(42) # детерминированный seed для воспроизводимости
failures: list[tuple[str, str, str]] = []
for _ in range(100):
cut = rng.randint(0, len(VALID_SVG))
prefix = VALID_SVG[:cut]
# иногда добавляем "хвост" (мусор после валидной части)
if rng.random() < 0.3:
tail = "".join(rng.choices(string.ascii_letters + " \n<>", k=rng.randint(1, 20)))
prefix = prefix + tail
# иногда добавляем reasoning перед svg
if rng.random() < 0.2:
prefix = "Thinking Process: blah blah\n" + prefix
try:
result = parse_to_valid(prefix)
except Exception as exc: # noqa: BLE001
failures.append((prefix, f"exception: {exc}", ""))
continue
try:
etree.fromstring(result.encode("utf-8"))
except etree.XMLSyntaxError as exc:
failures.append((prefix, str(exc), result))
if failures:
msg = "\n".join(
f" prefix={p!r}\n err={err}\n out={out[:200]!r}"
for p, err, out in failures[:5]
)
pytest.fail(
f"{len(failures)} из 100 рандомных префиксов дали невалидный XML:\n{msg}"
)
# ---------------------------------------------------------------------------
# 8. stream_to_valid (опциональный хелпер)
# ---------------------------------------------------------------------------
def test_stream_to_valid_yields_snapshots():
"""stream_to_valid принимает Iterable стримовых чанков и возвращает Iterator
валидных SVG-snapshots."""
chunks = [
"<svg",
' viewBox="0 0 64 64"',
">",
'<circle cx="32"',
' cy="32" r="20"',
"/>",
"</svg>",
]
snapshots = list(stream_to_valid(chunks))
# хотя бы для каждого чанка должен быть валидный snapshot
assert len(snapshots) == len(chunks)
for snap in snapshots:
# Каждый snapshot — валидный SVG
root = etree.fromstring(snap.encode("utf-8"))
assert root.tag.endswith("svg")
def test_stream_to_valid_last_chunk_is_complete():
"""После последнего полного чанка stream_to_valid должен вернуть полный SVG."""
chunks = ["<svg", ' viewBox="0 0 64 64"', ">", "</svg>"]
snapshots = list(stream_to_valid(chunks))
final = snapshots[-1]
root = etree.fromstring(final.encode("utf-8"))
assert root.get("viewBox") == "0 0 64 64"
# ---------------------------------------------------------------------------
# 9. Невалидные / edge-case входы
# ---------------------------------------------------------------------------
def test_input_is_none_treated_as_empty():
"""`None` обрабатывается как пустая строка → валидный пустой SVG."""
# parse_to_valid явно требует str, но проверим, что если кто-то передал None
# (защитное программирование), мы не падаем.
try:
result = parse_to_valid(None) # type: ignore[arg-type]
except (TypeError, AttributeError):
# если тип строго str — TypeError допустим
pytest.skip("parse_to_valid строго требует str (None → TypeError)")
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
def test_ensure_root_false_returns_fragment():
"""`ensure_root=False` → результат может быть фрагментом (без обёртки в <svg>).
Если в prefix уже есть валидный <svg>...</svg>, ensure_root=False
должен вернуть его (или его фрагмент) без добавления лишнего root.
"""
result = parse_to_valid(VALID_SVG, ensure_root=False)
root = _assert_valid_svg(result)
# Должен быть валидным
assert root.tag.endswith("svg")
def test_unicode_in_svg_text():
"""Unicode в комментариях и тексте — не ломает."""
prefix = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">'
'<!-- ёжик в тумане -->'
'<text x="10" y="32">Привет, мир!</text>'
"</svg>"
)
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
def test_svg_with_doctype_stripped():
"""Если в prefix есть DOCTYPE-преамбула, она отбрасывается/чистится."""
prefix = (
'<?xml version="1.0" encoding="UTF-8"?>'
'<!DOCTYPE svg PUBLIC "-//W3C//DTD SVG 1.1//EN" "http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd">'
+ VALID_SVG
)
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.get("viewBox") == "0 0 64 64"
def test_path_d_with_numbers_only_no_letter():
"""`<path d="1234.5 6789` (без `M` и без `>`) — путь не может стартовать с числа.
Должен либо обрезать, либо как-то починить, главное — XML валидный.
"""
prefix = '<svg viewBox="0 0 64 64"><path d="1234.5 6789'
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
def test_deeply_nested_partial():
"""Глубокая вложенность частичных тегов: всё закрывается."""
prefix = (
'<svg viewBox="0 0 64 64">'
"<g>"
"<defs>"
'<linearGradient id="g1">'
'<stop offset="0" stop-color="red"'
# </linearGradient>, </defs>, </g>, </svg> — всё отсутствует
)
result = parse_to_valid(prefix)
root = _assert_valid_svg(result)
assert root.tag.endswith("svg")
-651
View File
@@ -1,651 +0,0 @@
"""Юнит-тесты для lm_client.py.
Покрывают:
- успешный ответ: парсинг N=4 choice'ов, текст, finish_reasons, model, usage
- таймаут: httpx.TimeoutException → LMStudioUnavailable с упоминанием timeout
- 5xx от LM Studio: status_code 503 → LMStudioUnavailable
- 4xx: status_code 401 / 404 → LMStudioUnavailable (по контракту: ошибка клиента
или сервера — нам всё равно)
- пустой / битый JSON
- ответ без choices → LMStudioUnavailable
- ответ с tool_use (content=None, tool_calls есть) → raw_texts содержит пустую строку
- ответ с несколькими ```svg блоками: parse_svg идёт в валидаторе, тут проверим,
что chat() возвращает всю строку модели как есть в raw_text
- generate_svg() высокоуровневая обёртка: проверка сообщений и prompt guard
- encode_pil_to_data_url(): кодирует PIL.Image в data: URL
- validate_image(): PIL.Image проходит / отклоняется по размеру / формату
httpx мокается через unittest.mock — клиент создаётся внутри `chat()`,
поэтому патчим `httpx.Client`.
Запуск: `python -m pytest tests/test_lm_client.py -v`
"""
from __future__ import annotations
import base64
import io
import sys
from pathlib import Path
from typing import Any
from unittest.mock import MagicMock, patch
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
import httpx # noqa: E402
from lm_client import ( # noqa: E402
DEFAULT_BASE_URL,
LMStudioUnavailable,
chat,
encode_pil_to_data_url,
generate_svg,
validate_image,
)
# ---------------------------------------------------------------------------
# Хелперы для построения mock-ответа httpx
# ---------------------------------------------------------------------------
def _make_response(
*,
status_code: int = 200,
json_payload: dict | None = None,
text: str = "",
) -> MagicMock:
"""Создаёт mock httpx.Response с заданным status_code и JSON-телом."""
resp = MagicMock(spec=httpx.Response)
resp.status_code = status_code
if json_payload is not None:
resp.json.return_value = json_payload
else:
resp.json.side_effect = ValueError("not json")
resp.text = text
return resp
def _ok_payload(texts: list[str], *, model: str = "test-model") -> dict:
"""Стандартный OpenAI-style ответ с N choice'ами."""
return {
"id": "chatcmpl-test",
"object": "chat.completion",
"model": model,
"choices": [
{
"index": i,
"message": {"role": "assistant", "content": t},
"finish_reason": "stop",
}
for i, t in enumerate(texts)
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
# ---------------------------------------------------------------------------
# Успешный путь
# ---------------------------------------------------------------------------
def test_chat_success_returns_n_texts_and_metadata():
"""N=4: возвращаем 4 текста, finish_reasons, model, usage."""
payload = _ok_payload(["a", "b", "c", "d"], model="my-model")
resp = _make_response(json_payload=payload)
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
result = chat(
messages=[{"role": "user", "content": "hi"}],
model="my-model",
n=4,
temperature=0.5,
base_url="http://mock:1234/v1",
)
assert isinstance(result.raw_texts, list)
assert len(result.raw_texts) == 4
assert result.raw_texts == ["a", "b", "c", "d"]
assert result.finish_reasons == ["stop"] * 4
assert result.model == "my-model"
assert result.usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert result.elapsed_s >= 0.0
# Проверяем, что URL и заголовки формируются правильно.
call = mock_inst.post.call_args
url = call.args[0] if call.args else call.kwargs["url"]
assert url == "http://mock:1234/v1/chat/completions"
headers = call.kwargs["headers"]
assert headers["Content-Type"] == "application/json"
assert headers["Authorization"] == "Bearer lm-studio"
body = call.kwargs["json"]
assert body["model"] == "my-model"
assert body["n"] == 4
assert body["temperature"] == 0.5
assert body["stream"] is False
def test_chat_success_single_candidate_default():
"""Дефолт n=1 — возвращаем ровно один текст."""
resp = _make_response(json_payload=_ok_payload(["only one"]))
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
result = chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
assert len(result.raw_texts) == 1
assert result.raw_texts[0] == "only one"
def test_chat_server_returns_fewer_choices_pads_with_empty():
"""Сервер вернул 1 из 4 — добиваем пустыми строками и 'missing'."""
payload = _ok_payload(["one"]) # n=4, но пришёл только 1
resp = _make_response(json_payload=payload)
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
result = chat(
messages=[{"role": "user", "content": "x"}],
n=4,
base_url="http://m:1/v1",
)
assert len(result.raw_texts) == 4
assert result.raw_texts[0] == "one"
assert result.raw_texts[1:] == ["", "", ""]
assert result.finish_reasons == ["stop", "missing", "missing", "missing"]
# ---------------------------------------------------------------------------
# Ошибочные пути
# ---------------------------------------------------------------------------
def test_chat_timeout_raises_lmstudio_unavailable():
"""httpx.TimeoutException → LMStudioUnavailable, в тексте есть 'timeout'."""
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.side_effect = httpx.TimeoutException("timed out")
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
n=1,
base_url="http://m:1/v1",
timeout_s=5.0,
)
msg = str(excinfo.value).lower()
assert "timeout" in msg
# Код форматирует "5.0с" / "5с" — проверим, что число таймаута попало в сообщение.
assert "5" in msg
assert "с" in str(excinfo.value) # "превысил 5с"
def test_chat_5xx_raises_lmstudio_unavailable():
resp = _make_response(status_code=503, text="service unavailable")
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
msg = str(excinfo.value)
assert "503" in msg
assert "service unavailable"[:30] in msg or "service unav" in msg
def test_chat_500_raises_lmstudio_unavailable():
resp = _make_response(status_code=500, text="internal error")
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
assert "500" in str(excinfo.value)
def test_chat_4xx_raises_lmstudio_unavailable():
"""4xx — наша ошибка, но клиент всё равно бросает LMStudioUnavailable."""
resp = _make_response(status_code=401, text="unauthorized")
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
assert "401" in str(excinfo.value)
def test_chat_network_error_raises_lmstudio_unavailable():
"""Любой httpx.HTTPError (ConnectionError и пр.) → LMStudioUnavailable."""
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.side_effect = httpx.ConnectError("connection refused")
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
msg = str(excinfo.value).lower()
assert "недоступен" in msg or "unavailable" in msg
assert "connection refused" in msg
def test_chat_invalid_json_raises_lmstudio_unavailable():
"""200 OK, но тело — не JSON."""
resp = _make_response(status_code=200, text="<html>not json</html>")
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
assert "не-JSON" in str(excinfo.value) or "json" in str(excinfo.value).lower()
def test_chat_empty_choices_raises_lmstudio_unavailable():
"""200 OK, choices пустой → ошибка."""
payload = {"choices": []}
resp = _make_response(json_payload=payload)
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
with pytest.raises(LMStudioUnavailable) as excinfo:
chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
)
assert "choice" in str(excinfo.value).lower()
# ---------------------------------------------------------------------------
# Специфические кейсы: tool_use, контент None, несколько svg-блоков
# ---------------------------------------------------------------------------
def test_chat_tool_use_response_yields_empty_string_candidate():
"""Модель вернула tool_calls без content → raw_texts содержит '' для этого кандидата."""
payload = {
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "1",
"type": "function",
"function": {"name": "x", "arguments": "{}"},
}
],
},
"finish_reason": "tool_calls",
},
{
"index": 1,
"message": {
"role": "assistant",
"content": "<svg viewBox='0 0 64 64'/>",
},
"finish_reason": "stop",
},
]
}
resp = _make_response(json_payload=payload)
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
result = chat(
messages=[{"role": "user", "content": "x"}],
n=2,
base_url="http://m:1/v1",
)
# 1-й кандидат — пусто, 2-й — реальный svg.
assert result.raw_texts[0] == ""
assert result.raw_texts[1] == "<svg viewBox='0 0 64 64'/>"
assert result.finish_reasons[0] == "tool_calls"
def test_chat_response_with_multiple_svg_blocks_kept_as_raw_text():
"""Модель вернула ответ с несколькими ```svg блоками внутри — chat() должен
сохранить текст as-is. Парсинг — забота validator.extract_svg()."""
multi_svg_text = (
"Here are some variants:\n"
"```svg\n<svg viewBox='0 0 64 64'><circle cx='10' cy='10' r='5'/></svg>\n```\n"
"And another:\n"
"```svg\n<svg viewBox='0 0 64 64'><rect width='10' height='10'/></svg>\n```"
)
resp = _make_response(json_payload=_ok_payload([multi_svg_text] * 2))
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
result = chat(
messages=[{"role": "user", "content": "x"}],
n=2,
base_url="http://m:1/v1",
)
# chat() НЕ парсит SVG — оба кандидата идентичны.
assert len(result.raw_texts) == 2
for raw in result.raw_texts:
assert raw.count("<svg") == 2
assert "Here are some variants" in raw
def test_chat_content_as_list_of_parts_joined():
"""content может быть list[dict] (мультимодальный ответ) — склеиваем текстовые части."""
payload = {
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": [
{"type": "text", "text": "Hello "},
{"type": "text", "text": "world"},
],
},
"finish_reason": "stop",
}
]
}
resp = _make_response(json_payload=payload)
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
result = chat(
messages=[{"role": "user", "content": "x"}],
n=1,
base_url="http://m:1/v1",
)
assert result.raw_texts == ["Hello world"]
# ---------------------------------------------------------------------------
# generate_svg() — высокоуровневая обёртка
# ---------------------------------------------------------------------------
def test_generate_svg_rejects_empty_prompt():
with pytest.raises(ValueError, match="1 до 1000"):
generate_svg("", model="x")
def test_generate_svg_rejects_too_long_prompt():
with pytest.raises(ValueError, match="1 до 1000"):
generate_svg("x" * 1001, model="x")
def test_generate_svg_with_image_sends_multimodal_payload():
"""С image_b64 идёт мультимодальный user-message: text + image_url."""
svg_text = "<svg viewBox='0 0 64 64'/>"
resp = _make_response(json_payload=_ok_payload([svg_text]))
with patch("lm_client.chat") as mock_chat:
mock_chat.return_value = type("R", (), {
"raw_texts": [svg_text],
"elapsed_s": 0.05,
"model": "x",
"usage": None,
"finish_reasons": ["stop"],
})()
results = generate_svg(
"a fox",
image_b64="data:image/png;base64,AAA",
mode="icon",
n=1,
temperature=0.4,
model="x",
system_prompt="<system>icon</system>",
)
assert len(results) == 1
assert results[0]["svg_text"] == svg_text
assert results[0]["raw"] == svg_text
assert results[0]["finish_reason"] == "stop"
# Что ушло в chat()
call = mock_chat.call_args
msgs = call.kwargs["messages"]
assert msgs[0]["role"] == "system"
assert msgs[0]["content"] == "<system>icon</system>"
assert msgs[1]["role"] == "user"
user_content = msgs[1]["content"]
assert isinstance(user_content, list)
assert user_content[0]["type"] == "text"
assert user_content[0]["text"] == "a fox"
assert user_content[1]["type"] == "image_url"
assert user_content[1]["image_url"]["url"] == "data:image/png;base64,AAA"
assert call.kwargs["model"] == "x"
assert call.kwargs["n"] == 1
assert call.kwargs["temperature"] == 0.4
def test_generate_svg_text_only_sends_string_content():
"""Без image_b64 user.content — просто строка."""
with patch("lm_client.chat") as mock_chat:
mock_chat.return_value = type("R", (), {
"raw_texts": ["<svg/>"],
"elapsed_s": 0.0,
"model": "x",
"usage": None,
"finish_reasons": ["stop"],
})()
generate_svg("hello", n=1, model="x")
msgs = mock_chat.call_args.kwargs["messages"]
user_msg = msgs[0] # без system_prompt единственное user-сообщение
assert user_msg["role"] == "user"
assert user_msg["content"] == "hello"
def test_generate_svg_propagates_lmstudio_error():
"""Если chat() упал, generate_svg пробрасывает LMStudioUnavailable."""
with patch("lm_client.chat") as mock_chat:
mock_chat.side_effect = LMStudioUnavailable("upstream timeout")
with pytest.raises(LMStudioUnavailable, match="upstream timeout"):
generate_svg("hello", n=1, model="x")
def test_chat_rejects_n_less_than_one():
with pytest.raises(ValueError, match="n должно быть"):
chat(messages=[{"role": "user", "content": "x"}], n=0, base_url="http://m:1/v1")
# ---------------------------------------------------------------------------
# encode_pil_to_data_url
# ---------------------------------------------------------------------------
def test_encode_pil_to_data_url_produces_data_url_with_png_mime():
from PIL import Image
img = Image.new("RGB", (10, 10), color=(255, 0, 0))
url = encode_pil_to_data_url(img, mime="image/png")
assert url.startswith("data:image/png;base64,")
# base64 должен корректно декодироваться обратно
payload = url.split(",", 1)[1]
decoded = base64.b64decode(payload)
assert decoded.startswith(b"\x89PNG")
# и это валидный PNG
Image.open(io.BytesIO(decoded))
def test_encode_pil_to_data_url_jpeg_mime_converts_rgba():
"""При mime=jpeg и RGBA → конвертируем в RGB (иначе JPEG не съест)."""
from PIL import Image
img = Image.new("RGBA", (8, 8), color=(0, 255, 0, 128))
url = encode_pil_to_data_url(img, mime="image/jpeg")
assert url.startswith("data:image/jpeg;base64,")
payload = url.split(",", 1)[1]
decoded = base64.b64decode(payload)
# JPEG стартует с FFD8
assert decoded.startswith(b"\xff\xd8")
def test_encode_pil_to_data_url_png_mime_adds_alpha_if_other_mode():
"""PNG с mime=png и mode=L (grayscale) → конвертируем в RGBA."""
from PIL import Image
img = Image.new("L", (4, 4), color=128)
url = encode_pil_to_data_url(img, mime="image/png")
assert url.startswith("data:image/png;base64,")
# Просто проверим, что получили валидный PNG
payload = url.split(",", 1)[1]
decoded = base64.b64decode(payload)
assert decoded.startswith(b"\x89PNG")
# ---------------------------------------------------------------------------
# validate_image
# ---------------------------------------------------------------------------
def test_validate_image_passes_normal_png():
from PIL import Image
img = Image.new("RGB", (100, 100), color=(0, 0, 0))
validate_image(img) # не бросает
def test_validate_image_rejects_oversized_dimensions():
from PIL import Image
img = Image.new("RGB", (5000, 100), color=(0, 0, 0))
with pytest.raises(ValueError, match="слишком большое"):
validate_image(img, max_side=4096)
def test_validate_image_rejects_too_many_bytes():
"""Картинка проходит по стороне, но approx-байты > max_bytes.
Берём 4000x4000 (под max_side=4096) и max_bytes=1MB:
4000*4000*4 + 64KB = 64MB+ >> 1MB → должно сработать байтовое ограничение.
"""
from PIL import Image
img = Image.new("RGB", (4000, 4000), color=(0, 0, 0))
with pytest.raises(ValueError, match="МБ"):
validate_image(img, max_bytes=1 * 1024 * 1024)
def test_validate_image_rejects_unsupported_format():
from PIL import Image
img = Image.new("RGB", (10, 10), color=(0, 0, 0))
img.format = "BMP" # притворяемся BMP
with pytest.raises(ValueError, match="неподдерживаемый формат"):
validate_image(img)
def test_validate_image_rejects_none():
with pytest.raises(ValueError, match="не передано"):
validate_image(None) # type: ignore[arg-type]
def test_validate_image_rejects_non_pil():
with pytest.raises(ValueError, match="ожидался PIL.Image"):
validate_image("not-an-image") # type: ignore[arg-type]
# ---------------------------------------------------------------------------
# Минорный: payload включает n/temperature/max_tokens/stream=False
# ---------------------------------------------------------------------------
def test_chat_payload_contains_required_fields():
resp = _make_response(json_payload=_ok_payload(["x"]))
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
chat(
messages=[{"role": "user", "content": "x"}],
n=2,
temperature=0.7,
max_tokens=2048,
base_url="http://m:1/v1",
)
body = mock_inst.post.call_args.kwargs["json"]
assert body["n"] == 2
assert body["temperature"] == 0.7
assert body["max_tokens"] == 2048
assert body["stream"] is False
assert body["model"] # non-empty
def test_chat_uses_default_base_url_when_env_unset(monkeypatch):
"""Если base_url=None и env не задан, идём на DEFAULT_BASE_URL."""
monkeypatch.delenv("LM_STUDIO_BASE_URL", raising=False)
resp = _make_response(json_payload=_ok_payload(["x"]))
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
mock_inst.post.return_value = resp
MockClient.return_value.__enter__.return_value = mock_inst
chat(messages=[{"role": "user", "content": "x"}], n=1)
url = mock_inst.post.call_args.args[0]
assert url == f"{DEFAULT_BASE_URL}/chat/completions"
-458
View File
@@ -1,458 +0,0 @@
"""Юнит-тесты для `stream_chat()` и `StreamEvent` в lm_client.py.
Покрывают:
- Mock SSE-ответ: 3 data-чанка + [DONE] → 3 delta-события + 1 end
- HTTP 500: stream_chat() бросает LMStudioUnavailable
- HTTP 4xx: stream_chat() бросает LMStudioUnavailable
- Timeout: stream_chat() бросает LMStudioUnavailable (с упоминанием timeout)
- Пустой стрим (только [DONE]): 0 delta'ов, 1 end
- n>1 в стриме: WARNING в логах, payload уходит с n=1
- n=0 валидация
- Модель/usage/finish_reason приходят в end из последнего чанка
- Skip не-SSE строк (комментариев/heartbeat)
- Битый JSON в SSE → LMStudioUnavailable
- reasoning-токены из content НЕ отделяются (проходят как обычный delta)
httpx мокается через unittest.mock: подменяем `httpx.Client`,
а у экземпляра — `.stream(...)` (это context manager, отдаёт response
с iter_lines).
Запуск: `python -m pytest tests/test_lm_streaming.py -v`
"""
from __future__ import annotations
import logging
import sys
from contextlib import contextmanager
from pathlib import Path
from typing import Any, Iterator
from unittest.mock import MagicMock, patch
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
import httpx # noqa: E402
from lm_client import ( # noqa: E402
LMStudioUnavailable,
StreamEvent,
stream_chat,
)
# ---------------------------------------------------------------------------
# Хелперы для построения mock-стрима
# ---------------------------------------------------------------------------
def _sse_chunk(*deltas: str, model: str = "qwen3.5", finish_reason: str | None = None) -> list[dict]:
"""Собирает фиктивный JSON-чанк в стиле OpenAI-SSE.
Каждый delta — это content, который LM Studio кладёт в
`choices[0].delta.content`. Если задан finish_reason, он попадает
в этот же чанк (как в финальном чанке OpenAI-стрима).
"""
choice: dict[str, Any] = {"index": 0, "delta": {"role": "assistant"}}
if deltas:
choice["delta"]["content"] = "".join(deltas) if len(deltas) > 1 else deltas[0]
if finish_reason:
choice["finish_reason"] = finish_reason
choice["delta"]["content"] = "".join(deltas) if deltas else ""
return [{"id": "cmpl-x", "object": "chat.completion.chunk", "model": model, "choices": [choice]}]
def _make_stream_response(
*,
sse_lines: list[str] | None = None,
sse_iter: Iterator[str] | None = None,
status_code: int = 200,
error_body: str = "",
) -> MagicMock:
"""Создаёт mock-объект, имитирующий httpx.Response внутри `client.stream(...)`.
Args:
sse_lines: фиксированный список строк (для простых случаев) — имитация
тела SSE. Каждая строка — это ровно одна строка из `iter_lines()`.
sse_iter: кастомный итератор (если хотим имитировать ошибку посреди стрима).
status_code: HTTP status.
error_body: тело для случая status_code >= 400 (читается через .read()).
"""
resp = MagicMock(spec=httpx.Response)
resp.status_code = status_code
if sse_iter is not None:
resp.iter_lines.return_value = sse_iter
else:
resp.iter_lines.return_value = sse_lines or []
if status_code >= 400:
# resp.read() вызывается в stream_chat для превью ошибки.
resp.read.return_value = error_body.encode("utf-8")
else:
resp.read.return_value = b""
return resp
@contextmanager
def _patched_stream(response: MagicMock | None = None, side_effect: Exception | None = None):
"""Патчит `httpx.Client` так, что `client.stream(...)` отдаёт заданный response.
Args:
response: mock-Response (то, что отдаёт __enter__ контекст-менеджера).
side_effect: если задан — `client.stream(...)` бросает это исключение
ДО входа в context (имитация timeout на этапе open).
"""
with patch("lm_client.httpx.Client") as MockClient:
mock_inst = MagicMock()
MockClient.return_value.__enter__.return_value = mock_inst
if side_effect is not None:
mock_inst.stream.side_effect = side_effect
else:
# .stream(...) возвращает context manager, у которого __enter__
# возвращает наш response.
cm = MagicMock()
cm.__enter__.return_value = response
cm.__exit__.return_value = False
mock_inst.stream.return_value = cm
yield mock_inst
# ---------------------------------------------------------------------------
# Успешный стрим: 3 чанка + [DONE]
# ---------------------------------------------------------------------------
def test_stream_chat_parses_three_sse_chunks():
"""SSE из 3 чанков + [DONE] → 3 StreamEvent(type='delta') + 1 end."""
sse_lines = [
'data: {"id":"1","choices":[{"index":0,"delta":{"content":"<"}}]}',
"",
'data: {"id":"2","choices":[{"index":0,"delta":{"content":"svg "}}]}',
"",
'data: {"id":"3","choices":[{"index":0,"delta":{"content":"xmlns=..."}}]}',
"",
"data: [DONE]",
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp) as mock_inst:
events = list(stream_chat(
messages=[{"role": "user", "content": "hi"}],
base_url="http://m:1/v1",
))
# 3 delta + 1 end = 4 события.
assert len(events) == 4
assert all(e.type == "delta" for e in events[:3])
assert events[3].type == "end"
# Содержимое каждого delta.
assert events[0].content == "<"
assert events[1].content == "svg "
assert events[2].content == "xmlns=..."
# Конкатенация = полный текст.
full = "".join(e.content for e in events if e.type == "delta")
assert full == "<svg xmlns=..."
# End-event: пустой content, но type='end'.
assert events[3].content == ""
assert events[3].finish_reason == ""
def test_stream_chat_payload_uses_stream_true_and_n_one():
"""Payload должен содержать stream=True и n=1."""
sse_lines = ['data: {"choices":[{"delta":{"content":"x"}}]}', "", "data: [DONE]", ""]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp) as mock_inst:
list(stream_chat(
messages=[{"role": "user", "content": "hi"}],
n=1,
temperature=0.7,
max_tokens=1024,
base_url="http://m:1/v1",
))
call = mock_inst.stream.call_args
body = call.kwargs["json"]
assert body["stream"] is True
assert body["n"] == 1
assert body["temperature"] == 0.7
assert body["max_tokens"] == 1024
assert body["model"] # non-empty
# URL сформирован правильно.
method, url = call.args[0], call.args[1]
assert method == "POST"
assert url == "http://m:1/v1/chat/completions"
headers = call.kwargs["headers"]
assert headers["Authorization"] == "Bearer lm-studio"
assert headers["Content-Type"] == "application/json"
def test_stream_chat_collects_model_and_usage_in_end_event():
"""Последний чанк содержит model/usage/finish_reason — это попадает в end."""
sse_lines = [
'data: {"model":"qwen3.5","choices":[{"delta":{"content":"Hel"}}]}',
"",
'data: {"model":"qwen3.5","choices":[{"delta":{"content":"lo"}}],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}',
"",
'data: {"choices":[{"delta":{},"finish_reason":"stop"}]}',
"",
"data: [DONE]",
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
assert len(events) == 3 # 2 delta + 1 end
end = events[-1]
assert end.type == "end"
assert end.model == "qwen3.5"
assert end.usage == {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}
assert end.finish_reason == "stop"
def test_stream_chat_handles_empty_stream_only_done():
"""Стрим сразу [DONE] → 0 delta'ов, 1 end."""
sse_lines = ["data: [DONE]", ""]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
assert len(events) == 1
assert events[0].type == "end"
assert events[0].content == ""
# ---------------------------------------------------------------------------
# Ошибки HTTP / network / timeout
# ---------------------------------------------------------------------------
def test_stream_chat_500_raises_lmstudio_unavailable():
"""HTTP 500 → LMStudioUnavailable, в сообщении есть 500 и body preview."""
resp = _make_stream_response(status_code=500, error_body="internal error")
with _patched_stream(response=resp):
with pytest.raises(LMStudioUnavailable) as excinfo:
list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
msg = str(excinfo.value)
assert "500" in msg
assert "internal error" in msg
def test_stream_chat_4xx_raises_lmstudio_unavailable():
"""HTTP 401 → LMStudioUnavailable (4xx — тоже клиентская ошибка)."""
resp = _make_stream_response(status_code=401, error_body="unauthorized")
with _patched_stream(response=resp):
with pytest.raises(LMStudioUnavailable) as excinfo:
list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
msg = str(excinfo.value)
assert "401" in msg
assert "unauthorized" in msg
def test_stream_chat_timeout_raises_lmstudio_unavailable():
"""httpx.TimeoutException на open стрима → LMStudioUnavailable с 'timeout'."""
with _patched_stream(side_effect=httpx.TimeoutException("stream timed out")):
with pytest.raises(LMStudioUnavailable) as excinfo:
list(stream_chat(
messages=[{"role": "user", "content": "x"}],
base_url="http://m:1/v1",
timeout_s=5.0,
))
msg = str(excinfo.value).lower()
assert "timeout" in msg
def test_stream_chat_network_error_raises_lmstudio_unavailable():
"""Любой httpx.HTTPError → LMStudioUnavailable."""
with _patched_stream(side_effect=httpx.ConnectError("connection refused")):
with pytest.raises(LMStudioUnavailable) as excinfo:
list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
msg = str(excinfo.value).lower()
assert "недоступен" in msg or "unavailable" in msg
assert "connection refused" in msg
def test_stream_chat_broken_sse_json_raises_lmstudio_unavailable():
"""Битый JSON в SSE-чанке → LMStudioUnavailable."""
sse_lines = [
'data: {"choices":[{"delta":{"content":"Hel"}}]}',
"",
'data: this is not json',
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
with pytest.raises(LMStudioUnavailable) as excinfo:
list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
msg = str(excinfo.value).lower()
assert "не-json" in msg or "sse" in msg
# ---------------------------------------------------------------------------
# n>1 в стриме → WARNING + n=1 в payload
# ---------------------------------------------------------------------------
def test_stream_chat_n_greater_than_one_warns_and_uses_n_one(caplog):
"""n>1 в стриме: логируется WARNING, payload уходит с n=1."""
sse_lines = ['data: {"choices":[{"delta":{"content":"x"}}]}', "", "data: [DONE]", ""]
resp = _make_stream_response(sse_lines=sse_lines)
with caplog.at_level(logging.WARNING, logger="lm_client"):
with _patched_stream(response=resp) as mock_inst:
events = list(stream_chat(
messages=[{"role": "user", "content": "x"}],
n=4,
base_url="http://m:1/v1",
))
# Payload ушёл с n=1.
body = mock_inst.stream.call_args.kwargs["json"]
assert body["n"] == 1
# В логах было WARNING.
assert any(
"n=" in rec.message and "stream" in rec.message.lower()
for rec in caplog.records if rec.levelno == logging.WARNING
)
# Стрим всё равно отработал.
assert any(e.type == "delta" for e in events)
assert events[-1].type == "end"
def test_stream_chat_rejects_n_less_than_one():
"""n=0 → ValueError до обращения к сети."""
with pytest.raises(ValueError, match="n должно быть"):
list(stream_chat(
messages=[{"role": "user", "content": "x"}],
n=0,
base_url="http://m:1/v1",
))
# ---------------------------------------------------------------------------
# Малые кейсы: пропуск мусорных строк, reasoning в content
# ---------------------------------------------------------------------------
def test_stream_chat_skips_non_data_lines():
"""Строки без префикса `data:` (комментарии, event:, id:) — пропускаются."""
sse_lines = [
": this is a comment", # SSE-комментарий
"event: message", # event-поле
"id: 42", # id-поле
'data: {"choices":[{"delta":{"content":"Hi"}}]}',
"",
"data: [DONE]",
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
deltas = [e for e in events if e.type == "delta"]
assert len(deltas) == 1
assert deltas[0].content == "Hi"
assert events[-1].type == "end"
def test_stream_chat_reasoning_tokens_pass_through_unchanged():
"""qwen3.5 кладёт reasoning в основной content — мы НЕ отделяем."""
reasoning_then_svg = (
"Let me think about this. I need an icon of a fox. The viewBox is 64x64. "
"<svg viewBox='0 0 64 64'><circle cx='32' cy='32' r='10'/></svg>"
)
# Один большой delta с reasoning+svg внутри.
sse_lines = [
f'data: {{"choices":[{{"delta":{{"content":{reasoning_then_svg!r}}}}}]}}',
"",
"data: [DONE]",
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
full = "".join(e.content for e in events if e.type == "delta")
assert "Let me think" in full
assert "<svg" in full
assert "</svg>" in full
def test_stream_chat_handles_content_as_list_of_parts():
"""content может быть list[dict] (мультимодальный стрим) — склеиваем в строку."""
sse_lines = [
'data: {"choices":[{"delta":{"content":[{"type":"text","text":"Hel"},{"type":"text","text":"lo"}]}}]}',
"",
"data: [DONE]",
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
deltas = [e for e in events if e.type == "delta"]
assert len(deltas) == 1
assert deltas[0].content == "Hello"
def test_stream_chat_emits_end_even_without_done_marker():
"""Стрим без [DONE] всё равно получает end-event (с WARNING в логе)."""
sse_lines = [
'data: {"choices":[{"delta":{"content":"x"},"finish_reason":"stop"}]}',
"",
# нет [DONE] — обрыв по исчерпанию итератора
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
assert events[-1].type == "end"
assert events[-1].finish_reason == "stop"
# ---------------------------------------------------------------------------
# Sanity: StreamEvent — frozen dataclass с нужными полями
# ---------------------------------------------------------------------------
def test_stream_event_is_frozen_dataclass():
"""StreamEvent — frozen: попытка изменения атрибута → FrozenInstanceError."""
ev = StreamEvent(type="delta", content="x")
assert ev.type == "delta"
assert ev.content == "x"
assert ev.usage is None
assert ev.model == ""
assert ev.finish_reason == ""
with pytest.raises(Exception): # FrozenInstanceError
ev.type = "end" # type: ignore[misc]
def test_stream_chat_strip_data_prefix_with_or_without_space():
"""SSE-префикс `data:` допускает опциональный пробел после двоеточия."""
sse_lines = [
'data:{"choices":[{"delta":{"content":"A"}}]}', # без пробела
"",
'data: {"choices":[{"delta":{"content":"B"}}]}', # с пробелом
"",
"data:[DONE]", # без пробела
"",
]
resp = _make_stream_response(sse_lines=sse_lines)
with _patched_stream(response=resp):
events = list(stream_chat(messages=[{"role": "user", "content": "x"}], base_url="http://m:1/v1"))
deltas = [e for e in events if e.type == "delta"]
assert [d.content for d in deltas] == ["A", "B"]
assert events[-1].type == "end"
-345
View File
@@ -1,345 +0,0 @@
"""Юнит-тесты для prompts.py.
Покрывают:
- load_system_prompt("icon") → строка из system_icon.txt
- load_system_prompt("illustration") → строка из system_illustration.txt
- load_system_prompt("junk") → ValueError
- load_few_shot() → пары {role, content}, валидные, идут в нужном порядке
- build_messages(icon) → первое сообщение system (правильный файл),
затем few-shot пары, затем финальный user.
- build_messages(illustration) → system — из system_illustration.txt.
- build_messages с image_b64 → последний user — list[dict] с text + image_url,
текст содержит "Recreate the visual content...".
- build_messages без image_b64 → последний user — простая строка.
- build_messages с palette → "Palette: ..." добавлено в текст.
- build_messages с пустым промптом → ValueError.
- build_messages с неизвестным mode → ValueError.
- Контракт: system промпт подставляется ИЗ ПРАВИЛЬНОГО ФАЙЛА (содержимое
различается для icon и illustration).
- Few-shot включён: в messages их пары.
Запуск: `python -m pytest tests/test_prompts.py -v`
"""
from __future__ import annotations
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
from prompts import ( # noqa: E402
build_messages,
load_few_shot,
load_system_prompt,
)
# ---------------------------------------------------------------------------
# load_system_prompt
# ---------------------------------------------------------------------------
def test_load_system_prompt_icon_returns_icon_text():
"""system_icon.txt: первая строка упоминает 'icon designer'."""
text = load_system_prompt("icon")
assert isinstance(text, str)
assert len(text) > 50
assert "icon" in text.lower()
# Уникальное для icon: "OmniSVG-Icon" (различается от "OmniSVG-Illustration")
assert "OmniSVG-Icon" in text
def test_load_system_prompt_illustration_returns_illustration_text():
"""system_illustration.txt: первая строка упоминает illustrator."""
text = load_system_prompt("illustration")
assert isinstance(text, str)
assert len(text) > 50
assert "illustration" in text.lower()
assert "OmniSVG-Illustration" in text
def test_load_system_prompt_distinct_for_each_mode():
"""Содержимое system для icon и illustration — разные файлы."""
icon_txt = load_system_prompt("icon")
ill_txt = load_system_prompt("illustration")
assert icon_txt != ill_txt
def test_load_system_prompt_unknown_mode_raises():
with pytest.raises(ValueError, match="неизвестный mode"):
load_system_prompt("portrait")
with pytest.raises(ValueError, match="неизвестный mode"):
load_system_prompt("")
def test_load_system_prompt_strips_trailing_whitespace():
"""Файл может заканчиваться на перевод строки — strip() его убирает."""
text = load_system_prompt("icon")
assert not text.endswith("\n")
assert not text.endswith(" ")
# ---------------------------------------------------------------------------
# load_few_shot
# ---------------------------------------------------------------------------
def test_load_few_shot_returns_list_of_message_dicts():
msgs = load_few_shot()
assert isinstance(msgs, list)
assert len(msgs) > 0
for m in msgs:
assert isinstance(m, dict)
assert m["role"] in ("user", "assistant")
assert isinstance(m["content"], str)
assert len(m["content"]) > 0
def test_load_few_shot_alternates_user_assistant():
"""USER → ASSISTANT → USER → ASSISTANT … (по дизайну)."""
msgs = load_few_shot()
for i, m in enumerate(msgs):
expected = "user" if i % 2 == 0 else "assistant"
assert m["role"] == expected, (
f"индекс {i}: ожидался role={expected}, получен {m['role']}"
)
def test_load_few_shot_includes_refusal_example():
"""Последний assistant — это refusal без <svg> (по дизайну)."""
msgs = load_few_shot()
# Ищем любой assistant-блок, который не содержит <svg>
refusals = [m for m in msgs if m["role"] == "assistant" and "<svg" not in m["content"]]
assert len(refusals) >= 1, "few-shot должен включать хотя бы один refusal"
def test_load_few_shot_includes_svg_examples():
"""Среди assistant'ов есть и реальные <svg>...</svg>."""
msgs = load_few_shot()
svg_assistants = [m for m in msgs if m["role"] == "assistant" and "<svg" in m["content"]]
assert len(svg_assistants) >= 2 # минимум 2 иконки (по дизайну)
def test_load_few_shot_is_stable_across_calls():
"""load_few_shot() детерминирован — не возвращает разное при повторе."""
first = load_few_shot()
second = load_few_shot()
assert first == second
# ---------------------------------------------------------------------------
# build_messages — icon
# ---------------------------------------------------------------------------
def test_build_messages_icon_has_system_first():
"""messages[0] — system с правильным содержимым (system_icon.txt)."""
msgs = build_messages("hello", mode="icon", n=1, temperature=0.4)
assert msgs[0]["role"] == "system"
# Содержимое — ровно из system_icon.txt
assert msgs[0]["content"] == load_system_prompt("icon")
def test_build_messages_icon_has_few_shot_pairs_after_system():
"""После system идут few-shot пары."""
msgs = build_messages("hello", mode="icon", n=1, temperature=0.4)
few_shot = load_few_shot()
# Сразу после system и до последнего user
assert msgs[1:1 + len(few_shot)] == few_shot
def test_build_messages_icon_last_message_is_user_with_prompt():
"""Последнее сообщение — user, content = prompt (без image)."""
msgs = build_messages("the magnifier", mode="icon", n=1, temperature=0.4)
last = msgs[-1]
assert last["role"] == "user"
assert last["content"] == "the magnifier"
def test_build_messages_icon_prompt_is_stripped():
"""Пробелы по краям — strip'аются."""
msgs = build_messages(" hello \n", mode="icon")
assert msgs[-1]["content"] == "hello"
# ---------------------------------------------------------------------------
# build_messages — illustration
# ---------------------------------------------------------------------------
def test_build_messages_illustration_has_illustration_system():
"""system для illustration — из правильного файла."""
msgs = build_messages("a fox", mode="illustration", n=1, temperature=0.4)
assert msgs[0]["role"] == "system"
assert msgs[0]["content"] == load_system_prompt("illustration")
assert "OmniSVG-Illustration" in msgs[0]["content"]
def test_build_messages_illustration_uses_same_few_shot():
"""Few-shot одинаков для обоих mode (только system различается)."""
icon_msgs = build_messages("hi", mode="icon")
ill_msgs = build_messages("hi", mode="illustration")
# Содержимое system разное
assert icon_msgs[0]["content"] != ill_msgs[0]["content"]
# Few-shot — одинаковый (без system: индекс 1..1+len(fs))
fs = load_few_shot()
assert icon_msgs[1:1 + len(fs)] == ill_msgs[1:1 + len(fs)]
assert icon_msgs[1:1 + len(fs)] == fs
# ---------------------------------------------------------------------------
# build_messages — image (multimodal)
# ---------------------------------------------------------------------------
def test_build_messages_with_image_uses_multimodal_content():
"""image_b64 задан → последний user.content — list[dict] с text+image_url."""
msgs = build_messages(
"describe",
mode="icon",
image_b64="data:image/png;base64,ZZZ",
n=1,
temperature=0.4,
)
last = msgs[-1]
assert last["role"] == "user"
assert isinstance(last["content"], list)
assert last["content"][0]["type"] == "text"
assert last["content"][1]["type"] == "image_url"
assert last["content"][1]["image_url"]["url"] == "data:image/png;base64,ZZZ"
def test_build_messages_with_image_text_includes_recreate_instruction():
"""С image_b64 в text идёт преамбула 'Recreate the visual content...'."""
msgs = build_messages(
"my prompt",
mode="icon",
image_b64="data:image/png;base64,X",
)
last = msgs[-1]
text_part = last["content"][0]["text"]
assert "Recreate the visual content" in text_part
assert "my prompt" in text_part
def test_build_messages_without_image_uses_string_content():
"""Без image_b64 — последний user.content = строка."""
msgs = build_messages("just text", mode="icon")
last = msgs[-1]
assert isinstance(last["content"], str)
assert last["content"] == "just text"
# И НЕ содержит "Recreate" (этот преамбул только для image-режима)
assert "Recreate" not in last["content"]
# ---------------------------------------------------------------------------
# build_messages — palette
# ---------------------------------------------------------------------------
def test_build_messages_with_palette_appends_palette_line():
"""palette добавляется отдельной строкой 'Palette: ...' в user text."""
msgs = build_messages(
"a fox",
mode="icon",
palette="blue and teal",
)
last = msgs[-1]
text = last["content"]
assert "a fox" in text
assert "Palette: blue and teal" in text
def test_build_messages_with_palette_in_multimodal():
"""Palette добавляется и в multimodal-режиме (в text-часть)."""
msgs = build_messages(
"a fox",
mode="illustration",
image_b64="data:image/png;base64,X",
palette="warm autumn",
)
last = msgs[-1]
text_part = last["content"][0]["text"]
assert "a fox" in text_part
assert "Palette: warm autumn" in text_part
assert "Recreate" in text_part
def test_build_messages_palette_stripped():
"""palette.strip() — пробелы по краям убираются."""
msgs = build_messages("x", mode="icon", palette=" blue \n")
assert "Palette: blue" in msgs[-1]["content"]
# ---------------------------------------------------------------------------
# build_messages — валидация входа
# ---------------------------------------------------------------------------
def test_build_messages_empty_prompt_raises():
with pytest.raises(ValueError, match="промпт пустой"):
build_messages("", mode="icon")
def test_build_messages_whitespace_only_prompt_raises():
"""Промпт из одних пробелов — тоже пустой."""
with pytest.raises(ValueError, match="промпт пустой"):
build_messages(" \n\t ", mode="icon")
def test_build_messages_unknown_mode_raises():
with pytest.raises(ValueError, match="неизвестный mode"):
build_messages("hi", mode="portrait")
# ---------------------------------------------------------------------------
# build_messages — параметры n и temperature
# ---------------------------------------------------------------------------
def test_build_messages_n_and_temperature_do_not_appear_in_messages():
"""n и temperature — параметры API, не подставляются в content.
(По дизайну: они идут в payload, а в messages их нет.)
"""
msgs = build_messages("hi", mode="icon", n=4, temperature=0.7)
full_text = " ".join(
m["content"] if isinstance(m["content"], str) else str(m["content"])
for m in msgs
)
assert "n=4" not in full_text
assert "temperature" not in full_text or "temperature" in load_system_prompt("icon")
# Проверяем, что temperature вообще не появляется в user-сообщениях.
for m in msgs:
if m["role"] == "user" and m is not msgs[0]: # system пропускаем
content = m["content"]
if isinstance(content, str):
assert "temperature" not in content
# ---------------------------------------------------------------------------
# Полная структура: порядок сообщений
# ---------------------------------------------------------------------------
def test_build_messages_message_order():
"""Структура: system, few-shot пары, финальный user."""
msgs = build_messages("final", mode="icon")
fs = load_few_shot()
# Длина = 1 (system) + len(few_shot) + 1 (final user)
assert len(msgs) == 1 + len(fs) + 1
# system — первый
assert msgs[0]["role"] == "system"
# Последний — user
assert msgs[-1]["role"] == "user"
# Последний user содержит наш промпт
assert "final" in (
msgs[-1]["content"]
if isinstance(msgs[-1]["content"], str)
else msgs[-1]["content"][0]["text"]
)
-256
View File
@@ -1,256 +0,0 @@
"""Юнит-тесты для renderer.py.
Покрывают:
- render_png() с валидным SVG → либо PNG bytes, либо None (если cairo недоступен)
Документируем: функция НЕ бросает — всегда возвращает bytes | None.
- render_png() с пустым/None входом → None
- render_png() с невалидным SVG (синтаксически сломан) → None (cairo упадёт,
renderer ловит)
- render_png() с SVG без xmlns → корректно добавляет xmlns (cairo требует)
- save_png() создаёт каталог при отсутствии и пишет файл с ожидаемым путём
- save_png() с произвольными PNG-байтами сохраняет as-is (без декодирования)
- Интеграция: render_png + save_png end-to-end (если cairo есть) либо
save_png-fallback (если cairo нет)
ЗАМЕЧАНИЕ ПОВЕДЕНИЯ (документируем в этом тесте):
Если `cairo-2.dll`/`libcairo` недоступен — `render_png()` возвращает `None`,
не бросает. Это явный контракт из дизайн-доки §6: "Если cairo недоступен
… `render_png()` возвращает `None`". Невалидный SVG: если cairo есть — cairosvg
бросит Exception, renderer ловит его через `except Exception` и возвращает
`None`. Если cairo нет — мы и так сразу возвращаем `None`.
Запуск: `python -m pytest tests/test_renderer.py -v`
"""
from __future__ import annotations
import io
import sys
import tempfile
from pathlib import Path
from unittest.mock import patch
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
from renderer import _get_cairosvg, render_png, save_png # noqa: E402
# ---------------------------------------------------------------------------
# Документация поведения: пустой/None вход
# ---------------------------------------------------------------------------
def test_render_png_empty_string_returns_none():
"""Пустая строка → None, без исключений."""
assert render_png("") is None
def test_render_png_whitespace_only_returns_none():
"""Строка из пробелов → None."""
assert render_png(" \n \t ") is None
# ---------------------------------------------------------------------------
# Документация поведения: валидный SVG
# ---------------------------------------------------------------------------
VALID_ICON_SVG = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">'
'<circle cx="32" cy="32" r="20" fill="#FF0000"/>'
'</svg>'
)
def test_render_png_valid_svg_returns_png_or_none():
"""Контракт: возвращает либо bytes, либо None.
Если cairo установлен — bytes (PNG, начинается с magic bytes).
Если cairo недоступен — None.
В обоих случаях НЕ бросает.
"""
result = render_png(VALID_ICON_SVG, size=(64, 64))
if result is None:
# cairo недоступен — это допустимо, см. дизайн-док §6.
assert _get_cairosvg() is None, (
"_get_cairosvg() вернул модуль, но render_png() вернул None"
)
else:
# cairo есть — должны получить валидный PNG.
assert isinstance(result, bytes)
assert len(result) > 0
# PNG magic: 89 50 4E 47 0D 0A 1A 0A
assert result.startswith(b"\x89PNG\r\n\x1a\n"), (
f"вывод render_png() не начинается с PNG magic: {result[:16]!r}"
)
def test_render_png_uses_size_width_as_output():
"""Если cairo доступен, проверяем output_width через прямой вызов."""
cairosvg = _get_cairosvg()
if cairosvg is None:
pytest.skip("cairo недоступен — нельзя проверить output_width")
result = render_png(VALID_ICON_SVG, size=(256, 256))
assert result is not None
# Проверяем, что PNG имеет ширину 256 (viewBox 64x64 → квадрат 256x256)
from PIL import Image
img = Image.open(io.BytesIO(result))
assert img.size == (256, 256), f"ожидался 256x256, получено {img.size}"
# ---------------------------------------------------------------------------
# Документация поведения: невалидный SVG
# ---------------------------------------------------------------------------
def test_render_png_invalid_svg_returns_none_or_raises_cairo():
"""Невалидный SVG → либо None, либо cairosvg бросает.
Наш контракт: render_png() ВСЕГДА возвращает bytes|None, не бросает наружу.
Это валидируется в `test_render_png_does_not_propagate_cairo_exceptions`
(через mock).
"""
broken = '<svg xmlns="http://www.w3.org/2000/svg"><circle cx="32" cy="32" r="999999999" fill="oops"/></svg>'
# Независимо от состояния cairo, не должны получить необработанный raise наружу.
result = render_png(broken, size=(64, 64))
assert result is None # либо cairo нет, либо cairo бросил и мы вернули None
def test_render_png_does_not_propagate_renderer_exceptions():
"""Даже если все бэкенды бросают — render_png() возвращает None, а не raise.
Сейчас renderer использует resvg-py как первый бэкенд, cairosvg — fallback.
Мокаем ОБА, чтобы render_png дошёл до конца и вернул None без exception.
"""
class FakeRenderer:
@staticmethod
def svg_to_bytes(**kwargs):
raise RuntimeError("simulated resvg failure")
class FakeCairo:
@staticmethod
def svg2png(**kwargs):
raise RuntimeError("simulated cairo failure")
with patch("renderer._get_resvg", return_value=FakeRenderer()), \
patch("renderer._get_cairosvg", return_value=FakeCairo()):
result = render_png(VALID_ICON_SVG, size=(64, 64))
assert result is None
# ---------------------------------------------------------------------------
# xmlns-инъекция
# ---------------------------------------------------------------------------
def test_render_png_adds_xmlns_if_missing():
"""Если модель не поставила xmlns, renderer добавляет его в первую <svg>."""
# Прямо проверяем _ensure_xmlns (внутренняя, но контрактная)
from renderer import _ensure_xmlns
no_xmlns = '<svg viewBox="0 0 64 64"><rect/></svg>'
fixed = _ensure_xmlns(no_xmlns)
assert 'xmlns="http://www.w3.org/2000/svg"' in fixed.split(">", 1)[0]
already = '<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64"></svg>'
assert _ensure_xmlns(already) == already
# В первом теге xmlns, но в href — не считается (спека смотрит ТОЛЬКО в head)
in_attr = '<svg viewBox="0 0 64 64" href="http://x.com"><rect/></svg>'
fixed_in_attr = _ensure_xmlns(in_attr)
head = fixed_in_attr.split(">", 1)[0]
assert 'xmlns="http://www.w3.org/2000/svg"' in head
# ---------------------------------------------------------------------------
# save_png()
# ---------------------------------------------------------------------------
def test_save_png_creates_directory_and_writes_file(tmp_path: Path):
"""save_png создаёт каталог и пишет файл с ожидаемым именем."""
fake_png = b"\x89PNG\r\n\x1a\n" + b"\x00" * 100 # валидный префикс + мусор
previews_dir = tmp_path / "previews"
out = save_png(
fake_png,
previews_dir=previews_dir,
record_id=42,
candidate_index=0,
)
assert out.exists()
assert out.is_file()
assert out == previews_dir / "42_0.png"
assert out.read_bytes() == fake_png
def test_save_png_does_not_decode_or_validate_png_bytes(tmp_path: Path):
"""save_png пишет байты as-is — никаких проверок PNG-формата."""
previews_dir = tmp_path / "p"
# Мусорные байты — не PNG, но save_png это не волнует.
out = save_png(
b"not actually png",
previews_dir=previews_dir,
record_id=1,
candidate_index=3,
)
assert out == previews_dir / "1_3.png"
assert out.read_bytes() == b"not actually png"
def test_save_png_appends_filename_with_zero_padded_index(tmp_path: Path):
"""record_id=5, candidate_index=7 → '5_7.png'."""
previews_dir = tmp_path / "p"
out = save_png(b"x", previews_dir=previews_dir, record_id=5, candidate_index=7)
assert out.name == "5_7.png"
def test_save_png_overwrites_existing_file(tmp_path: Path):
"""Повторный save с тем же (record_id, candidate_index) перезаписывает."""
previews_dir = tmp_path / "p"
out1 = save_png(b"first", previews_dir=previews_dir, record_id=1, candidate_index=0)
out2 = save_png(b"second", previews_dir=previews_dir, record_id=1, candidate_index=0)
assert out1 == out2
assert out1.read_bytes() == b"second"
# ---------------------------------------------------------------------------
# Интеграция: render + save (если cairo есть)
# ---------------------------------------------------------------------------
def test_render_and_save_integration(tmp_path: Path):
"""Сквозной кейс: render → save → файл существует и непустой."""
result = render_png(VALID_ICON_SVG, size=(64, 64))
if result is None:
pytest.skip("cairo недоступен — интеграционный тест render+save пропущен")
out = save_png(
result,
previews_dir=tmp_path / "previews",
record_id=10,
candidate_index=0,
)
assert out.exists()
assert out.stat().st_size > 0
# Содержимое — валидный PNG
from PIL import Image
img = Image.open(out)
img.verify() # поднимает, если битый PNG
# ---------------------------------------------------------------------------
# Граничный случай: путь с пробелами / Unicode
# ---------------------------------------------------------------------------
def test_save_png_works_with_unicode_in_path(tmp_path: Path):
"""save_png принимает каталог с юникодом (актуально для Windows)."""
previews_dir = tmp_path / "превью" # кириллица
out = save_png(b"x", previews_dir=previews_dir, record_id=1, candidate_index=0)
assert out.exists()
assert out.name == "1_0.png"
-213
View File
@@ -1,213 +0,0 @@
"""Юнит-тесты для validator.py.
Покрывают контрактные правила из design.md §5 и явные требования задачи:
- валидный SVG → ok=True
- отсутствует viewBox → ok=False (missing_viewbox)
- есть <script> → ok=False (disallowed_tag)
- есть <foreignObject> → ok=False (disallowed_tag)
- http:// ссылка в href → ok=False (external_ref)
- on*= атрибут → ok=False (event_handler)
Плюс несколько дополнительных кейсов, чтобы покрыть size-limit, парсинг
из обёртки с пояснением и edge cases.
Запуск:
python -m pytest tests/ -v
"""
from __future__ import annotations
import sys
from pathlib import Path
# Делаем корень проекта доступным как `import validator`.
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import pytest # noqa: E402
from validator import ( # noqa: E402
MAX_BYTES,
validate,
validate_svg,
extract_svg,
ValidatorError,
)
# ---------------------------------------------------------------------------
# Положительные кейсы
# ---------------------------------------------------------------------------
VALID_ICON = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64">
<circle cx="26" cy="26" r="14" fill="#3B82F6"/>
<rect x="36" y="34" width="6" height="20" rx="3" fill="#3B82F6" transform="rotate(45 39 44)"/>
</svg>"""
def test_valid_svg_returns_ok_and_cleaned():
ok, reason, cleaned = validate_svg(VALID_ICON, mode="icon")
assert ok is True
assert reason == ""
assert "<svg" in cleaned
assert cleaned.endswith("</svg>")
# viewBox должен сохраниться после сериализации.
assert 'viewBox="0 0 64 64"' in cleaned
def test_valid_svg_wrapped_in_prose_is_extracted():
"""Модель иногда отвечает `Here is your icon: <svg>...</svg>`. Должны вытащить."""
wrapped = (
"Sure! Here is your SVG icon:\n"
f"{VALID_ICON}\n"
"Hope that helps."
)
ok, reason, cleaned = validate_svg(wrapped, mode="icon")
assert ok is True, f"expected ok, got reason={reason!r}"
assert reason == ""
assert "<svg" in cleaned
def test_valid_illustration_with_filter():
"""В illustration разрешены <filter>, <feGaussianBlur> и т.п."""
ill = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 512 512" width="512" height="512">
<defs>
<filter id="b" x="-10%" y="-10%" width="120%" height="120%">
<feGaussianBlur stdDeviation="2"/>
</filter>
</defs>
<rect x="100" y="100" width="200" height="200" fill="#ff0000" filter="url(#b)"/>
</svg>"""
ok, reason, _ = validate_svg(ill, mode="illustration")
assert ok is True, f"reason={reason!r}"
# ---------------------------------------------------------------------------
# Негативные кейсы
# ---------------------------------------------------------------------------
def test_missing_viewbox_rejected():
bad = """<svg xmlns="http://www.w3.org/2000/svg">
<circle cx="10" cy="10" r="5"/>
</svg>"""
ok, reason, cleaned = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "missing_viewbox"
assert cleaned == ""
def test_bad_viewbox_rejected():
bad = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="oops">
<circle cx="10" cy="10" r="5"/>
</svg>"""
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "bad_viewbox"
def test_script_tag_rejected():
bad = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<script>alert(1)</script>
</svg>"""
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "disallowed_tag"
def test_foreign_object_rejected():
bad = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<foreignObject width="100" height="100">
<div xmlns="http://www.w3.org/1999/xhtml">Hi</div>
</foreignObject>
</svg>"""
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "disallowed_tag"
def test_http_ref_rejected():
bad = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<use href="http://example.com/sprite.svg#icon"/>
</svg>"""
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "external_ref"
def test_onclick_attribute_rejected():
bad = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<rect x="0" y="0" width="10" height="10" onclick="alert(1)"/>
</svg>"""
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "event_handler"
def test_url_with_https_rejected():
"""https:// в url(...) тоже блокируется."""
bad = """<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<defs>
<pattern id="p" width="4" height="4" patternUnits="userSpaceOnUse">
<rect width="4" height="4" fill="url('https://example.com/tex.png')"/>
</pattern>
</defs>
</svg>"""
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
# Сейчас external_ref ловится на `url(` независимо от http(s).
assert reason in ("external_ref", "disallowed_tag", "unknown_tag")
def test_too_large_rejected():
"""SVG > MAX_BYTES[mode] → too_large."""
# Соберём раздутый SVG с большим комментарием.
pad = "<!-- " + ("x" * (MAX_BYTES["icon"] + 100)) + " -->"
bad = (
'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">'
+ pad
+ "</svg>"
)
ok, reason, _ = validate_svg(bad, mode="icon")
assert ok is False
assert reason == "too_large"
def test_empty_input_rejected():
ok, reason, _ = validate_svg("", mode="icon")
assert ok is False
assert reason == "empty_input"
def test_no_svg_block_rejected():
ok, reason, _ = validate_svg("Sorry, I cannot help with that.", mode="icon")
assert ok is False
assert reason == "not_svg"
# ---------------------------------------------------------------------------
# Прямые проверки API (validate() бросает, extract_svg() утилитарный)
# ---------------------------------------------------------------------------
def test_validate_strict_raises():
with pytest.raises(ValidatorError) as excinfo:
validate("plain text", mode="icon")
assert excinfo.value.code == "not_svg"
def test_extract_svg_finds_block():
text = "noise <svg viewBox='0 0 1 1'></svg> noise"
found = extract_svg(text)
assert found.startswith("<svg")
assert found.endswith("</svg>")
def test_validate_svg_returns_tuple_for_unexpected_exception():
"""Даже если что-то странное, обёртка возвращает кортеж (ok=False, ...)."""
ok, reason, cleaned = validate_svg(None, mode="icon") # type: ignore[arg-type]
assert ok is False
assert reason in ("empty_input", "internal:TypeError")
assert cleaned == ""
-228
View File
@@ -1,228 +0,0 @@
"""Валидация SVG-ответов модели.
Делает две вещи:
1) `extract_svg()` — вытаскивает <svg>...</svg> из произвольного текста
(модель иногда отвечает с пояснениями или code-fence).
2) `validate(svg_text, mode=...)` — строгая проверка по правилам дизайна.
3) `validate_svg(svg_text) -> (ok, reason, cleaned_svg)` — обёртка для UI,
описанная в задаче: всегда возвращает кортеж, не бросает исключений.
Правила (соответствуют §5 design.md и §6 spec.md):
- well-formed XML (lxml);
- ровно один корневой <svg>;
- наличие `viewBox` (4 числа);
- теги — только из ALLOWED_TAGS;
- <script>, <foreignObject>, <image> — запрещены;
- http://, https:// — запрещены в href/xlink:href и в `url(...)`;
- on* event-handler атрибуты — запрещены (доп. защита по задаче);
- размер файла <= лимита mode (32/256 KB).
"""
from __future__ import annotations
import logging
import re
from dataclasses import dataclass
from lxml import etree
log = logging.getLogger(__name__)
SVG_NS = "http://www.w3.org/2000/svg"
ALLOWED_TAGS: set[str] = {
"svg",
"g",
"defs",
"symbol",
"use",
"path",
"rect",
"circle",
"ellipse",
"line",
"polygon",
"polyline",
"linearGradient",
"radialGradient",
"stop",
"text",
"tspan",
"textPath",
"title",
"desc",
"filter",
"feGaussianBlur",
"feOffset",
"feMerge",
"feMergeNode",
"feColorMatrix",
"feBlend",
"feFlood",
"feComposite",
"clipPath",
"mask",
"pattern",
}
DISALLOWED_TAGS: set[str] = {"script", "foreignObject", "image", "iframe", "style"}
MAX_BYTES: dict[str, int] = {
"icon": 32 * 1024,
"illustration": 256 * 1024,
}
class ValidatorError(Exception):
"""Ошибка валидации с машино-читаемым кодом и человеческим пояснением."""
def __init__(self, code: str, detail: str = ""):
super().__init__(f"{code}: {detail}" if detail else code)
self.code = code
self.detail = detail
_SVG_RE = re.compile(
r"<svg\b[^>]*>.*?</svg\s*>",
re.IGNORECASE | re.DOTALL,
)
def extract_svg(text: str) -> str:
"""Достаёт первое вхождение <svg>...</svg> из произвольного ответа модели."""
if not text or "<svg" not in text.lower():
raise ValidatorError("not_svg", "в ответе нет <svg>...</svg>")
match = _SVG_RE.search(text)
if not match:
raise ValidatorError("not_svg", "не удалось выделить <svg>...</svg> regex'ом")
return match.group(0).strip()
def _local(tag: object) -> str:
"""Снимает namespace: {ns}local -> local. Терпимо к non-str (cyfunction)."""
s = str(tag)
if "}" in s:
return s.split("}", 1)[1]
return s
def _iter_real_elements(root: etree._Element):
"""Итератор только по Element-узлам (пропускает комменты и PI)."""
for el in root.iter():
# lxml._Element и etree._Element имеют tag как str; cyfunction — это
# callable у комментариев/PI. Отфильтруем через isinstance.
if isinstance(getattr(el, "tag", None), str):
yield el
def _check_viewbox(root: etree._Element) -> None:
vb = root.get("viewBox")
if not vb:
raise ValidatorError("missing_viewbox", "атрибут viewBox отсутствует")
try:
parts = [float(x) for x in vb.split()]
except ValueError as exc:
raise ValidatorError("bad_viewbox", f"viewBox='{vb}'") from exc
if len(parts) != 4:
raise ValidatorError("bad_viewbox", f"viewBox='{vb}' (нужно 4 числа)")
def _check_tags_and_attrs(root: etree._Element) -> None:
for el in _iter_real_elements(root):
local = _local(el.tag)
if local in DISALLOWED_TAGS:
raise ValidatorError("disallowed_tag", local)
if local not in ALLOWED_TAGS:
raise ValidatorError("unknown_tag", local)
for attr, val in el.attrib.items():
if not val:
continue
attr_local = _local(attr).lower()
if attr_local.startswith("on"):
raise ValidatorError("event_handler", f"{attr_local}={val[:80]}")
if "http://" in val or "https://" in val:
if attr_local.endswith("href") or "url(" in val.lower():
raise ValidatorError("external_ref", f"{attr_local}={val[:80]}")
def _ensure_size(svg_text: str, mode: str) -> None:
limit = MAX_BYTES.get(mode)
if limit is None:
return
if len(svg_text.encode("utf-8")) > limit:
raise ValidatorError(
"too_large",
f"{len(svg_text.encode('utf-8'))} байт > {limit} байт (mode={mode})",
)
def _serialize_clean(root: etree._Element) -> str:
if not root.get("xmlns"):
new_root = etree.Element(_local(root.tag), nsmap={None: SVG_NS})
for k, v in root.attrib.items():
new_root.set(_local(k), v)
for child in root:
new_root.append(child)
root = new_root
xml_bytes = etree.tostring(
root, pretty_print=False, xml_declaration=False, encoding="utf-8"
)
return xml_bytes.decode("utf-8")
def validate(svg_text: str, *, mode: str = "icon") -> etree._Element:
"""Строгая валидация по правилам дизайна. Бросает ValidatorError."""
svg_text = extract_svg(svg_text)
try:
root = etree.fromstring(svg_text.encode("utf-8"))
except etree.XMLSyntaxError as exc:
raise ValidatorError("malformed_xml", str(exc)) from exc
if _local(root.tag) != "svg":
raise ValidatorError("not_svg", root.tag)
_check_viewbox(root)
_check_tags_and_attrs(root)
_ensure_size(svg_text, mode)
return root
@dataclass(frozen=True)
class ValidationResult:
ok: bool
reason: str
cleaned_svg: str
def validate_svg(svg_text: str, *, mode: str = "icon") -> tuple[bool, str, str]:
"""Удобная обёртка над `validate()`: всегда возвращает кортеж.
Returns:
Кортеж `(ok, reason, cleaned_svg)`.
"""
if not svg_text or not svg_text.strip():
return False, "empty_input", ""
try:
root = validate(svg_text, mode=mode)
except ValidatorError as exc:
log.info("validate_svg: reject (%s)", exc.code)
return False, exc.code, ""
except Exception as exc: # noqa: BLE001
log.exception("validate_svg: неожиданная ошибка")
return False, f"internal:{type(exc).__name__}", ""
return True, "", _serialize_clean(root)
__all__ = [
"ALLOWED_TAGS",
"DISALLOWED_TAGS",
"MAX_BYTES",
"SVG_NS",
"ValidationResult",
"ValidatorError",
"extract_svg",
"validate",
"validate_svg",
]