feat(image-gen): add service layer with safe artifact saving
This commit is contained in:
@@ -0,0 +1,197 @@
|
||||
"""Image generation service: resolve model, dispatch adapter, save safely."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any
|
||||
|
||||
from .adapters.base import ImageGenAdapter, ImageGenError
|
||||
from .adapters.gemini import GeminiImageAdapter
|
||||
from .adapters.openai import OpenAIImageAdapter
|
||||
from .config import ImageModelEntry, load_image_generation_settings
|
||||
|
||||
ADAPTER_CLASSES: dict[str, type] = {
|
||||
"openai": OpenAIImageAdapter,
|
||||
"gemini": GeminiImageAdapter,
|
||||
}
|
||||
|
||||
ALLOWED_INPUT_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"}
|
||||
MAX_INPUT_BYTES = 20 * 1024 * 1024
|
||||
DEFAULT_MIME_TYPE = "image/png"
|
||||
|
||||
|
||||
def _config_path() -> Path:
|
||||
from EvoScientist.config.settings import get_config_path
|
||||
|
||||
return get_config_path()
|
||||
|
||||
|
||||
def _available_hint(models: list[ImageModelEntry]) -> str:
|
||||
names = [entry.display_name() for entry in models]
|
||||
return f" Available image models: {', '.join(names)}." if names else ""
|
||||
|
||||
|
||||
def _resolve_entry(model: str | None) -> tuple[ImageModelEntry, float]:
|
||||
settings = load_image_generation_settings(config_path=_config_path())
|
||||
if not settings.models:
|
||||
raise ImageGenError(
|
||||
"There are no image models configured; add an image_generation "
|
||||
"section to config.yaml."
|
||||
)
|
||||
if model:
|
||||
entry = settings.find_model(model)
|
||||
if entry is None:
|
||||
raise ImageGenError(
|
||||
f"Unknown image model {model!r}."
|
||||
+ _available_hint(settings.models)
|
||||
)
|
||||
return entry, settings.timeout_seconds
|
||||
default = settings.default_model or settings.models[0].id
|
||||
entry = settings.find_model(default)
|
||||
if entry is None:
|
||||
raise ImageGenError(
|
||||
f"Default image model {default!r} is not in the models list."
|
||||
+ _available_hint(settings.models)
|
||||
)
|
||||
return entry, settings.timeout_seconds
|
||||
|
||||
|
||||
def _adapter_for(entry: ImageModelEntry, timeout: float) -> ImageGenAdapter:
|
||||
adapter_cls = ADAPTER_CLASSES.get(entry.provider)
|
||||
if adapter_cls is None:
|
||||
raise ImageGenError(f"Unknown image provider {entry.provider!r}.")
|
||||
return adapter_cls(entry, timeout=timeout)
|
||||
|
||||
|
||||
def _clean_logical_path(value: str) -> str:
|
||||
raw = value.replace("\\", "/").strip()
|
||||
if not raw:
|
||||
raise ImageGenError("Path is required")
|
||||
if raw.startswith("/") or re.match(r"^[A-Za-z]:/", raw):
|
||||
raise ImageGenError("Absolute paths are not allowed")
|
||||
path = PurePosixPath(raw)
|
||||
if any(part in {"", ".", ".."} for part in path.parts):
|
||||
raise ImageGenError("Path traversal is not allowed")
|
||||
return str(path)
|
||||
|
||||
|
||||
def _normalize_output_paths(
|
||||
output_path: str | None, *, default_stem: str, count: int
|
||||
) -> list[str]:
|
||||
if output_path is None or not output_path.strip():
|
||||
stamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
base = f"artifacts/{default_stem}_{stamp}.png"
|
||||
else:
|
||||
base = _clean_logical_path(output_path)
|
||||
if not base.startswith("artifacts/"):
|
||||
raise ImageGenError("output_path must start with artifacts/")
|
||||
suffix = PurePosixPath(base).suffix.lower()
|
||||
if suffix and suffix != ".png":
|
||||
raise ImageGenError("Only PNG image outputs are supported")
|
||||
if not suffix:
|
||||
base = f"{base}.png"
|
||||
if count <= 1:
|
||||
return [base]
|
||||
path = PurePosixPath(base)
|
||||
return [f"{path.with_suffix('')}_{idx}{path.suffix}" for idx in range(1, count + 1)]
|
||||
|
||||
|
||||
def _resolve_input_image(workspace: Path, logical_path: str) -> Path:
|
||||
clean = _clean_logical_path(logical_path)
|
||||
if PurePosixPath(clean).suffix.lower() not in ALLOWED_INPUT_EXTENSIONS:
|
||||
raise ImageGenError("Unsupported image extension")
|
||||
root = workspace.resolve()
|
||||
target = (root / clean).resolve()
|
||||
if not target.is_relative_to(root) or not target.is_file():
|
||||
raise ImageGenError("Input image not found")
|
||||
if target.stat().st_size > MAX_INPUT_BYTES:
|
||||
raise ImageGenError("Input image is too large")
|
||||
return target
|
||||
|
||||
|
||||
def _save_images(
|
||||
workspace: Path,
|
||||
images: list[bytes],
|
||||
*,
|
||||
output_path: str | None,
|
||||
default_stem: str,
|
||||
) -> list[str]:
|
||||
(workspace / "artifacts").mkdir(parents=True, exist_ok=True)
|
||||
logical_paths = _normalize_output_paths(
|
||||
output_path, default_stem=default_stem, count=len(images)
|
||||
)
|
||||
root = workspace.resolve()
|
||||
saved: list[str] = []
|
||||
for logical_path, data in zip(logical_paths, images, strict=True):
|
||||
target = (root / logical_path).resolve()
|
||||
if not target.is_relative_to(root):
|
||||
raise ImageGenError("Invalid output_path")
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
target.write_bytes(data)
|
||||
saved.append(logical_path)
|
||||
return saved
|
||||
|
||||
|
||||
def _success(paths: list[str], *, model: str, size: str) -> dict[str, Any]:
|
||||
return {
|
||||
"ok": True,
|
||||
"path": paths[0] if paths else None,
|
||||
"paths": paths,
|
||||
"mime_type": DEFAULT_MIME_TYPE,
|
||||
"model": model,
|
||||
"size": size,
|
||||
}
|
||||
|
||||
|
||||
async def generate_for_workspace(
|
||||
workspace: Path,
|
||||
*,
|
||||
prompt: str,
|
||||
model: str | None = None,
|
||||
size: str = "1024x1024",
|
||||
quality: str = "auto",
|
||||
background: str = "auto",
|
||||
output_path: str | None = None,
|
||||
n: int = 1,
|
||||
) -> dict[str, Any]:
|
||||
entry, timeout = _resolve_entry(model)
|
||||
if not entry.supports_generation:
|
||||
raise ImageGenError(f"Image model {entry.display_name()!r} cannot generate.")
|
||||
adapter = _adapter_for(entry, timeout)
|
||||
images = await adapter.generate(
|
||||
prompt=prompt, size=size, quality=quality, background=background, n=n
|
||||
)
|
||||
saved = _save_images(
|
||||
workspace, images, output_path=output_path, default_stem="generated"
|
||||
)
|
||||
return _success(saved, model=entry.id, size=size)
|
||||
|
||||
|
||||
async def edit_for_workspace(
|
||||
workspace: Path,
|
||||
*,
|
||||
image_path: str,
|
||||
prompt: str,
|
||||
model: str | None = None,
|
||||
mask_path: str | None = None,
|
||||
size: str = "1024x1024",
|
||||
quality: str = "auto",
|
||||
output_path: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
entry, timeout = _resolve_entry(model)
|
||||
if not entry.supports_edit:
|
||||
raise ImageGenError(
|
||||
f"Image model {entry.display_name()!r} does not support edit."
|
||||
)
|
||||
image = _resolve_input_image(workspace, image_path)
|
||||
mask = _resolve_input_image(workspace, mask_path) if mask_path else None
|
||||
adapter = _adapter_for(entry, timeout)
|
||||
images = await adapter.edit(
|
||||
image=image, mask=mask, prompt=prompt, size=size, quality=quality
|
||||
)
|
||||
saved = _save_images(
|
||||
workspace, images, output_path=output_path, default_stem="edited"
|
||||
)
|
||||
return _success(saved, model=entry.id, size=size)
|
||||
@@ -0,0 +1,165 @@
|
||||
"""Tests for image_gen.service."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.image_gen import service
|
||||
from EvoScientist.image_gen.adapters.base import ImageGenError
|
||||
|
||||
PNG_BYTES = b"\x89PNG\r\n\x1a\n" + b"svc"
|
||||
|
||||
|
||||
class _FakeAdapter:
|
||||
"""Records calls; returns fixed bytes."""
|
||||
|
||||
instances: list["_FakeAdapter"] = []
|
||||
|
||||
def __init__(self, entry, *, client=None, timeout=120.0):
|
||||
self.entry = entry
|
||||
self.calls: list[dict] = []
|
||||
_FakeAdapter.instances.append(self)
|
||||
|
||||
async def generate(self, **kwargs):
|
||||
self.calls.append({"kind": "generate", **kwargs})
|
||||
return [PNG_BYTES]
|
||||
|
||||
async def edit(self, **kwargs):
|
||||
self.calls.append({"kind": "edit", **kwargs})
|
||||
return [PNG_BYTES]
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def fake_adapters(monkeypatch, tmp_path):
|
||||
_FakeAdapter.instances = []
|
||||
monkeypatch.setattr(
|
||||
service, "ADAPTER_CLASSES", {"openai": _FakeAdapter, "gemini": _FakeAdapter}
|
||||
)
|
||||
config = tmp_path / "config.yaml"
|
||||
config.write_text(
|
||||
"""
|
||||
image_generation:
|
||||
default_model: gpt-image-2
|
||||
models:
|
||||
- id: gpt-image-2
|
||||
name: GPT Image 2
|
||||
provider: openai
|
||||
api_key: sk-x
|
||||
- id: imagen-4.0-generate-001
|
||||
provider: gemini
|
||||
api_key: sk-y
|
||||
supports_edit: false
|
||||
""",
|
||||
encoding="utf-8",
|
||||
)
|
||||
monkeypatch.setattr(service, "_config_path", lambda: config)
|
||||
return config
|
||||
|
||||
|
||||
async def test_generate_saves_to_artifacts(tmp_path):
|
||||
result = await service.generate_for_workspace(tmp_path, prompt="a cat")
|
||||
assert result["ok"] is True
|
||||
assert result["model"] == "gpt-image-2"
|
||||
assert result["paths"][0].startswith("artifacts/generated_")
|
||||
saved = tmp_path / result["paths"][0]
|
||||
assert saved.read_bytes() == PNG_BYTES
|
||||
adapter = _FakeAdapter.instances[0]
|
||||
assert adapter.entry.id == "gpt-image-2"
|
||||
assert adapter.calls[0]["prompt"] == "a cat"
|
||||
|
||||
|
||||
async def test_generate_explicit_output_path(tmp_path):
|
||||
result = await service.generate_for_workspace(
|
||||
tmp_path, prompt="x", output_path="artifacts/cover.png"
|
||||
)
|
||||
assert result["paths"] == ["artifacts/cover.png"]
|
||||
assert (tmp_path / "artifacts" / "cover.png").is_file()
|
||||
|
||||
|
||||
async def test_output_path_traversal_rejected(tmp_path):
|
||||
with pytest.raises(ImageGenError):
|
||||
await service.generate_for_workspace(
|
||||
tmp_path, prompt="x", output_path="../escape.png"
|
||||
)
|
||||
with pytest.raises(ImageGenError):
|
||||
await service.generate_for_workspace(
|
||||
tmp_path, prompt="x", output_path="/abs/path.png"
|
||||
)
|
||||
|
||||
|
||||
async def test_output_path_must_be_artifacts_png(tmp_path):
|
||||
with pytest.raises(ImageGenError):
|
||||
await service.generate_for_workspace(
|
||||
tmp_path, prompt="x", output_path="other/cover.png"
|
||||
)
|
||||
with pytest.raises(ImageGenError):
|
||||
await service.generate_for_workspace(
|
||||
tmp_path, prompt="x", output_path="artifacts/cover.jpg"
|
||||
)
|
||||
|
||||
|
||||
async def test_unknown_model_lists_available(tmp_path):
|
||||
with pytest.raises(ImageGenError) as excinfo:
|
||||
await service.generate_for_workspace(tmp_path, prompt="x", model="nope")
|
||||
message = str(excinfo.value)
|
||||
assert "GPT Image 2" in message and "imagen-4.0-generate-001" in message
|
||||
|
||||
|
||||
async def test_no_models_configured(tmp_path, monkeypatch):
|
||||
empty = tmp_path / "empty.yaml"
|
||||
empty.write_text("webui_port: 4716\n", encoding="utf-8")
|
||||
monkeypatch.setattr(service, "_config_path", lambda: empty)
|
||||
with pytest.raises(ImageGenError, match="no image models"):
|
||||
await service.generate_for_workspace(tmp_path, prompt="x")
|
||||
|
||||
|
||||
async def test_n_greater_than_one_numbered(tmp_path):
|
||||
class _TwoAdapter(_FakeAdapter):
|
||||
async def generate(self, **kwargs):
|
||||
return [PNG_BYTES, PNG_BYTES]
|
||||
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
monkeypatch.setattr(
|
||||
service, "ADAPTER_CLASSES", {"openai": _TwoAdapter, "gemini": _TwoAdapter}
|
||||
)
|
||||
result = await service.generate_for_workspace(
|
||||
tmp_path, prompt="x", output_path="artifacts/v.png", n=2
|
||||
)
|
||||
assert result["paths"] == ["artifacts/v_1.png", "artifacts/v_2.png"]
|
||||
monkeypatch.undo()
|
||||
|
||||
|
||||
async def test_edit_reads_input_and_saves(tmp_path):
|
||||
src = tmp_path / "artifacts" / "src.png"
|
||||
src.parent.mkdir(parents=True)
|
||||
src.write_bytes(PNG_BYTES)
|
||||
result = await service.edit_for_workspace(
|
||||
tmp_path, image_path="artifacts/src.png", prompt="make blue"
|
||||
)
|
||||
assert result["ok"] is True
|
||||
assert result["paths"][0].startswith("artifacts/edited_")
|
||||
adapter = _FakeAdapter.instances[0]
|
||||
assert adapter.calls[0]["kind"] == "edit"
|
||||
|
||||
|
||||
async def test_edit_rejects_bad_input(tmp_path):
|
||||
with pytest.raises(ImageGenError):
|
||||
await service.edit_for_workspace(
|
||||
tmp_path, image_path="artifacts/missing.png", prompt="x"
|
||||
)
|
||||
outside = tmp_path / "outside.txt"
|
||||
outside.write_text("hi", encoding="utf-8")
|
||||
with pytest.raises(ImageGenError):
|
||||
await service.edit_for_workspace(
|
||||
tmp_path, image_path="outside.txt", prompt="x"
|
||||
)
|
||||
|
||||
|
||||
async def test_edit_blocked_when_model_lacks_support(tmp_path):
|
||||
with pytest.raises(ImageGenError, match="edit"):
|
||||
await service.edit_for_workspace(
|
||||
tmp_path,
|
||||
image_path="artifacts/whatever.png",
|
||||
prompt="x",
|
||||
model="imagen-4.0-generate-001",
|
||||
)
|
||||
Reference in New Issue
Block a user