| """Image generation — mock placeholders or FLUX Klein backend.""" |
|
|
| from __future__ import annotations |
|
|
| import base64 |
| import os |
| import random |
| from io import BytesIO |
| from pathlib import Path |
| from typing import Any |
|
|
| import httpx |
| from PIL import Image, ImageDraw |
|
|
| from src.image_utils import ( |
| bait_prompt, |
| chroma_key_green, |
| fish_edit_prompt, |
| fish_greenscreen_prompt, |
| fish_underwater_prompt, |
| image_to_png_bytes, |
| trim_transparent, |
| ) |
|
|
| ROOT = Path(__file__).resolve().parent.parent |
| IMAGE_SIZE = int(os.getenv("IMAGE_SIZE", "384")) |
|
|
|
|
| class ImageClient: |
| def __init__(self) -> None: |
| self.mode = os.getenv("IMAGE_MODE", "mock").lower() |
| self.api_base = os.getenv("IMAGE_API_BASE", "").rstrip("/") |
| self.flux_model = os.getenv("FLUX_MODEL", "black-forest-labs/FLUX.2-klein-9B") |
| self.size = IMAGE_SIZE |
|
|
| def _placeholder(self, label: str, *, green: bool = False) -> bytes: |
| rnd = random.Random(label) |
| body = tuple(rnd.randint(90, 230) for _ in range(3)) + (255,) |
| fin = tuple(max(0, c - 50) for c in body[:3]) + (255,) |
| img = Image.new( |
| "RGBA", (256, 256), (0, 255, 0, 255) if green else (26, 74, 106, 255) |
| ) |
| draw = ImageDraw.Draw(img) |
| if not green: |
| for _ in range(8): |
| x, y = rnd.randint(0, 255), rnd.randint(0, 255) |
| draw.ellipse((x, y, x + 6, y + 6), outline=(180, 220, 235, 120)) |
| draw.polygon([(196, 128), (240, 92), (240, 164)], fill=fin) |
| draw.ellipse((40, 88, 208, 168), fill=body) |
| draw.polygon([(110, 92), (140, 60), (160, 94)], fill=fin) |
| draw.ellipse((62, 108, 86, 132), fill=(255, 255, 255, 255)) |
| draw.ellipse((70, 116, 82, 128), fill=(20, 20, 30, 255)) |
| draw.arc((44, 118, 92, 156), 20, 110, fill=(20, 20, 30, 200), width=3) |
| if green: |
| img = chroma_key_green(img) |
| return image_to_png_bytes(img) |
|
|
| def _generate_remote(self, prompt: str) -> Image.Image: |
| if not self.api_base: |
| raise RuntimeError("IMAGE_API_BASE required for remote image mode") |
| with httpx.Client(timeout=300) as client: |
| url = self.api_base.rstrip("/") |
| if not url.endswith(".run"): |
| url = f"{url}/generate" |
| r = client.post( |
| url, |
| json={"prompt": prompt, "width": self.size, "height": self.size}, |
| ) |
| r.raise_for_status() |
| b64 = r.json()["image_b64"] |
| return Image.open(BytesIO(base64.b64decode(b64))) |
|
|
| def _generate_local(self, prompt: str) -> Image.Image: |
| import torch |
| from diffusers import Flux2KleinPipeline |
|
|
| pipe = Flux2KleinPipeline.from_pretrained( |
| self.flux_model, |
| torch_dtype=torch.bfloat16, |
| ) |
| pipe.enable_model_cpu_offload() |
| result = pipe( |
| prompt=prompt, |
| num_inference_steps=4, |
| width=self.size, |
| height=self.size, |
| ) |
| return result.images[0] |
|
|
| def _generate_modal(self, prompt: str, source_png: bytes | None = None) -> Image.Image: |
| import modal |
|
|
| app_name = os.getenv("MODAL_APP_NAME", "llm-fishing") |
| cls_name = os.getenv("MODAL_IMAGE_CLASS", "ImageGenerator") |
| ImageGenerator = modal.Cls.from_name(app_name, cls_name) |
| b64 = ImageGenerator().generate.remote( |
| prompt, |
| width=self.size, |
| height=self.size, |
| image_b64=base64.b64encode(source_png).decode("ascii") if source_png else None, |
| ) |
| return Image.open(BytesIO(base64.b64decode(b64))) |
|
|
| def placeholder_png(self, label: str, *, green: bool = False) -> bytes: |
| return self._placeholder(label, green=green) |
|
|
| def _generate( |
| self, prompt: str, *, greenscreen: bool, source_png: bytes | None = None |
| ) -> bytes: |
| try: |
| if self.mode == "mock": |
| return self._placeholder(prompt[:20], green=greenscreen) |
| if self.mode == "modal": |
| img = self._generate_modal(prompt, source_png) |
| elif self.mode == "remote": |
| img = self._generate_remote(prompt) |
| elif self.mode == "local": |
| img = self._generate_local(prompt) |
| else: |
| raise ValueError(f"Unknown IMAGE_MODE: {self.mode}") |
|
|
| if greenscreen: |
| img = trim_transparent(chroma_key_green(img)) |
| return image_to_png_bytes(img) |
| except Exception: |
| return self._placeholder(prompt[:20], green=greenscreen) |
|
|
| def generate_bait(self, bait: str) -> bytes: |
| return self._generate(bait_prompt(bait), greenscreen=True) |
|
|
| def generate_fish_underwater(self, fish: dict[str, Any], bait: str = "") -> bytes: |
| return self._generate(fish_underwater_prompt(fish, bait), greenscreen=False) |
|
|
| def generate_fish_sprite( |
| self, fish: dict[str, Any], underwater_png: bytes | None = None, bait: str = "" |
| ) -> bytes: |
| if underwater_png and self.mode == "modal": |
| return self._generate( |
| fish_edit_prompt(fish, bait), |
| greenscreen=True, |
| source_png=underwater_png, |
| ) |
| return self._generate(fish_greenscreen_prompt(fish, bait), greenscreen=True) |
|
|
| def warmup(self) -> None: |
| """Tiny FLUX render to wake the image container.""" |
| if self.mode == "mock": |
| return |
| if self.mode == "modal": |
| import modal |
|
|
| app_name = os.getenv("MODAL_APP_NAME", "llm-fishing") |
| cls_name = os.getenv("MODAL_IMAGE_CLASS", "ImageGenerator") |
| ImageGenerator = modal.Cls.from_name(app_name, cls_name) |
| ImageGenerator().generate.remote( |
| "solid bright green square pixel art, game asset", |
| width=128, |
| height=128, |
| ) |
| return |
| self._generate("solid green pixel", greenscreen=True) |
|
|