261 lines
9.5 KiB
Python
261 lines
9.5 KiB
Python
"""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())
|