Add cloud sprite providers and browser demos
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
"""Client for Polza.ai's media API.
|
||||
|
||||
The client deliberately uses the standard library so enabling cloud generation
|
||||
does not add a runtime dependency to the local Diffusers installation.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
import mimetypes
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.request import Request, urlopen
|
||||
|
||||
|
||||
API_URL = "https://polza.ai/api/v1"
|
||||
MAX_REFERENCE_BYTES = 50 * 1024 * 1024
|
||||
|
||||
|
||||
class PolzaError(RuntimeError):
|
||||
"""An API or transport error returned by Polza.ai."""
|
||||
|
||||
|
||||
class PolzaClient:
|
||||
def __init__(self, api_key: Optional[str] = None, api_url: str = API_URL):
|
||||
self.api_key = api_key or os.environ.get("POLZA_API_KEY")
|
||||
self.api_url = api_url.rstrip("/")
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
if not self.api_key:
|
||||
raise PolzaError(
|
||||
"POLZA_API_KEY is not configured. Set it in the MCP server environment."
|
||||
)
|
||||
return {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
def _request(self, method: str, path: str, payload: Optional[dict] = None) -> dict:
|
||||
body = json.dumps(payload).encode("utf-8") if payload is not None else None
|
||||
request = Request(
|
||||
f"{self.api_url}{path}", data=body, headers=self._headers(), method=method
|
||||
)
|
||||
try:
|
||||
with urlopen(request, timeout=125) as response:
|
||||
result = json.loads(response.read().decode("utf-8"))
|
||||
except HTTPError as exc:
|
||||
detail = exc.read().decode("utf-8", errors="replace")
|
||||
try:
|
||||
detail = json.loads(detail).get("error", {}).get("message", detail)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
raise PolzaError(f"Polza API returned HTTP {exc.code}: {detail}") from exc
|
||||
except (URLError, TimeoutError) as exc:
|
||||
raise PolzaError(f"Could not reach Polza API: {exc}") from exc
|
||||
|
||||
if result.get("error"):
|
||||
error = result["error"]
|
||||
raise PolzaError(error.get("message", str(error)))
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def reference_payload(reference: str) -> dict[str, str]:
|
||||
"""Convert an HTTPS URL, data URI, or local image path for Media API."""
|
||||
if reference.startswith(("https://", "http://")):
|
||||
return {"type": "url", "data": reference}
|
||||
if reference.startswith("data:"):
|
||||
return {"type": "base64", "data": reference}
|
||||
|
||||
path = Path(reference).expanduser()
|
||||
if not path.is_file():
|
||||
raise PolzaError(f"Reference image does not exist: {reference}")
|
||||
if path.stat().st_size > MAX_REFERENCE_BYTES:
|
||||
raise PolzaError(f"Reference image exceeds the 50 MB API limit: {reference}")
|
||||
mime_type = mimetypes.guess_type(path.name)[0] or "application/octet-stream"
|
||||
encoded = base64.b64encode(path.read_bytes()).decode("ascii")
|
||||
return {"type": "base64", "data": f"data:{mime_type};base64,{encoded}"}
|
||||
|
||||
def create_image(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
prompt: str,
|
||||
reference_images: Optional[list[str]] = None,
|
||||
count: int = 1,
|
||||
aspect_ratio: Optional[str] = None,
|
||||
seed: Optional[int] = None,
|
||||
quality: Optional[str] = None,
|
||||
output_format: str = "png",
|
||||
background: Optional[str] = None,
|
||||
provider: Optional[dict] = None,
|
||||
wait: bool = True,
|
||||
) -> dict:
|
||||
if not 1 <= count <= 10:
|
||||
raise PolzaError("count must be between 1 and 10")
|
||||
input_data: dict[str, Any] = {
|
||||
"prompt": prompt,
|
||||
"max_images": count,
|
||||
"output_format": output_format,
|
||||
}
|
||||
if reference_images:
|
||||
input_data["images"] = [self.reference_payload(ref) for ref in reference_images]
|
||||
for name, value in {
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"seed": seed,
|
||||
"quality": quality,
|
||||
"background": background,
|
||||
}.items():
|
||||
if value is not None:
|
||||
input_data[name] = value
|
||||
|
||||
payload: dict[str, Any] = {"model": model, "input": input_data, "async": not wait}
|
||||
if provider:
|
||||
payload["provider"] = provider
|
||||
return self._request("POST", "/media", payload)
|
||||
|
||||
def get_status(self, generation_id: str) -> dict:
|
||||
return self._request("GET", f"/media/{generation_id}")
|
||||
|
||||
def wait_for_completion(
|
||||
self, generation_id: str, timeout_seconds: int = 115, poll_seconds: float = 2
|
||||
) -> dict:
|
||||
deadline = time.monotonic() + timeout_seconds
|
||||
while True:
|
||||
status = self.get_status(generation_id)
|
||||
state = status.get("status", "").lower()
|
||||
if state in {"completed", "succeeded", "success"}:
|
||||
return status
|
||||
if state in {"failed", "error", "cancelled", "canceled"}:
|
||||
error = status.get("error") or {}
|
||||
raise PolzaError(error.get("message", f"Generation {generation_id} {state}"))
|
||||
if time.monotonic() >= deadline:
|
||||
return status
|
||||
time.sleep(poll_seconds)
|
||||
|
||||
@staticmethod
|
||||
def image_sources(response: dict) -> list[str]:
|
||||
"""Extract CDN URLs or base64 images from documented response variants."""
|
||||
found: list[str] = []
|
||||
|
||||
def visit(value: Any) -> None:
|
||||
if isinstance(value, dict):
|
||||
for key, item in value.items():
|
||||
if key in {"url", "b64_json", "base64"} and isinstance(item, str):
|
||||
found.append(item)
|
||||
elif key in {"data", "output", "images", "result", "results"}:
|
||||
visit(item)
|
||||
elif isinstance(value, list):
|
||||
for item in value:
|
||||
visit(item)
|
||||
|
||||
visit(response)
|
||||
return list(dict.fromkeys(found))
|
||||
Reference in New Issue
Block a user