Add dual-GPU FLUX 9B image pipeline
This commit is contained in:
@@ -0,0 +1,260 @@
|
||||
"""Hermes image generation/edit provider for the local Athena router.
|
||||
|
||||
The provider deliberately rejects public destinations. Prompts and generated
|
||||
images may only travel to a loopback or private-network address.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import ipaddress
|
||||
import json
|
||||
import os
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from agent.image_gen_provider import (
|
||||
DEFAULT_ASPECT_RATIO,
|
||||
ImageGenProvider,
|
||||
error_response,
|
||||
normalize_reference_images,
|
||||
resolve_aspect_ratio,
|
||||
save_b64_image,
|
||||
success_response,
|
||||
)
|
||||
from agent.secret_scope import get_secret
|
||||
|
||||
|
||||
_SIZES = {
|
||||
"landscape": "1536x1024",
|
||||
"square": "1024x1024",
|
||||
"portrait": "1024x1536",
|
||||
}
|
||||
_DEFAULT_BASE_URL = "http://192.168.1.212:8081/v1"
|
||||
_DEFAULT_MODEL = "FLUX.2-klein-9B-fp8-beta"
|
||||
_MAX_IMAGE_BYTES = 20 * 1024 * 1024
|
||||
|
||||
|
||||
def _base_url() -> str:
|
||||
"""Use the dedicated image URL and never inherit an unrelated chat URL."""
|
||||
override = os.environ.get("ATHENA_IMAGE_BASE_URL", "").strip()
|
||||
return (override or _DEFAULT_BASE_URL).rstrip("/")
|
||||
|
||||
|
||||
def _model() -> str:
|
||||
return os.environ.get("ATHENA_IMAGE_MODEL", "").strip() or _DEFAULT_MODEL
|
||||
|
||||
|
||||
def _api_key() -> str:
|
||||
"""Prefer a scoped key; accept the existing router key for migration."""
|
||||
return (
|
||||
get_secret("ATHENA_IMAGE_API_KEY", "")
|
||||
or get_secret("ROUTER_API_KEY", "")
|
||||
or get_secret("HERMES_CUSTOM_192_168_1_212_8081_API_KEY", "")
|
||||
or ""
|
||||
).strip()
|
||||
|
||||
|
||||
def _private_destination(url: str) -> bool:
|
||||
"""Fail closed unless the configured endpoint is local/private."""
|
||||
try:
|
||||
parsed = urllib.parse.urlparse(url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||
return False
|
||||
if parsed.hostname == "localhost":
|
||||
return True
|
||||
address = ipaddress.ip_address(parsed.hostname)
|
||||
return address.is_private or address.is_loopback
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
|
||||
|
||||
def _load_private_image(ref: str) -> bytes:
|
||||
"""Load a local/data/private-LAN image without contacting public hosts."""
|
||||
ref = ref.strip()
|
||||
lower = ref.lower()
|
||||
if lower.startswith("data:image/"):
|
||||
_, separator, payload = ref.partition(",")
|
||||
if not separator:
|
||||
raise ValueError("invalid image data URI")
|
||||
data = base64.b64decode(payload, validate=True)
|
||||
elif lower.startswith(("http://", "https://")):
|
||||
if not _private_destination(ref):
|
||||
raise ValueError("public reference-image URLs are blocked")
|
||||
request = urllib.request.Request(
|
||||
ref, headers={"User-Agent": "Hermes-Athena-Image/2.0"})
|
||||
with urllib.request.urlopen(request, timeout=60) as response:
|
||||
data = response.read(_MAX_IMAGE_BYTES + 1)
|
||||
else:
|
||||
from agent.file_safety import raise_if_read_blocked
|
||||
raise_if_read_blocked(ref)
|
||||
with open(ref, "rb") as image_file:
|
||||
data = image_file.read(_MAX_IMAGE_BYTES + 1)
|
||||
if not data or len(data) > _MAX_IMAGE_BYTES:
|
||||
raise ValueError("reference image is empty or exceeds 20 MiB")
|
||||
return data
|
||||
|
||||
|
||||
class AthenaLocalImageProvider(ImageGenProvider):
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "athena-local"
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
return "Athena Local (FLUX.2 Klein)"
|
||||
|
||||
def is_available(self) -> bool:
|
||||
return bool(_api_key()) and _private_destination(_base_url())
|
||||
|
||||
def list_models(self) -> List[Dict[str, Any]]:
|
||||
return [{
|
||||
"id": _model(),
|
||||
"display": "FLUX.2 Klein 9B FP8 Beta on Athena",
|
||||
"speed": "local",
|
||||
"strengths": "Private local generation and multi-reference editing",
|
||||
"price": "local / no cloud",
|
||||
}]
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
return _model()
|
||||
|
||||
def capabilities(self) -> Dict[str, Any]:
|
||||
return {"modalities": ["text", "image"], "max_reference_images": 3}
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"name": "Athena Local (FLUX.2 Klein)",
|
||||
"badge": "local",
|
||||
"tag": "Private image generation on Athena; public endpoints are rejected",
|
||||
"env_vars": [
|
||||
{"key": "ATHENA_IMAGE_API_KEY", "prompt": "Athena router API key"},
|
||||
{
|
||||
"key": "ATHENA_IMAGE_BASE_URL",
|
||||
"prompt": "Athena image API base URL",
|
||||
"default": _DEFAULT_BASE_URL,
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
def generate(
|
||||
self,
|
||||
prompt: str,
|
||||
aspect_ratio: str = DEFAULT_ASPECT_RATIO,
|
||||
*,
|
||||
image_url: Optional[str] = None,
|
||||
reference_image_urls: Optional[List[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Dict[str, Any]:
|
||||
clean_prompt = (prompt or "").strip()
|
||||
aspect = resolve_aspect_ratio(aspect_ratio)
|
||||
base_url = _base_url()
|
||||
|
||||
if not clean_prompt:
|
||||
return error_response(
|
||||
error="Prompt is required.", error_type="invalid_argument",
|
||||
provider=self.name, aspect_ratio=aspect)
|
||||
if not _private_destination(base_url):
|
||||
return error_response(
|
||||
error=("Athena image endpoint is not a private-network "
|
||||
"destination; request blocked."),
|
||||
error_type="unsafe_destination", provider=self.name,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
model = _model()
|
||||
api_key = _api_key()
|
||||
if not api_key:
|
||||
return error_response(
|
||||
error=("No Athena router key is configured. Set "
|
||||
"ATHENA_IMAGE_API_KEY or reuse "
|
||||
"HERMES_CUSTOM_192_168_1_212_8081_API_KEY."),
|
||||
error_type="auth_required", provider=self.name, model=model,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
sources: List[str] = []
|
||||
if isinstance(image_url, str) and image_url.strip():
|
||||
sources.append(image_url.strip())
|
||||
sources.extend(normalize_reference_images(reference_image_urls) or [])
|
||||
sources = sources[:4]
|
||||
try:
|
||||
encoded_sources = [
|
||||
base64.b64encode(_load_private_image(source)).decode("ascii")
|
||||
for source in sources
|
||||
]
|
||||
except Exception as exc:
|
||||
return error_response(
|
||||
error=f"Reference image could not be loaded locally: {exc}",
|
||||
error_type="io_error", provider=self.name, model=model,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
request_data = {
|
||||
"model": model,
|
||||
"prompt": clean_prompt,
|
||||
"size": _SIZES[aspect],
|
||||
"n": 1,
|
||||
"quality": "standard",
|
||||
"steps": 4,
|
||||
"guidance": 1.0,
|
||||
"response_format": "b64_json",
|
||||
}
|
||||
endpoint = "generations"
|
||||
if encoded_sources:
|
||||
endpoint = "edits"
|
||||
request_data["image_b64"] = encoded_sources[0]
|
||||
request_data["reference_images_b64"] = encoded_sources[1:]
|
||||
request = urllib.request.Request(
|
||||
f"{base_url}/images/{endpoint}",
|
||||
data=json.dumps(request_data).encode("utf-8"), method="POST",
|
||||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
})
|
||||
|
||||
try:
|
||||
with urllib.request.urlopen(request, timeout=900) as response:
|
||||
result = json.load(response)
|
||||
except urllib.error.HTTPError as exc:
|
||||
try:
|
||||
detail = exc.read(4096).decode("utf-8", errors="replace")
|
||||
except Exception:
|
||||
detail = ""
|
||||
return error_response(
|
||||
error=f"Athena image request failed (HTTP {exc.code}): {detail[:500]}",
|
||||
error_type="api_error", provider=self.name, model=model,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
except (OSError, TimeoutError, ValueError, json.JSONDecodeError) as exc:
|
||||
return error_response(
|
||||
error=f"Athena image request failed: {exc}",
|
||||
error_type="connection_error", provider=self.name, model=model,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
items = result.get("data") if isinstance(result, dict) else None
|
||||
first = items[0] if isinstance(items, list) and items else None
|
||||
b64_data = first.get("b64_json") if isinstance(first, dict) else None
|
||||
if not isinstance(b64_data, str) or not b64_data:
|
||||
return error_response(
|
||||
error="Athena returned no image data.",
|
||||
error_type="empty_response", provider=self.name, model=model,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
try:
|
||||
saved = save_b64_image(b64_data, prefix="athena_flux2")
|
||||
except Exception as exc:
|
||||
return error_response(
|
||||
error=f"Generated image could not be saved: {exc}",
|
||||
error_type="io_error", provider=self.name, model=model,
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
return success_response(
|
||||
image=str(saved), model=model, prompt=clean_prompt,
|
||||
aspect_ratio=aspect, provider=self.name,
|
||||
modality="image" if encoded_sources else "text",
|
||||
extra={"size": _SIZES[aspect], "local_only": True,
|
||||
"reference_images": len(encoded_sources)})
|
||||
|
||||
|
||||
def register(ctx) -> None:
|
||||
ctx.register_image_gen_provider(AthenaLocalImageProvider())
|
||||
Reference in New Issue
Block a user