Spaces:
Sleeping
Sleeping
| """Duck-typed model provider configuration. | |
| Each ModelProvider is a plain data holder (name, api_url, api_token_env, params) | |
| plus a swappable `call` callable. There is no abstract base class / provider | |
| subclass hierarchy - any function matching the `call` signature | |
| `(provider: ModelProvider, image_path: str) -> bytes` can be attached to a | |
| provider instance, which is how providers with different API contracts | |
| (direct synchronous response vs. async job/poll) coexist without inheritance. | |
| """ | |
| from __future__ import annotations | |
| import base64 | |
| import io | |
| import json | |
| import logging | |
| import mimetypes | |
| import time | |
| from dataclasses import dataclass, field, replace | |
| from typing import Callable | |
| import requests | |
| from PIL import Image, ImageOps | |
| logger = logging.getLogger(__name__) | |
| # fal.ai's wan-i2v endpoint only supports these fixed aspect ratios; passing | |
| # "auto" can resolve to an unsupported computed size (422) for some input | |
| # image dimensions, so "auto" in a provider's params is a sentinel meaning | |
| # "pick the closest of these from the actual uploaded image" (see | |
| # _closest_supported_aspect_ratio below). | |
| _SUPPORTED_ASPECT_RATIOS = { | |
| "16:9": 16 / 9, | |
| "9:16": 9 / 16, | |
| "1:1": 1.0, | |
| } | |
| _MIME_TO_PIL_FORMAT = { | |
| "image/jpeg": "JPEG", | |
| "image/png": "PNG", | |
| "image/webp": "WEBP", | |
| } | |
| def _load_oriented_image(image_path: str) -> Image.Image: | |
| """Opens `image_path` and bakes in its EXIF orientation, if any. | |
| Phone photos are frequently stored with a raw pixel buffer in one | |
| orientation plus an EXIF "Orientation" tag telling viewers to rotate it | |
| (e.g. a portrait photo stored as landscape pixels + "rotate 90"). Image | |
| viewers and Gradio's own preview apply that tag automatically, but a | |
| naive `Image.open(...).size` read (and the raw bytes sent to fal.ai) | |
| does not - so without correcting for it here, the aspect ratio we | |
| compute and the pixels we send to the video model both end up | |
| reflecting the wrong (physical, not visual) orientation, producing a | |
| sideways video. | |
| """ | |
| image = Image.open(image_path) | |
| return ImageOps.exif_transpose(image) or image | |
| def _closest_supported_aspect_ratio(image_path: str) -> str: | |
| with _load_oriented_image(image_path) as image: | |
| width, height = image.size | |
| ratio = width / height if height else 1.0 | |
| return min(_SUPPORTED_ASPECT_RATIOS, key=lambda label: abs(_SUPPORTED_ASPECT_RATIOS[label] - ratio)) | |
| def _format_http_error(response: requests.Response) -> str: | |
| """Builds a concise, useful error from an HTTP response body.""" | |
| try: | |
| body = response.json() | |
| except ValueError: | |
| text = (response.text or "").strip() | |
| return text or f"HTTP {response.status_code}" | |
| if isinstance(body, dict): | |
| for key in ("detail", "error", "message"): | |
| value = body.get(key) | |
| if value: | |
| return str(value) | |
| return str(body) | |
| def _check_response(response: requests.Response, provider_name: str, context: str) -> None: | |
| """Raises RuntimeError with a short, UI-friendly message if `response` | |
| is not 2xx/3xx - but first logs full diagnostic detail (status line, | |
| response headers, and the full response body) at ERROR level. | |
| The exception message is deliberately short since it's shown directly | |
| in the UI's result card, but that's often not enough to actually debug | |
| a provider-side failure (e.g. which exact field/value it rejected) - | |
| the full request/response detail this logs is what you need for that, | |
| so check the server logs (not just the UI) when a provider call fails. | |
| """ | |
| if response.ok: | |
| return | |
| logger.error( | |
| "%s: %s failed - %s %s\nrequest: %s %s\nresponse headers: %s\nresponse body: %s", | |
| provider_name, | |
| context, | |
| response.status_code, | |
| response.reason, | |
| response.request.method if response.request else "?", | |
| response.url, | |
| dict(response.headers), | |
| (response.text or "")[:4000], | |
| ) | |
| detail = _format_http_error(response) | |
| raise RuntimeError(f"{context} failed for '{provider_name}' ({response.status_code}): {detail}") | |
| class ModelProvider: | |
| name: str | |
| api_url: str | |
| api_token_env: str | |
| api_token_value: str = "" | |
| params: dict = field(default_factory=dict) | |
| call: Callable[["ModelProvider", str], bytes] = None | |
| # Optional hook for providers whose params don't have a direct | |
| # "duration"/"durationSeconds" key (see apply_global_settings below) - | |
| # given (params, duration_seconds), returns updated params translating | |
| # the duration into whatever shape that provider's API expects. | |
| apply_duration: Callable[[dict, int], dict] | None = None | |
| # This provider's own valid "resolution" enum values, ordered lowest to | |
| # highest quality. Providers' resolution enums do NOT overlap uniformly | |
| # (e.g. ltx-2 only accepts 1080p/1440p/2160p - no 720p at all), so the | |
| # global resolution setting can't just be duck-typed in directly like | |
| # prompt can - see apply_global_settings's use of _nearest_resolution. | |
| resolution_choices: tuple[str, ...] | None = None | |
| # Names of additional (non-secret) per-user config values this provider | |
| # needs beyond api_token_env - e.g. Azure resources are deployed | |
| # per-user/per-resource, so a fixed api_url isn't enough; the user must | |
| # also supply their own resource endpoint URL. Every provider that isn't | |
| # a simple fixed-API-url + key needs this. Values are supplied via | |
| # extra_config_values (populated from the UI, same BYOK section as the | |
| # API key) and read back with .extra_config(name). | |
| extra_config_envs: tuple[str, ...] = () | |
| extra_config_values: dict = field(default_factory=dict) | |
| def api_token(self) -> str: | |
| """Returns the per-request API token supplied by the current browser user.""" | |
| return self.api_token_value or "" | |
| def extra_config(self, name: str) -> str: | |
| """Returns the per-request value of one of this provider's | |
| extra_config_envs, supplied by the current browser user.""" | |
| return self.extra_config_values.get(name, "") | |
| # Canonical quality ranking used to translate the UI's single global | |
| # resolution choice into whichever concrete value each provider actually | |
| # supports (see _nearest_resolution) - 2160p and 4k are treated as the same | |
| # tier since they're the same quality level under different provider naming. | |
| _RESOLUTION_RANK = {"720p": 0, "1080p": 1, "1440p": 2, "2160p": 3, "4k": 3} | |
| def _nearest_resolution(choices: tuple[str, ...], requested: str) -> str: | |
| """Picks the entry in `choices` whose quality tier is closest to the | |
| requested resolution's tier (e.g. requesting "720p" against a provider | |
| that only offers ("1080p", "1440p", "2160p") resolves to "1080p", its | |
| lowest/closest available option, instead of sending an unsupported | |
| value the provider's API would reject).""" | |
| requested_rank = _RESOLUTION_RANK.get(requested, 1) | |
| return min(choices, key=lambda choice: abs(_RESOLUTION_RANK.get(choice, 1) - requested_rank)) | |
| def apply_global_settings( | |
| provider: ModelProvider, | |
| prompt: str | None, | |
| duration_seconds: int | None, | |
| resolution: str | None, | |
| ) -> ModelProvider: | |
| """Returns a copy of `provider` with the UI's global settings (prompt, | |
| duration, resolution) applied, duck-typed against whatever params | |
| keys/hooks each provider actually supports. | |
| - prompt: overridden if the provider's params include a "prompt" key. | |
| - resolution: overridden if the provider's params include a | |
| "resolution" key, translated to the closest value in the provider's | |
| own `resolution_choices` (providers' resolution enums differ and | |
| don't all overlap, e.g. ltx-2 has no "720p" option at all - see | |
| _nearest_resolution). If a provider has no `resolution_choices` | |
| declared, the raw value is used as-is. | |
| - duration_seconds: applied via `provider.apply_duration` if the | |
| provider defines one (e.g. wan-i2v, which has no direct duration | |
| param and instead derives duration from num_frames/fps, or veo, | |
| which only accepts an exact 4/6/8 value); otherwise applied directly | |
| to whichever of "duration" / "durationSeconds" the provider's params | |
| actually has, if any. | |
| Providers that don't have a matching key/hook for a given setting are | |
| left untouched for that setting - this is what keeps a provider's | |
| unsupported fields from being silently (and incorrectly) injected. | |
| """ | |
| params = dict(provider.params) | |
| if prompt and "prompt" in params: | |
| params["prompt"] = prompt | |
| if resolution and "resolution" in params: | |
| if provider.resolution_choices: | |
| params["resolution"] = _nearest_resolution(provider.resolution_choices, resolution) | |
| else: | |
| params["resolution"] = resolution | |
| if duration_seconds is not None: | |
| if provider.apply_duration: | |
| params = provider.apply_duration(params, duration_seconds) | |
| elif "duration" in params: | |
| params["duration"] = duration_seconds | |
| elif "durationSeconds" in params: | |
| params["durationSeconds"] = duration_seconds | |
| return replace(provider, params=params) | |
| def generic_sync_call(provider: ModelProvider, image_path: str) -> bytes: | |
| """Template for providers with a direct synchronous response contract. | |
| POSTs the image as multipart/form-data along with `provider.params`, | |
| and returns the response body (expected to be video bytes) directly. | |
| """ | |
| headers = {} | |
| if provider.api_token: | |
| headers["Authorization"] = f"Bearer {provider.api_token}" | |
| with open(image_path, "rb") as image_file: | |
| response = requests.post( | |
| provider.api_url, | |
| headers=headers, | |
| data=provider.params, | |
| files={"image": image_file}, | |
| timeout=120, | |
| ) | |
| response.raise_for_status() | |
| return response.content | |
| def polling_call( | |
| provider: ModelProvider, | |
| image_path: str, | |
| poll_interval_seconds: float = 2.0, | |
| max_wait_seconds: float = 300.0, | |
| ) -> bytes: | |
| """Template for providers with an async job-queue contract (submit -> poll -> download). | |
| Expects: | |
| - POST provider.api_url returns JSON with a "status_url" (and optionally an | |
| immediate "output_url"). | |
| - Polling "status_url" returns JSON with a "status" field, becoming "succeeded" | |
| (with "output_url") or "failed" (with "error"). | |
| This is a template to copy/adapt per real provider contract, not a universal | |
| implementation - concrete async APIs differ in field names and shapes. | |
| """ | |
| headers = {} | |
| if provider.api_token: | |
| headers["Authorization"] = f"Bearer {provider.api_token}" | |
| with open(image_path, "rb") as image_file: | |
| submit_response = requests.post( | |
| provider.api_url, | |
| headers=headers, | |
| data=provider.params, | |
| files={"image": image_file}, | |
| timeout=30, | |
| ) | |
| submit_response.raise_for_status() | |
| job = submit_response.json() | |
| status_url = job["status_url"] | |
| deadline = time.monotonic() + max_wait_seconds | |
| while time.monotonic() < deadline: | |
| status_response = requests.get(status_url, headers=headers, timeout=30) | |
| status_response.raise_for_status() | |
| status = status_response.json() | |
| if status.get("status") == "succeeded": | |
| output_url = status["output_url"] | |
| video_response = requests.get(output_url, headers=headers, timeout=120) | |
| video_response.raise_for_status() | |
| return video_response.content | |
| if status.get("status") == "failed": | |
| raise RuntimeError(status.get("error", "Provider job failed")) | |
| time.sleep(poll_interval_seconds) | |
| raise TimeoutError(f"Provider '{provider.name}' job timed out after {max_wait_seconds}s") | |
| FAL_QUEUE_BASE = "https://queue.fal.run" | |
| def _encode_oriented_image(image_path: str) -> tuple[str, bytes]: | |
| """Returns (mime_type, image_bytes) for `image_path` with EXIF | |
| orientation baked in (see _load_oriented_image). Shared by every | |
| provider that needs the image as raw/base64 bytes rather than a | |
| multipart file upload. | |
| """ | |
| mime_type, _ = mimetypes.guess_type(image_path) | |
| mime_type = mime_type or "image/png" | |
| with _load_oriented_image(image_path) as image: | |
| pil_format = _MIME_TO_PIL_FORMAT.get(mime_type, image.format or "PNG") | |
| if pil_format == "JPEG" and image.mode in ("RGBA", "P"): | |
| image = image.convert("RGB") | |
| buffer = io.BytesIO() | |
| image.save(buffer, format=pil_format) | |
| return mime_type, buffer.getvalue() | |
| def _image_to_data_uri(image_path: str) -> str: | |
| mime_type, image_bytes = _encode_oriented_image(image_path) | |
| encoded = base64.b64encode(image_bytes).decode("ascii") | |
| return f"data:{mime_type};base64,{encoded}" | |
| def fal_queue_call( | |
| provider: ModelProvider, | |
| image_path: str, | |
| poll_interval_seconds: float = 2.0, | |
| max_wait_seconds: float = 600.0, | |
| ) -> bytes: | |
| """Calls a fal.ai queue-based model endpoint (submit -> poll -> fetch result). | |
| `provider.api_url` must be the fal app id, e.g. "fal-ai/ltx-2/image-to-video". | |
| `provider.params` is sent as the request body (merged with the uploaded | |
| image as a base64 data URI, since fal's queue API expects a publicly | |
| accessible URL or a data URI rather than a multipart file upload). | |
| See https://docs.fal.ai for the queue submit/status/result contract. | |
| """ | |
| headers = { | |
| "Authorization": f"Key {provider.api_token}", | |
| "Content-Type": "application/json", | |
| } | |
| payload = {**provider.params, "image_url": _image_to_data_uri(image_path)} | |
| if payload.get("aspect_ratio") == "auto": | |
| payload["aspect_ratio"] = _closest_supported_aspect_ratio(image_path) | |
| logger.debug("fal.ai %s: submitting job to %s", provider.name, provider.api_url) | |
| submit_response = requests.post( | |
| f"{FAL_QUEUE_BASE}/{provider.api_url}", | |
| headers=headers, | |
| json=payload, | |
| timeout=30, | |
| ) | |
| _check_response(submit_response, provider.name, "fal.ai job submission") | |
| submission = submit_response.json() | |
| status_url = submission["status_url"] | |
| response_url = submission["response_url"] | |
| deadline = time.monotonic() + max_wait_seconds | |
| while time.monotonic() < deadline: | |
| status_response = requests.get(status_url, headers=headers, timeout=30) | |
| _check_response(status_response, provider.name, "fal.ai status poll") | |
| status = status_response.json() | |
| if status.get("status") == "COMPLETED": | |
| logger.debug("fal.ai %s: job completed", provider.name) | |
| break | |
| if status.get("status") == "FAILED": | |
| logger.error("fal.ai %s: job failed: %s", provider.name, status) | |
| raise RuntimeError(f"fal.ai job failed for '{provider.name}': {status}") | |
| time.sleep(poll_interval_seconds) | |
| else: | |
| logger.error("fal.ai %s: job timed out after %.0fs", provider.name, max_wait_seconds) | |
| raise TimeoutError(f"fal.ai job for '{provider.name}' timed out after {max_wait_seconds}s") | |
| result_response = requests.get(response_url, headers=headers, timeout=30) | |
| _check_response(result_response, provider.name, "fal.ai result fetch") | |
| result = result_response.json() | |
| video_url = result["video"]["url"] | |
| video_response = requests.get(video_url, timeout=180) | |
| _check_response(video_response, provider.name, "fal.ai video download") | |
| return video_response.content | |
| REPLICATE_API_BASE = "https://api.replicate.com/v1" | |
| def _extract_replicate_output_url(output) -> str | None: | |
| """Normalizes Replicate's `output` field to a single file URL. | |
| Different model schemas represent a single output file as a bare | |
| string URL, a one-item list, or (rarely) a dict with a "url"/"video" | |
| key - this isn't a universal contract, just the shapes seen in | |
| practice, so adapt if a new provider's schema differs. | |
| """ | |
| if isinstance(output, str): | |
| return output | |
| if isinstance(output, list) and output: | |
| return _extract_replicate_output_url(output[0]) | |
| if isinstance(output, dict): | |
| return output.get("url") or output.get("video") | |
| return None | |
| def replicate_call( | |
| provider: ModelProvider, | |
| image_path: str, | |
| poll_interval_seconds: float = 2.0, | |
| max_wait_seconds: float = 600.0, | |
| ) -> bytes: | |
| """Calls a Replicate model via its official REST API (create prediction | |
| -> poll -> fetch output), authenticated with the user's own Replicate | |
| API token (BYOK). | |
| `provider.api_url` must be the Replicate model id, e.g. | |
| "alibaba/happyhorse-1.1". `provider.params` is sent as the prediction's | |
| "input", merged with the uploaded image as a base64 data URI in an | |
| "images" array - Replicate's API accepts a data URI directly for | |
| file-type inputs, so no public upload step is needed. | |
| See https://replicate.com/docs/reference/http for the create/poll | |
| prediction contract. | |
| """ | |
| headers = { | |
| "Authorization": f"Bearer {provider.api_token}", | |
| "Content-Type": "application/json", | |
| } | |
| payload = {"input": {**provider.params, "images": [_image_to_data_uri(image_path)]}} | |
| logger.debug("replicate %s: creating prediction for %s", provider.name, provider.api_url) | |
| create_response = requests.post( | |
| f"{REPLICATE_API_BASE}/models/{provider.api_url}/predictions", | |
| headers=headers, | |
| json=payload, | |
| timeout=30, | |
| ) | |
| _check_response(create_response, provider.name, "replicate prediction creation") | |
| prediction = create_response.json() | |
| status_url = prediction["urls"]["get"] | |
| deadline = time.monotonic() + max_wait_seconds | |
| while True: | |
| status_response = requests.get(status_url, headers=headers, timeout=30) | |
| _check_response(status_response, provider.name, "replicate status poll") | |
| prediction = status_response.json() | |
| status = prediction.get("status") | |
| if status == "succeeded": | |
| logger.debug("replicate %s: prediction succeeded", provider.name) | |
| break | |
| if status in ("failed", "canceled"): | |
| logger.error("replicate %s: prediction %s: %s", provider.name, status, prediction.get("error")) | |
| raise RuntimeError(f"replicate prediction {status} for '{provider.name}': {prediction.get('error')}") | |
| if time.monotonic() >= deadline: | |
| logger.error("replicate %s: prediction timed out after %.0fs", provider.name, max_wait_seconds) | |
| raise TimeoutError(f"replicate prediction for '{provider.name}' timed out after {max_wait_seconds}s") | |
| time.sleep(poll_interval_seconds) | |
| video_url = _extract_replicate_output_url(prediction.get("output")) | |
| if not video_url: | |
| raise RuntimeError(f"replicate prediction for '{provider.name}' succeeded but returned no output URL") | |
| # Per Replicate's docs, output file URLs require the Authorization header | |
| # to fetch, unlike fal.ai's (unauthenticated) delivery URLs above. | |
| video_response = requests.get(video_url, headers=headers, timeout=180) | |
| _check_response(video_response, provider.name, "replicate video download") | |
| return video_response.content | |
| GEMINI_API_BASE = "https://generativelanguage.googleapis.com/v1beta" | |
| def _extract_veo_video_uri(node) -> str | None: | |
| """Recursively searches a Gemini API operation response for a video | |
| file URI. | |
| The documented Instance/Params schema for Veo on the Gemini API is | |
| stable (see VideoGenerationModelInstance/-Params, shared with Vertex), | |
| but the exact nesting of the *output* file reference inside | |
| `operation.response` for the API-key-based Gemini API (as opposed to | |
| Vertex's GCS-only `VideoGenerationModelResult.gcsUris`) wasn't | |
| confirmable from public docs at the time this was written - recursing | |
| for any `{"video": {"uri": ...}}` (or bare `{"uri": ...}`) shape is a | |
| deliberate hedge against minor key-nesting differences, not a documented | |
| contract. | |
| """ | |
| if isinstance(node, dict): | |
| video = node.get("video") | |
| if isinstance(video, dict) and video.get("uri"): | |
| return video["uri"] | |
| if "uri" in node and isinstance(node["uri"], str): | |
| return node["uri"] | |
| for value in node.values(): | |
| found = _extract_veo_video_uri(value) | |
| if found: | |
| return found | |
| elif isinstance(node, list): | |
| for item in node: | |
| found = _extract_veo_video_uri(item) | |
| if found: | |
| return found | |
| return None | |
| def _extract_veo_filter_reason(node) -> str | None: | |
| """Recursively searches a Gemini API operation response for a | |
| safety-filter rejection reason (e.g. `raiMediaFilteredReasons`), which | |
| is the most common real-world cause of a "done" operation with no | |
| video URI - the prompt/image got silently filtered rather than the | |
| response shape being unexpected.""" | |
| if isinstance(node, dict): | |
| reasons = node.get("raiMediaFilteredReasons") | |
| if reasons: | |
| return "; ".join(str(reason) for reason in reasons) | |
| for value in node.values(): | |
| found = _extract_veo_filter_reason(value) | |
| if found: | |
| return found | |
| elif isinstance(node, list): | |
| for item in node: | |
| found = _extract_veo_filter_reason(item) | |
| if found: | |
| return found | |
| return None | |
| def veo_call( | |
| provider: ModelProvider, | |
| image_path: str, | |
| poll_interval_seconds: float = 10.0, | |
| max_wait_seconds: float = 600.0, | |
| ) -> bytes: | |
| """Calls Google's Veo image-to-video model via the Gemini API (create | |
| long-running prediction -> poll -> download), authenticated with the | |
| user's own Gemini API key (BYOK). | |
| `provider.api_url` must be the Veo model id, e.g. | |
| "veo-3.1-generate-preview". `provider.params` is sent as the | |
| prediction's "parameters" (except "prompt", which goes in "instances" | |
| alongside the uploaded image). See | |
| https://ai.google.dev/gemini-api/docs/veo for the documented | |
| prompt/image/parameter contract; this uses the generic | |
| models.predictLongRunning REST method (see | |
| https://ai.google.dev/api/models#method:-models.predictlongrunning), | |
| which every Gemini API model (including Veo) shares. | |
| """ | |
| headers = { | |
| "x-goog-api-key": provider.api_token, | |
| "Content-Type": "application/json", | |
| } | |
| mime_type, image_bytes = _encode_oriented_image(image_path) | |
| params = dict(provider.params) | |
| prompt = params.pop("prompt", "") | |
| instance = { | |
| "prompt": prompt, | |
| "image": { | |
| "mimeType": mime_type, | |
| "bytesBase64Encoded": base64.b64encode(image_bytes).decode("ascii"), | |
| }, | |
| } | |
| payload = {"instances": [instance], "parameters": params} | |
| logger.debug("veo %s: creating prediction for model %s", provider.name, provider.api_url) | |
| create_response = requests.post( | |
| f"{GEMINI_API_BASE}/models/{provider.api_url}:predictLongRunning", | |
| headers=headers, | |
| json=payload, | |
| timeout=30, | |
| ) | |
| _check_response(create_response, provider.name, "veo prediction creation") | |
| operation = create_response.json() | |
| operation_name = operation["name"] | |
| deadline = time.monotonic() + max_wait_seconds | |
| while not operation.get("done"): | |
| if time.monotonic() >= deadline: | |
| logger.error("veo %s: operation timed out after %.0fs", provider.name, max_wait_seconds) | |
| raise TimeoutError(f"veo operation for '{provider.name}' timed out after {max_wait_seconds}s") | |
| time.sleep(poll_interval_seconds) | |
| status_response = requests.get(f"{GEMINI_API_BASE}/{operation_name}", headers=headers, timeout=30) | |
| _check_response(status_response, provider.name, "veo operation poll") | |
| operation = status_response.json() | |
| if operation.get("error"): | |
| logger.error("veo %s: operation failed: %s", provider.name, operation["error"]) | |
| raise RuntimeError(f"veo operation failed for '{provider.name}': {operation['error']}") | |
| video_uri = _extract_veo_video_uri(operation.get("response")) | |
| if not video_uri: | |
| # Log the full response so this is actually debuggable next time - | |
| # the shape of a "done" operation with no video can vary (safety | |
| # filtering, quota/partial failures, an undocumented response | |
| # nesting, etc.) and a bare exception message throws that | |
| # information away. | |
| logger.error( | |
| "veo %s: operation done but no video URI found; full response=%s", | |
| provider.name, | |
| json.dumps(operation, indent=2)[:4000], | |
| ) | |
| filter_reason = _extract_veo_filter_reason(operation.get("response")) | |
| if filter_reason: | |
| raise RuntimeError( | |
| f"veo operation for '{provider.name}' was rejected by Google's safety filters: {filter_reason}" | |
| ) | |
| raise RuntimeError(f"veo operation for '{provider.name}' succeeded but returned no video URI") | |
| video_response = requests.get(video_uri, headers=headers, timeout=180) | |
| _check_response(video_response, provider.name, "veo video download") | |
| return video_response.content | |
| def _wan_apply_duration(params: dict, duration_seconds: int) -> dict: | |
| """wan-i2v has no direct duration param - it derives duration from | |
| num_frames / frames_per_second, so translate the global duration | |
| (seconds) setting into num_frames using its own fps.""" | |
| params = dict(params) | |
| fps = params.get("frames_per_second", 16) | |
| params["num_frames"] = max(1, round(duration_seconds * fps)) | |
| return params | |
| def _veo_apply_duration(params: dict, duration_seconds: int) -> dict: | |
| """Veo only accepts an exact durationSeconds of 4, 6, or 8 (confirmed | |
| against the live API - a 400 otherwise) - snap the requested global | |
| duration to the nearest supported value.""" | |
| allowed = (4, 6, 8) | |
| params = dict(params) | |
| params["durationSeconds"] = min(allowed, key=lambda value: abs(value - duration_seconds)) | |
| return params | |
| def _azure_sora_apply_duration(params: dict, duration_seconds: int) -> dict: | |
| """Azure Sora-2's duration field is "seconds" (a string), not | |
| "duration"/"durationSeconds". Confirmed live against the API that it | |
| only accepts exactly 4, 8, or 12 - any other value (e.g. "6") is | |
| rejected with a 400 ("Invalid value... Supported values are: '4', '8', | |
| and '12'.") - so snap the requested global duration to the nearest | |
| supported value, the same way _veo_apply_duration does.""" | |
| allowed = (4, 8, 12) | |
| params = dict(params) | |
| params["seconds"] = str(min(allowed, key=lambda value: abs(value - duration_seconds))) | |
| return params | |
| def _orient_size(size: str, image: Image.Image) -> str: | |
| """Swaps width/height in `size` ("WIDTHxHEIGHT") if needed so its | |
| orientation (portrait/landscape) matches `image`'s, keeping the same | |
| resolution class (e.g. a 1280x720 default becomes 720x1280 for a | |
| portrait photo).""" | |
| width, height = (int(part) for part in size.lower().split("x")) | |
| if (image.height > image.width) != (height > width): | |
| width, height = height, width | |
| return f"{width}x{height}" | |
| def _encode_image_for_sora(image_path: str, size: str) -> tuple[str, bytes]: | |
| """Like _encode_oriented_image, but also resizes/center-crops the image | |
| to exactly match `size` ("WIDTHxHEIGHT"). | |
| Confirmed live against Azure's Sora-2 API: it rejects `input_reference` | |
| with "Inpaint image must match the requested width and height" unless | |
| the uploaded image's pixel dimensions are an exact match for the | |
| requested output `size` - real user photos essentially never happen to | |
| be exactly 1280x720/720x1280 already, so this must be done | |
| unconditionally, not just as a fallback. | |
| """ | |
| target_width, target_height = (int(part) for part in size.lower().split("x")) | |
| with _load_oriented_image(image_path) as image: | |
| scale = max(target_width / image.width, target_height / image.height) | |
| resized = image.resize( | |
| (round(image.width * scale), round(image.height * scale)), Image.LANCZOS | |
| ) | |
| left = (resized.width - target_width) // 2 | |
| top = (resized.height - target_height) // 2 | |
| cropped = resized.crop((left, top, left + target_width, top + target_height)) | |
| if cropped.mode in ("RGBA", "P"): | |
| cropped = cropped.convert("RGB") | |
| buffer = io.BytesIO() | |
| cropped.save(buffer, format="PNG") | |
| return "image/png", buffer.getvalue() | |
| def azure_sora_call( | |
| provider: ModelProvider, | |
| image_path: str, | |
| poll_interval_seconds: float = 5.0, | |
| max_wait_seconds: float = 600.0, | |
| ) -> bytes: | |
| """Calls Azure's Sora-2 video API (upload reference image -> create -> | |
| poll -> download content), authenticated with the user's own Azure AI | |
| Foundry resource (BYOK). | |
| Unlike every other provider here, an Azure Sora deployment is tied to a | |
| specific user resource, so a fixed `api_url` + key isn't enough - the | |
| user must also supply their own resource's video endpoint URL (e.g. | |
| "https://{resource}.openai.azure.com/openai/v1"), collected via | |
| `extra_config_envs`/`extra_config()` and entered in the same BYOK | |
| section as the API key. | |
| `provider.api_url` is the model id ("sora-2"), sent as the "model" | |
| field. `provider.params` holds "size" (e.g. "1280x720") and "seconds" | |
| (a string, e.g. "4"). | |
| This provider's actual contract was only discoverable by testing live | |
| against a real resource (its behavior didn't match either the generic | |
| curl example that inspired it, or Azure's separate/incompatible | |
| "video/generations/jobs" API which 404s on this resource entirely) - | |
| confirmed end-to-end (upload/create/poll/content all succeeded | |
| producing a real queued job): | |
| 1. POST {endpoint}/files (multipart, purpose="assistants") to upload | |
| the reference image, returning a file id. Sending the image | |
| directly as a multipart "input_reference" file on the create | |
| call (as fal/replicate/veo do for their own image fields) is | |
| REJECTED with "Invalid type for 'input_reference': expected an | |
| object, but got a file instead" - input_reference must instead be | |
| a JSON object referencing an uploaded file's id. | |
| 2. POST {endpoint}/videos as `application/json` (NOT multipart) with | |
| "input_reference": {"file_id": <id from step 1>}. The referenced | |
| image's pixel dimensions must exactly match the request's "size" | |
| ("Inpaint image must match the requested width and height"), so | |
| the image is resized/center-cropped to match first - see | |
| _orient_size/_encode_image_for_sora. | |
| """ | |
| endpoint = provider.extra_config("AZURE_SORA_ENDPOINT").rstrip("/") | |
| if not endpoint: | |
| raise RuntimeError( | |
| f"'{provider.name}' is missing its Azure endpoint URL - set it in the API Keys section." | |
| ) | |
| # Users commonly copy just the bare resource URL (e.g. | |
| # "https://{resource}.openai.azure.com" or ".../api/projects/{name}") | |
| # rather than the full video-API base - both confirmed live to 404 | |
| # ("Resource not found") without the "/openai/v1" segment, so normalize | |
| # it here instead of silently failing on an easy-to-make input mistake. | |
| if not endpoint.endswith("/openai/v1"): | |
| endpoint = f"{endpoint}/openai/v1" | |
| headers = {"Authorization": f"Bearer {provider.api_token}"} | |
| with _load_oriented_image(image_path) as oriented_image: | |
| size = _orient_size(provider.params.get("size", "1280x720"), oriented_image) | |
| mime_type, image_bytes = _encode_image_for_sora(image_path, size) | |
| extension = mimetypes.guess_extension(mime_type) or ".png" | |
| filename = f"reference{extension}" | |
| logger.debug("azure sora %s: uploading reference image", provider.name) | |
| upload_response = requests.post( | |
| f"{endpoint}/files", | |
| headers=headers, | |
| data={"purpose": "assistants"}, | |
| files={"file": (filename, image_bytes, mime_type)}, | |
| timeout=60, | |
| ) | |
| _check_response(upload_response, provider.name, "azure sora reference image upload") | |
| file_id = upload_response.json()["id"] | |
| body = {key: value for key, value in provider.params.items() if key != "prompt"} | |
| body["model"] = provider.api_url | |
| body["prompt"] = provider.params.get("prompt", "") | |
| body["size"] = size | |
| body["input_reference"] = {"file_id": file_id} | |
| logger.debug("azure sora %s: creating video job", provider.name) | |
| logger.debug("azure sora %s: request body=%s", provider.name, {**body, "input_reference": file_id}) | |
| create_response = requests.post( | |
| f"{endpoint}/videos", | |
| headers={**headers, "Content-Type": "application/json"}, | |
| json=body, | |
| timeout=60, | |
| ) | |
| _check_response(create_response, provider.name, "azure sora video creation") | |
| job = create_response.json() | |
| video_id = job["id"] | |
| deadline = time.monotonic() + max_wait_seconds | |
| while True: | |
| status_response = requests.get(f"{endpoint}/videos/{video_id}", headers=headers, timeout=30) | |
| _check_response(status_response, provider.name, "azure sora status poll") | |
| job = status_response.json() | |
| status = job.get("status") | |
| if status == "completed": | |
| logger.debug("azure sora %s: video completed", provider.name) | |
| break | |
| if status == "failed": | |
| logger.error("azure sora %s: video failed: %s", provider.name, job.get("error")) | |
| raise RuntimeError(f"azure sora video failed for '{provider.name}': {job.get('error')}") | |
| if time.monotonic() >= deadline: | |
| logger.error("azure sora %s: video timed out after %.0fs", provider.name, max_wait_seconds) | |
| raise TimeoutError(f"azure sora video for '{provider.name}' timed out after {max_wait_seconds}s") | |
| time.sleep(poll_interval_seconds) | |
| content_response = requests.get(f"{endpoint}/videos/{video_id}/content", headers=headers, timeout=180) | |
| _check_response(content_response, provider.name, "azure sora video download") | |
| return content_response.content | |
| # Registry of configured providers. The fal.ai providers use fal's queue API | |
| # (submit -> poll -> fetch result) via the shared `fal_queue_call`; the | |
| # Replicate provider uses Replicate's REST API via `replicate_call`; the Veo | |
| # provider uses the Gemini API's predictLongRunning method via `veo_call`. | |
| # Every provider is BYOK - the user's API key is supplied at request time in | |
| # the UI, never read from the server environment. | |
| PROVIDERS: list[ModelProvider] = [ | |
| ModelProvider( | |
| name="ltx-2-fal", | |
| api_url="fal-ai/ltx-2/image-to-video", | |
| api_token_env="FAL_KEY", | |
| params={ | |
| "prompt": "Animate this image with natural, smooth motion.", | |
| "duration": 6, | |
| "resolution": "1080p", | |
| "fps": 25, | |
| "generate_audio": True, | |
| }, | |
| # Confirmed against the live API: ltx-2 does NOT accept "720p" at | |
| # all (422 "Input should be '1080p', '1440p' or '2160p'"). | |
| resolution_choices=("1080p", "1440p", "2160p"), | |
| call=fal_queue_call, | |
| ), | |
| ModelProvider( | |
| name="wan2.1-i2v-720p-fal", | |
| api_url="fal-ai/wan-i2v", | |
| api_token_env="FAL_KEY", | |
| params={ | |
| "prompt": "Animate this image with natural, smooth motion.", | |
| "resolution": "720p", | |
| "num_frames": 81, | |
| "frames_per_second": 16, | |
| # Resolved to a concrete supported ratio (16:9 / 9:16 / 1:1) from | |
| # the actual uploaded image at call time - see | |
| # _closest_supported_aspect_ratio. | |
| "aspect_ratio": "auto", | |
| }, | |
| resolution_choices=("720p", "1080p"), | |
| apply_duration=_wan_apply_duration, | |
| call=fal_queue_call, | |
| ), | |
| ModelProvider( | |
| name="happyhorse-1.1-replicate", | |
| api_url="alibaba/happyhorse-1.1", | |
| api_token_env="REPLICATE_API_TOKEN", | |
| params={ | |
| "prompt": "Animate this image with natural, smooth motion.", | |
| "resolution": "1080p", | |
| "duration": 5, | |
| # No aspect_ratio param: per the model's schema, aspect_ratio only | |
| # applies to text-to-video/reference-to-video - for image-to-video | |
| # (single image, our case) the image's own aspect ratio is used. | |
| }, | |
| resolution_choices=("720p", "1080p"), | |
| call=replicate_call, | |
| ), | |
| ModelProvider( | |
| name="veo-3.1-gemini", | |
| api_url="veo-3.1-generate-preview", | |
| api_token_env="GEMINI_API_KEY", | |
| params={ | |
| "prompt": "Animate this image with natural, smooth motion.", | |
| "resolution": "720p", | |
| # Despite being shown as quoted enum values ("4"/"6"/"8") in | |
| # Google's docs table, the API rejects a string here with a 400 | |
| # ("value type for durationSeconds needs to be a number") - | |
| # confirmed against the live API - must be a JSON number. | |
| "durationSeconds": 8, | |
| # Image-to-video only supports "allow_adult" (not "allow_all") | |
| # per the Veo API parameter table. | |
| "personGeneration": "allow_adult", | |
| # No aspectRatio param: when a single starting image is provided, | |
| # Veo uses the image's own aspect ratio. | |
| }, | |
| resolution_choices=("720p", "1080p", "4k"), | |
| apply_duration=_veo_apply_duration, | |
| call=veo_call, | |
| ), | |
| ModelProvider( | |
| name="sora-2-azure", | |
| api_url="sora-2", | |
| api_token_env="AZURE_SORA_API_KEY", | |
| # Azure Sora deployments are per-resource - the user must also | |
| # supply their own resource's video API endpoint (e.g. | |
| # "https://{resource}.openai.azure.com/openai/v1"), collected in | |
| # the same BYOK section as the API key. | |
| extra_config_envs=("AZURE_SORA_ENDPOINT",), | |
| params={ | |
| "prompt": "Animate this image with natural, smooth motion.", | |
| "size": "1280x720", | |
| "seconds": "4", | |
| }, | |
| apply_duration=_azure_sora_apply_duration, | |
| call=azure_sora_call, | |
| ), | |
| ] | |