Files
EvoScientist/EvoScientist/image_gen/service.py
T
m4 aae8d0a379 feat: workspace file references, read-file-images middleware, image model enabled flag
In-progress work committed to unblock the config import/export plan:
- prompts: FILE_REFERENCES section for workspace-relative file citation
- backends: resolve quoted virtual absolute paths onto the sandbox workspace
- middleware: read_file_images middleware; message_budget extensions
- image_gen/model_registry: image model 'enabled' flag refactor
- memory/launch, gateway/background_runs, tools/image follow-ons
- scripts: dev_backend.sh, release.sh
- tests for the above
2026-08-12 19:43:35 +08:00

237 lines
7.7 KiB
Python

"""Image generation service: resolve model, dispatch adapter, save safely."""
from __future__ import annotations
import asyncio
import re
import time
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)
)
if not entry.enabled:
raise ImageGenError(
f"Image model {entry.display_name()!r} is disabled."
)
return entry, settings.timeout_seconds
default = settings.default_model or settings.models[0].id
entry = settings.find_model(default)
if entry is not None and entry.enabled:
return entry, settings.timeout_seconds
for candidate in settings.models:
if candidate.enabled:
return candidate, settings.timeout_seconds
raise ImageGenError(
"There are no enabled image models; enable one in the "
"image_generation section of config.yaml."
)
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,
}
TEST_PROMPT = "A single red circle centered on a plain white background."
async def test_entry(entry: ImageModelEntry, *, timeout: float) -> int:
"""Generate one probe image against ``entry``; returns latency in ms.
The image bytes are discarded: this only proves the configured provider,
credentials and defaults can complete a real generation call.
"""
adapter = _adapter_for(entry, timeout)
started = time.monotonic()
await adapter.generate(
prompt=TEST_PROMPT,
size=entry.default_size,
quality=entry.default_quality,
background="auto",
n=1,
)
return int((time.monotonic() - started) * 1000)
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]:
# Config load and file I/O run in a thread: the langgraph dev server
# intercepts blocking calls (os.mkdir, read_text, stat) made directly on
# the event loop and aborts the tool call.
entry, timeout = await asyncio.to_thread(_resolve_entry, model)
adapter = _adapter_for(entry, timeout)
images = await adapter.generate(
prompt=prompt, size=size, quality=quality, background=background, n=n
)
saved = await asyncio.to_thread(
_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 = await asyncio.to_thread(_resolve_entry, model)
image = await asyncio.to_thread(_resolve_input_image, workspace, image_path)
mask = (
await asyncio.to_thread(_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 = await asyncio.to_thread(
_save_images,
workspace,
images,
output_path=output_path,
default_stem="edited",
)
return _success(saved, model=entry.id, size=size)