mirror of
https://github.com/calesthio/OpenMontage.git
synced 2026-08-17 05:56:29 +08:00
Merge remote-tracking branch 'origin/main' into codex/add-gemini-seedance-minimax-models
# Conflicts: # docs/PROVIDERS.md
This commit is contained in:
89
.agents/skills/atlas-cloud/SKILL.md
Normal file
89
.agents/skills/atlas-cloud/SKILL.md
Normal file
@@ -0,0 +1,89 @@
|
||||
---
|
||||
name: atlas-cloud
|
||||
description: Generate or edit images and videos through the Atlas Cloud gateway. Use for Atlas-hosted Seedance 2.5/2.0, Gemini Omni Flash, MiniMax H3, Seedream 5.0, GPT Image 2, Nano Banana 2, or when one ATLASCLOUD_API_KEY should access multiple media model families.
|
||||
---
|
||||
|
||||
# Atlas Cloud
|
||||
|
||||
Route complete productions through `image_selector` or `video_selector`. Call
|
||||
`atlas_image` or `atlas_video` directly when the user names Atlas Cloud or an
|
||||
exact Atlas model. Never substitute a direct vendor endpoint for an Atlas request.
|
||||
|
||||
Set `ATLASCLOUD_API_KEY`. The aliases `ATLAS_CLOUD_API_KEY` and `ATLAS_API_KEY`
|
||||
are accepted for compatibility.
|
||||
|
||||
## Preflight every paid call
|
||||
|
||||
1. Read `tool.get_info()["model_catalog"]`; do not infer availability from a
|
||||
search collection or invent a task suffix.
|
||||
2. Select an exact model id and operation supported by that catalog.
|
||||
3. Announce the Atlas tool, provider, exact model, request count, and estimated
|
||||
cost before submitting.
|
||||
4. Use local-path inputs when convenient. The tool uploads them through Atlas.
|
||||
5. Save outputs inside the active `projects/<project-id>/` tree and inspect them.
|
||||
|
||||
## Video routes
|
||||
|
||||
| Family | Operations | Exact route notes |
|
||||
|---|---|---|
|
||||
| Seedance 2.5 | text, image, reference | `bytedance/seedance-2.5/{text,image,reference}-to-video`; image mode uses `image` plus optional `last_image`; reference mode accepts up to 30 images, 10 videos, and 10 audios (including audio-only); 4–30s; $0.134/s |
|
||||
| Seedance 2.0 | text, image, reference | `bytedance/seedance-2.0/{text,image,reference}-to-video`; 4–15s; $0.112/s |
|
||||
| Gemini Omni Flash | text, image, reference, video edit | Standard routes use `google/gemini-omni-flash/...`; developer routes exist for text/image/reference only. Standard image mode uses one `image`; standard reference uses `images`; developer reference requires one `video_clips` object. |
|
||||
| MiniMax H3 | text, image, reference | `minimax/h3/{text,image,reference}-to-video`; image mode supports optional `end_image`; reference mode requires `refers` objects; 4–15s; 768P/2K; $0.10/s |
|
||||
|
||||
Use canonical OpenMontage fields:
|
||||
|
||||
```python
|
||||
tool.execute({
|
||||
"prompt": "...",
|
||||
"model": "minimax/h3/reference-to-video",
|
||||
"operation": "reference_to_video",
|
||||
"duration": 10,
|
||||
"resolution": "2K",
|
||||
"reference_images": ["projects/demo/character.png"],
|
||||
"reference_audios": ["projects/demo/performance.wav"],
|
||||
"output_path": "projects/demo/h3.mp4",
|
||||
})
|
||||
```
|
||||
|
||||
For Gemini developer reference video, pass `video_clips` objects with `url`,
|
||||
`start`, and `ends`. For H3, callers may pass normalized `reference_images`,
|
||||
`reference_videos`, and `reference_audios`; the tool converts them to `refers`.
|
||||
|
||||
Do not silently clamp duration, resolution, or ratio. Invalid values must fail
|
||||
before billing with the supported choices in the error.
|
||||
|
||||
## Image routes
|
||||
|
||||
| Family | Operations | Exact route notes |
|
||||
|---|---|---|
|
||||
| Seedream 5.0 Pro | generate, edit, decompose | `bytedance/seedream-v5.0-pro/text-to-image`, `/edit`, `/layer-decomposition`; edit accepts up to 10 images; sizes use `WIDTH*HEIGHT` |
|
||||
| Seedream 5.0 Lite | edit | `bytedance/seedream-v5.0-lite/edit`; no live Lite text-to-image sibling is assumed; accepts up to 14 images |
|
||||
| GPT Image 2 | generate, edit | `openai/gpt-image-2/text-to-image`, `/edit`; sizes use `WIDTHxHEIGHT`; edit accepts up to 10 images |
|
||||
| Nano Banana 2 | generate, edit | `google/nano-banana-2/text-to-image`, `/edit`; uses aspect ratio plus 1k/2k/4k resolution; edit accepts up to 14 images |
|
||||
|
||||
Set `generation_mode` to `generate`, `edit`, or `decompose`. Source images can
|
||||
be supplied as `image_path`, `image_paths`, `image_url`, or `image_urls`.
|
||||
|
||||
```python
|
||||
tool.execute({
|
||||
"prompt": "Replace the packaging with matte cobalt glass; preserve the logo",
|
||||
"model": "google/nano-banana-2/edit",
|
||||
"generation_mode": "edit",
|
||||
"image_paths": ["projects/campaign/source.png"],
|
||||
"resolution": "2k",
|
||||
"output_path": "projects/campaign/revised.png",
|
||||
})
|
||||
```
|
||||
|
||||
## Result and failure contract
|
||||
|
||||
Both tools submit to Atlas, poll the returned prediction id, download every
|
||||
output, and return `ToolResult` provenance containing the provider, exact model,
|
||||
prediction id, submitted parameters, source URL, artifact paths, and estimated
|
||||
cost. A missing key, unsupported route, invalid enum, missing media input, failed
|
||||
prediction, timeout, or download failure must return a failed result without
|
||||
switching providers.
|
||||
|
||||
Confirm current pricing and schemas from the machine-readable model page
|
||||
(`Accept: text/markdown`) immediately before quoting or spending on a batch.
|
||||
@@ -145,6 +145,7 @@ Three selector tools abstract multi-provider capabilities:
|
||||
| `tts_selector` | Text-to-speech | Ranks discovered providers by task fit, quality, control, reliability, cost, latency, and continuity |
|
||||
| `image_selector` | Image generation | Ranks discovered providers from the live registry; no hardcoded provider order |
|
||||
| `video_selector` | Video generation | Ranks discovered providers from the live registry; user preference is respected when explicitly provided |
|
||||
| `atlas_image` / `atlas_video` | Atlas Cloud generation | Exposes exact per-model route catalogs for image generation/editing and text/image/reference/video-edit generation |
|
||||
|
||||
Selectors route based on: user preference when explicitly set, then scored ranking across available providers. They adapt input schemas between providers transparently.
|
||||
|
||||
|
||||
@@ -54,6 +54,7 @@ AZURE_SPEECH_REGION= # Speech resource region, e.g. eastus
|
||||
# MULTI-MODEL GATEWAY (one key, 6+ tools)
|
||||
FAL_KEY= # FLUX, Recraft, Kling, Veo, MiniMax video
|
||||
MINIMAX_API_KEY= # MiniMax first-party image + MiniMax H3 video generation
|
||||
ATLASCLOUD_API_KEY= # Atlas Cloud image/video gateway
|
||||
|
||||
# KLING OFFICIAL DIRECT API
|
||||
KLING_API_KEY= # Official Kling video, image, TTS, avatar, lip sync
|
||||
@@ -403,6 +404,32 @@ choose them with `preferred_provider: "minimax"`.
|
||||
|
||||
---
|
||||
|
||||
### Atlas Cloud — Image and Video Gateway
|
||||
|
||||
**Tools:** `atlas_image`, `atlas_video`
|
||||
**Env var:** `ATLASCLOUD_API_KEY` (aliases: `ATLAS_CLOUD_API_KEY`, `ATLAS_API_KEY`)
|
||||
**Skill:** `.agents/skills/atlas-cloud/SKILL.md`
|
||||
|
||||
Atlas Cloud provides one endpoint and key for the following explicitly cataloged
|
||||
routes. OpenMontage validates each model's real schema instead of treating task
|
||||
suffixes or parameter names as interchangeable.
|
||||
|
||||
| Family | Supported routes | Current Atlas rate |
|
||||
|---|---|---:|
|
||||
| Seedance 2.5 | text/image/reference to video | $0.134/sec |
|
||||
| Seedance 2.0 | text/image/reference to video | $0.112/sec |
|
||||
| Gemini Omni Flash | text/image/reference to video; video edit; developer text/image/reference | $0.112–0.140/sec |
|
||||
| MiniMax H3 | text/image/reference to video | $0.100/sec |
|
||||
| Seedream 5.0 Pro | text to image; edit; layer decomposition | $0.022–0.045/image |
|
||||
| GPT Image 2 | text to image; edit | $0.009–0.010/image |
|
||||
| Nano Banana 2 | text to image; edit | $0.080/image |
|
||||
|
||||
Inspect `get_info()["model_catalog"]` for exact IDs, operations, media shapes,
|
||||
durations, and resolutions. Prices are estimates sourced from each model's
|
||||
machine-readable Atlas page and should be reconfirmed before a paid batch.
|
||||
|
||||
---
|
||||
|
||||
### Kling Official — Direct API
|
||||
|
||||
> **Official Kling path.** This is separate from `kling_video` via fal.ai: it uses Kling's official `Authorization: Bearer <KLING_API_KEY>` API, provider name `kling_official`, and direct Classic/Turbo/Omni task protocols.
|
||||
@@ -1330,6 +1357,7 @@ These tools require only FFmpeg or Python packages — no GPU, no API key.
|
||||
| **ElevenLabs** | `ELEVENLABS_API_KEY` | `elevenlabs_tts`, `music_gen` | Free tier + paid |
|
||||
| **fish.audio** | `FISH_AUDIO_API_KEY` | `fish_audio_tts` | Free tier (s2.1-pro-free) + paid |
|
||||
| **fal.ai** | `FAL_KEY` | `flux_image`, `recraft_image`, `kling_video`, `veo_video`, `seedance_video`, `gemini_omni_fal`, `minimax_fal_video` | Pay-as-you-go |
|
||||
| **Atlas Cloud** | `ATLASCLOUD_API_KEY` | `atlas_image`, `atlas_video` | Pay-as-you-go |
|
||||
| **Kling Official** | `KLING_API_KEY` | `kling_official_video`, `kling_official_image`, `kling_tts`, `kling_avatar`, `kling_lip_sync` | Pay-as-you-go |
|
||||
| **Volcengine Ark** | `ARK_API_KEY` | `seedance_ark` | Pay-as-you-go |
|
||||
| **MiniMax direct** | `MINIMAX_API_KEY` | `minimax_image`, `minimax_video` | Pay-as-you-go |
|
||||
|
||||
252
scripts/atlas_media_smoke.py
Normal file
252
scripts/atlas_media_smoke.py
Normal file
@@ -0,0 +1,252 @@
|
||||
"""Run the opt-in paid Atlas Cloud media smoke suite through OpenMontage selectors."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
from tools.graphics.image_selector import ImageSelector
|
||||
from tools.video.video_selector import VideoSelector
|
||||
|
||||
|
||||
IMAGE_CASES: tuple[dict[str, Any], ...] = (
|
||||
{
|
||||
"id": "seedream_campaign_key_visual",
|
||||
"skill": "campaign-key-visual",
|
||||
"model": "bytedance/seedream-v5.0-pro/text-to-image",
|
||||
"prompt": (
|
||||
"Create a premium vertical 3:4 campaign key visual for the fictional fragrance LUMA VALE. "
|
||||
"One faceted cobalt-glass bottle stands on pale travertine, surrounded by three translucent "
|
||||
"mineral ribbons and a single hard-edged pool of amber light. Reserve a calm upper-left field "
|
||||
"for exactly the headline LUMA VALE and the smaller line MINERAL LIGHT. Editorial art direction, "
|
||||
"precise readable typography, tactile glass, stone pores, controlled shadows, deep cobalt, warm "
|
||||
"amber and bone palette, 85mm product lens. No extra copy, duplicate bottle, fake certification, "
|
||||
"watermark, illegible letters, clutter, or generic luxury-ad styling. Fictional brand concept."
|
||||
),
|
||||
"width": 1530,
|
||||
"height": 2040,
|
||||
"output_format": "png",
|
||||
},
|
||||
{
|
||||
"id": "gpt_image_impossible_material_portrait",
|
||||
"skill": "impossible-material-fashion-portrait",
|
||||
"model": "openai/gpt-image-2/text-to-image",
|
||||
"prompt": (
|
||||
"Create a 3:4 editorial fashion portrait of one fictional adult model from the waist up, wearing "
|
||||
"a sculptural coat made from translucent smoked quartz that bends like tailored wool while retaining "
|
||||
"crystalline fracture planes. The face remains natural and unobstructed; the coat has one high collar, "
|
||||
"two sleeves, plausible seams and gravity. Charcoal cyclorama, narrow cool rim light, soft frontal fill, "
|
||||
"medium-format 100mm look, realistic skin texture, restrained slate and silver palette. No real-person "
|
||||
"likeness, extra limbs, fused hands, jewelry, text, logo, watermark, plastic CGI sheen, or broken anatomy."
|
||||
),
|
||||
"width": 1536,
|
||||
"height": 2048,
|
||||
"quality": "high",
|
||||
"output_format": "png",
|
||||
},
|
||||
{
|
||||
"id": "nano_banana_bottled_world",
|
||||
"skill": "bottled-miniature-world",
|
||||
"model": "google/nano-banana-2/text-to-image",
|
||||
"prompt": (
|
||||
"Create a vertical 3:4 cinematic illustration of one complete clear apothecary bottle on an old walnut "
|
||||
"desk, sealed with one rough cork. Inside is one coherent fictional 1:500 moonlit canal village: exactly "
|
||||
"one stone observatory, one arched bridge, three cottage clusters, one connected path, dark pines and "
|
||||
"warm lanterns, with no people or animals. Prove containment with a visible base, rounded shoulders, neck, "
|
||||
"thick rim, wall thickness, curved Fresnel highlights, edge refraction, caustics and tabletop contact. "
|
||||
"100mm macro, cool moonlight and warm windows. No duplicated landmark, floating architecture, scale drift, "
|
||||
"broken glass, cloudy walls, impossible refraction, label, logo, text, or watermark."
|
||||
),
|
||||
"width": 1536,
|
||||
"height": 2048,
|
||||
"resolution": "2k",
|
||||
"thinking_level": "high",
|
||||
"output_format": "png",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
VIDEO_CASES: tuple[dict[str, Any], ...] = (
|
||||
{
|
||||
"id": "seedance25_architectural_reveal",
|
||||
"skill": "architectural-sketch-to-space-reveal",
|
||||
"model": "bytedance/seedance-2.5/text-to-video",
|
||||
"prompt": (
|
||||
"One continuous ten-second architectural transformation. Begin inches above an architect's graphite "
|
||||
"section drawing on cream paper; the camera glides forward as drawn contour lines rise into warm limestone "
|
||||
"walls, pencil hatching becomes slatted oak, and a blue wash becomes a shallow reflecting pool. Without a "
|
||||
"cut, pass through the drawn doorway into the completed sunlit courtyard while construction lines remain "
|
||||
"faintly visible in the finished surfaces. End in a stable wide reveal beneath a circular oculus. Precise "
|
||||
"geometry, believable material transition, restrained museum atmosphere, synchronized pencil-scratch and "
|
||||
"room-tone audio. No people, captions, logos, teleporting objects, melting walls, jump cuts, or watermark."
|
||||
),
|
||||
"operation": "text_to_video",
|
||||
"duration": 10,
|
||||
"resolution": "720p",
|
||||
"aspect_ratio": "16:9",
|
||||
"generate_audio": True,
|
||||
},
|
||||
{
|
||||
"id": "seedance20_food_fantasy",
|
||||
"skill": "food-sensory-fantasy-film",
|
||||
"model": "bytedance/seedance-2.0/text-to-video",
|
||||
"prompt": (
|
||||
"A ten-second macro food fantasy in one fluid shot: an immaculate dark-chocolate sphere rests on black "
|
||||
"stone; a ribbon of hot espresso pours from above, the shell cracks cleanly and opens like petals, releasing "
|
||||
"a miniature saffron cloud and dozens of bright ruby pomegranate seeds that bounce with believable weight. "
|
||||
"The camera makes a slow 120-degree orbit and settles on the glossy molten center. High-speed detail, real "
|
||||
"fluid viscosity, crisp crumbs, appetizing steam, warm copper highlights, synchronized pour/crack/bounce audio. "
|
||||
"No hands, utensils, text, brand, duplicate fruit, dirty surface, implausible splash, jump cut, or watermark."
|
||||
),
|
||||
"operation": "text_to_video",
|
||||
"duration": 10,
|
||||
"resolution": "720p",
|
||||
"aspect_ratio": "16:9",
|
||||
"generate_audio": True,
|
||||
},
|
||||
{
|
||||
"id": "gemini_omni_greenhouse_take",
|
||||
"skill": "generated-continuous-take",
|
||||
"model": "google/gemini-omni-flash/text-to-video",
|
||||
"prompt": (
|
||||
"One unbroken ten-second Steadicam journey through a vast night greenhouse during a gentle storm. Start "
|
||||
"tight on rain sliding down one glass pane, pull backward through hanging vines, descend beside a narrow "
|
||||
"irrigation channel, then curve around a gardener's empty brass cart as hundreds of bioluminescent blue "
|
||||
"flowers open in a timed wave toward the lens. Finish at a wide symmetrical view of the glowing conservatory. "
|
||||
"Continuous geography and lighting, foreground occlusion motivates the move, realistic wet leaves and glass, "
|
||||
"subtle rain and metal resonance. No people, hidden cuts, camera collision, duplicated cart, text, or watermark."
|
||||
),
|
||||
"operation": "text_to_video",
|
||||
"duration": 10,
|
||||
"resolution": "720p",
|
||||
"aspect_ratio": "16:9",
|
||||
"thinking_level": "high",
|
||||
},
|
||||
{
|
||||
"id": "minimax_h3_creature_encounter",
|
||||
"skill": "creature-encounter-film",
|
||||
"model": "minimax/h3/text-to-video",
|
||||
"prompt": (
|
||||
"A ten-second cinematic wildlife encounter on a wind-scoured volcanic plateau at dawn. A small six-legged "
|
||||
"fictional basalt creature sprints beside the low tracking camera; each footfall kicks loose pumice with "
|
||||
"convincing mass. At second four it leaps across a narrow fissure, folds its stone plates midair, lands hard, "
|
||||
"slides, and braces while a sheet of dust overtakes the lens. The camera eases to a stop as the creature turns "
|
||||
"one luminous amber eye toward us. One continuous shot, consistent anatomy and scale, strong contact shadows, "
|
||||
"real inertia and debris physics. No attack, gore, extra creature, morphing limbs, text, logo, cut, or watermark."
|
||||
),
|
||||
"operation": "text_to_video",
|
||||
"duration": 10,
|
||||
"resolution": "2K",
|
||||
"aspect_ratio": "16:9",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _result_record(case: dict[str, Any], result: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"id": case["id"],
|
||||
"skill": case["skill"],
|
||||
"model": case["model"],
|
||||
"success": result.success,
|
||||
"error": result.error,
|
||||
"artifacts": result.artifacts,
|
||||
"cost_usd": result.cost_usd,
|
||||
"data": result.data,
|
||||
}
|
||||
|
||||
|
||||
def run(project_dir: Path, *, images: bool, videos: bool, allow_paid: bool) -> int:
|
||||
if not allow_paid:
|
||||
raise SystemExit("Refusing paid Atlas calls without --allow-paid")
|
||||
|
||||
project_dir.mkdir(parents=True, exist_ok=True)
|
||||
manifest_path = project_dir / "smoke_manifest.json"
|
||||
records: list[dict[str, Any]] = []
|
||||
if manifest_path.exists():
|
||||
previous = json.loads(manifest_path.read_text(encoding="utf-8"))
|
||||
records = [record for record in previous.get("records", []) if record.get("success")]
|
||||
|
||||
def checkpoint() -> None:
|
||||
manifest = {
|
||||
"provider": "atlascloud",
|
||||
"endpoint_policy": "OpenMontage selectors -> Atlas Cloud tools -> Atlas Cloud HTTP API",
|
||||
"estimated_batch_cost_usd": 4.844,
|
||||
"records": records,
|
||||
"success": len({record["id"] for record in records if record["success"]}) == len(IMAGE_CASES) + len(VIDEO_CASES),
|
||||
}
|
||||
manifest_path.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
||||
|
||||
completed_ids = {record["id"] for record in records}
|
||||
if images:
|
||||
selector = ImageSelector()
|
||||
for case in IMAGE_CASES:
|
||||
if case["id"] in completed_ids:
|
||||
continue
|
||||
inputs = {
|
||||
**case,
|
||||
"preferred_provider": "atlascloud",
|
||||
"allowed_providers": ["atlascloud"],
|
||||
"generation_mode": "generate",
|
||||
"output_path": str(project_dir / "images" / f"{case['id']}.png"),
|
||||
}
|
||||
inputs.pop("id")
|
||||
inputs.pop("skill")
|
||||
result = selector.execute(inputs)
|
||||
records.append(_result_record(case, result))
|
||||
checkpoint()
|
||||
if not result.success:
|
||||
break
|
||||
|
||||
if videos and all(record["success"] for record in records):
|
||||
selector = VideoSelector()
|
||||
for case in VIDEO_CASES:
|
||||
if case["id"] in completed_ids:
|
||||
continue
|
||||
inputs = {
|
||||
**case,
|
||||
"preferred_provider": "atlascloud",
|
||||
"allowed_providers": ["atlascloud"],
|
||||
"output_path": str(project_dir / "videos" / f"{case['id']}.mp4"),
|
||||
}
|
||||
inputs.pop("id")
|
||||
inputs.pop("skill")
|
||||
result = selector.execute(inputs)
|
||||
records.append(_result_record(case, result))
|
||||
checkpoint()
|
||||
if not result.success:
|
||||
break
|
||||
|
||||
checkpoint()
|
||||
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
|
||||
print(json.dumps({
|
||||
"success": manifest["success"],
|
||||
"manifest": str(manifest_path),
|
||||
"results": [{"id": item["id"], "success": item["success"], "error": item["error"]} for item in records],
|
||||
}, indent=2))
|
||||
return 0 if manifest["success"] else 1
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--project-dir", type=Path, required=True)
|
||||
parser.add_argument("--images-only", action="store_true")
|
||||
parser.add_argument("--videos-only", action="store_true")
|
||||
parser.add_argument("--allow-paid", action="store_true")
|
||||
args = parser.parse_args()
|
||||
return run(
|
||||
args.project_dir,
|
||||
images=not args.videos_only,
|
||||
videos=not args.images_only,
|
||||
allow_paid=args.allow_paid,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
131
tests/conftest.py
Normal file
131
tests/conftest.py
Normal file
@@ -0,0 +1,131 @@
|
||||
"""Session-wide test safety net.
|
||||
|
||||
**No test may open a network connection.** Provider tools bill per call, so a
|
||||
test that reaches a real endpoint costs the developer money — silently, and
|
||||
every time CI runs. This blocks outbound sockets for the whole test session.
|
||||
|
||||
The guard is at the socket layer on purpose. Patching `requests` only covers
|
||||
tools that use `requests`; the fleet also talks to vendor SDKs (google-cloud,
|
||||
openai, boto3), `httpx`, and raw `urllib`. Everything bottoms out in
|
||||
`socket.connect`, so that is where the wall goes.
|
||||
|
||||
Loopback is still allowed — local servers, ffmpeg RPC, and Backlot fixtures need it.
|
||||
|
||||
To write a test that genuinely hits a live API:
|
||||
|
||||
@pytest.mark.live_api
|
||||
def test_real_call():
|
||||
...
|
||||
|
||||
Marked tests are **skipped by default** and only run with the env flag set:
|
||||
|
||||
OPENMONTAGE_ALLOW_NETWORK=1 pytest -m live_api
|
||||
|
||||
Limitation: this guards the pytest process. A test that shells out to a
|
||||
subprocess (node, ffmpeg, npx) is outside its reach — don't call paid APIs
|
||||
from a subprocess in tests.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import socket
|
||||
|
||||
import pytest
|
||||
|
||||
_ALLOW_ENV_FLAG = "OPENMONTAGE_ALLOW_NETWORK"
|
||||
|
||||
_LOOPBACK_HOSTS = {"localhost", "127.0.0.1", "::1", "0.0.0.0", ""}
|
||||
|
||||
_real_connect = socket.socket.connect
|
||||
_real_connect_ex = socket.socket.connect_ex
|
||||
_real_create_connection = socket.create_connection
|
||||
|
||||
|
||||
class NetworkCallInTestError(RuntimeError):
|
||||
"""Raised when a test tries to open a non-loopback connection."""
|
||||
|
||||
|
||||
def _network_allowed() -> bool:
|
||||
return os.environ.get(_ALLOW_ENV_FLAG, "").strip().lower() in {"1", "true", "yes"}
|
||||
|
||||
|
||||
def _is_loopback(address) -> bool:
|
||||
"""True for loopback TCP/UDP targets and for AF_UNIX socket paths."""
|
||||
if isinstance(address, (str, bytes)):
|
||||
return True # AF_UNIX / abstract socket — local by definition
|
||||
if not isinstance(address, (tuple, list)) or not address:
|
||||
return True # unrecognised shape; let the real call decide
|
||||
host = address[0]
|
||||
if isinstance(host, bytes):
|
||||
host = host.decode("utf-8", "replace")
|
||||
if not isinstance(host, str):
|
||||
return False
|
||||
host = host.strip("[]").lower()
|
||||
if host in _LOOPBACK_HOSTS:
|
||||
return True
|
||||
return host.startswith("127.")
|
||||
|
||||
|
||||
def _blocked(address) -> NetworkCallInTestError:
|
||||
return NetworkCallInTestError(
|
||||
f"Blocked a network connection to {address!r} during a test.\n"
|
||||
f"\n"
|
||||
f"Tests must not call real endpoints — provider APIs bill per request.\n"
|
||||
f"Mock the transport instead (see tests/tools/test_atlas_video.py for the\n"
|
||||
f"fake-`requests` pattern), or mark the test @pytest.mark.live_api and run\n"
|
||||
f"it deliberately with {_ALLOW_ENV_FLAG}=1."
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
def _block_network():
|
||||
"""Refuse non-loopback sockets for the entire session."""
|
||||
if _network_allowed():
|
||||
yield
|
||||
return
|
||||
|
||||
def guarded_connect(self, address, *args, **kwargs):
|
||||
if not _is_loopback(address):
|
||||
raise _blocked(address)
|
||||
return _real_connect(self, address, *args, **kwargs)
|
||||
|
||||
def guarded_connect_ex(self, address, *args, **kwargs):
|
||||
if not _is_loopback(address):
|
||||
raise _blocked(address)
|
||||
return _real_connect_ex(self, address, *args, **kwargs)
|
||||
|
||||
def guarded_create_connection(address, *args, **kwargs):
|
||||
if not _is_loopback(address):
|
||||
raise _blocked(address)
|
||||
return _real_create_connection(address, *args, **kwargs)
|
||||
|
||||
socket.socket.connect = guarded_connect
|
||||
socket.socket.connect_ex = guarded_connect_ex
|
||||
socket.create_connection = guarded_create_connection
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
socket.socket.connect = _real_connect
|
||||
socket.socket.connect_ex = _real_connect_ex
|
||||
socket.create_connection = _real_create_connection
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"live_api: test performs a real, billable API call. Skipped unless "
|
||||
f"{_ALLOW_ENV_FLAG}=1 is set.",
|
||||
)
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
"""Skip live_api tests unless the operator explicitly opted in."""
|
||||
if _network_allowed():
|
||||
return
|
||||
skip = pytest.mark.skip(
|
||||
reason=f"live API test — costs money; set {_ALLOW_ENV_FLAG}=1 to run"
|
||||
)
|
||||
for item in items:
|
||||
if "live_api" in item.keywords:
|
||||
item.add_marker(skip)
|
||||
250
tests/contracts/test_atlas_tools.py
Normal file
250
tests/contracts/test_atlas_tools.py
Normal file
@@ -0,0 +1,250 @@
|
||||
"""Contract and exact-schema tests for the Atlas Cloud media gateway."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from tools import atlas_client
|
||||
from tools.atlas_models import IMAGE_MODELS, VIDEO_MODELS
|
||||
from tools.base_tool import BaseTool, ExecutionMode, ToolRuntime, ToolStability, ToolStatus, ToolTier
|
||||
from tools.graphics.atlas_image import AtlasImage
|
||||
from tools.graphics.image_selector import ImageSelector
|
||||
from tools.video.atlas_video import AtlasVideo
|
||||
from tools.video.video_selector import VideoSelector
|
||||
|
||||
TOOLS = [AtlasImage, AtlasVideo]
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_atlas_env(monkeypatch):
|
||||
for key in atlas_client.ENV_KEYS:
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("cls", TOOLS, ids=lambda cls: cls.name)
|
||||
class TestContract:
|
||||
def test_contract_identity(self, cls):
|
||||
tool = cls()
|
||||
assert issubclass(cls, BaseTool)
|
||||
assert tool.provider == "atlascloud"
|
||||
assert tool.tier == ToolTier.GENERATE
|
||||
assert tool.stability == ToolStability.BETA
|
||||
assert tool.runtime == ToolRuntime.API
|
||||
assert tool.execution_mode == ExecutionMode.SYNC
|
||||
assert "prompt" in tool.input_schema["required"]
|
||||
assert "atlas-cloud" in tool.agent_skills
|
||||
assert "env:ATLASCLOUD_API_KEY" in tool.dependencies
|
||||
|
||||
def test_status_and_key_aliases(self, cls, monkeypatch):
|
||||
assert cls().get_status() == ToolStatus.UNAVAILABLE
|
||||
for key in atlas_client.ENV_KEYS:
|
||||
monkeypatch.setenv(key, "test")
|
||||
assert cls().get_status() == ToolStatus.AVAILABLE
|
||||
monkeypatch.delenv(key)
|
||||
|
||||
def test_no_key_fails_before_network(self, cls):
|
||||
result = cls().execute({"prompt": "test"})
|
||||
assert not result.success
|
||||
assert "ATLASCLOUD_API_KEY" in result.error
|
||||
|
||||
def test_discovery_metadata_contains_catalog(self, cls):
|
||||
info = cls().get_info()
|
||||
assert info["model_catalog"]
|
||||
assert info["provider_matrix"]
|
||||
|
||||
|
||||
def test_layer3_skill_exists():
|
||||
skill = PROJECT_ROOT / ".agents" / "skills" / "atlas-cloud" / "SKILL.md"
|
||||
assert skill.exists()
|
||||
assert "ATLASCLOUD_API_KEY" in skill.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def test_registry_discovers_atlas_tools(isolated_tool_registry):
|
||||
isolated_tool_registry.discover()
|
||||
assert isolated_tool_registry.get("atlas_image") is not None
|
||||
assert isolated_tool_registry.get("atlas_video") is not None
|
||||
|
||||
|
||||
def test_exact_atlas_model_pins_video_selector_candidate():
|
||||
atlas = AtlasVideo()
|
||||
unrelated = type("Unrelated", (), {
|
||||
"input_schema": {"properties": {"model": {"enum": ["vendor/other"]}}},
|
||||
"get_info": lambda self: {},
|
||||
"supports": {"text_to_video": True},
|
||||
"is_operation_available": lambda self, operation: True,
|
||||
})()
|
||||
filtered = VideoSelector()._filter_candidates(
|
||||
{"model": "minimax/h3/text-to-video", "operation": "text_to_video"},
|
||||
[unrelated, atlas],
|
||||
)
|
||||
assert filtered == [atlas]
|
||||
|
||||
|
||||
def test_exact_atlas_model_pins_image_selector_candidate():
|
||||
atlas = AtlasImage()
|
||||
unrelated = type("Unrelated", (), {
|
||||
"input_schema": {"properties": {"model": {"enum": ["vendor/other"]}}},
|
||||
"get_info": lambda self: {},
|
||||
"supports": {},
|
||||
})()
|
||||
filtered = ImageSelector()._filter_candidates(
|
||||
{"model": "openai/gpt-image-2/text-to-image"},
|
||||
[unrelated, atlas],
|
||||
)
|
||||
assert filtered == [atlas]
|
||||
|
||||
|
||||
class TestVideoRoutes:
|
||||
@pytest.mark.parametrize(
|
||||
"family,operations",
|
||||
[
|
||||
("bytedance/seedance-2.5", {"text_to_video", "image_to_video", "reference_to_video"}),
|
||||
("bytedance/seedance-2.0", {"text_to_video", "image_to_video", "reference_to_video"}),
|
||||
("minimax/h3", {"text_to_video", "image_to_video", "reference_to_video"}),
|
||||
("google/gemini-omni-flash", {"text_to_video", "image_to_video", "reference_to_video", "video_edit"}),
|
||||
],
|
||||
)
|
||||
def test_family_operations_are_discoverable(self, family, operations):
|
||||
assert operations <= set(AtlasVideo.provider_matrix[family])
|
||||
|
||||
def test_exact_live_catalog(self):
|
||||
assert len(VIDEO_MODELS) == 16
|
||||
assert "minimax/h3/reference-to-video" in VIDEO_MODELS
|
||||
assert "google/gemini-omni-flash/video-edit" in VIDEO_MODELS
|
||||
|
||||
def test_seedance_25_payload(self):
|
||||
payload = AtlasVideo()._build_payload(
|
||||
{"prompt": "p", "duration": 10, "aspect_ratio": "21:9", "resolution": "720p", "generate_audio": True},
|
||||
"bytedance/seedance-2.5/text-to-video",
|
||||
)
|
||||
assert payload == {
|
||||
"model": "bytedance/seedance-2.5/text-to-video", "prompt": "p", "duration": 10,
|
||||
"ratio": "21:9", "resolution": "720p", "generate_audio": True,
|
||||
}
|
||||
|
||||
def test_seedance_i2v_uses_image_and_last_image(self):
|
||||
payload = AtlasVideo()._build_payload(
|
||||
{"prompt": "p", "duration": 10, "resolution": "720p", "image_url": "https://x/a.png", "last_image_url": "https://x/b.png"},
|
||||
"bytedance/seedance-2.5/image-to-video",
|
||||
)
|
||||
assert payload["image"] == "https://x/a.png"
|
||||
assert payload["last_image"] == "https://x/b.png"
|
||||
assert "image_url" not in payload
|
||||
|
||||
def test_seedance_25_accepts_audio_only_and_30_images(self):
|
||||
tool = AtlasVideo()
|
||||
audio = tool._build_payload(
|
||||
{"prompt": "cut to the beat", "duration": 10, "reference_audios": ["https://x/beat.mp3"]},
|
||||
"bytedance/seedance-2.5/reference-to-video",
|
||||
)
|
||||
assert audio["reference_audios"] == ["https://x/beat.mp3"]
|
||||
payload = tool._build_payload(
|
||||
{"prompt": "p", "duration": 10, "reference_images": [f"https://x/{i}.png" for i in range(30)]},
|
||||
"bytedance/seedance-2.5/reference-to-video",
|
||||
)
|
||||
assert len(payload["reference_images"]) == 30
|
||||
|
||||
def test_gemini_standard_and_developer_have_distinct_schemas(self):
|
||||
standard = AtlasVideo()._build_payload(
|
||||
{"prompt": "p", "duration": 10, "reference_images": ["https://x/a.png"]},
|
||||
"google/gemini-omni-flash/reference-to-video",
|
||||
)
|
||||
developer = AtlasVideo()._build_payload(
|
||||
{"prompt": "p", "duration": 10, "video_url": "https://x/a.mp4"},
|
||||
"google/gemini-omni-flash/reference-to-video-developer",
|
||||
)
|
||||
assert standard["images"] == ["https://x/a.png"]
|
||||
assert developer["video_clips"] == [{"url": "https://x/a.mp4", "start": 0, "ends": 10}]
|
||||
|
||||
def test_gemini_video_edit_uses_video(self):
|
||||
payload = AtlasVideo()._build_payload(
|
||||
{"prompt": "make it dusk", "video_url": "https://x/a.mp4"},
|
||||
"google/gemini-omni-flash/video-edit",
|
||||
)
|
||||
assert payload["video"] == "https://x/a.mp4"
|
||||
assert "duration" not in payload and "aspect_ratio" not in payload
|
||||
|
||||
def test_h3_refers_objects(self):
|
||||
payload = AtlasVideo()._build_payload(
|
||||
{"prompt": "p", "duration": 10, "resolution": "2K", "reference_images": ["https://x/a.png"], "reference_audios": ["https://x/a.mp3"]},
|
||||
"minimax/h3/reference-to-video",
|
||||
)
|
||||
assert payload["refers"] == [
|
||||
{"url": "https://x/a.png", "type": "image"},
|
||||
{"url": "https://x/a.mp3", "type": "audio"},
|
||||
]
|
||||
|
||||
def test_h3_i2v_uses_end_image(self):
|
||||
payload = AtlasVideo()._build_payload(
|
||||
{"prompt": "p", "duration": 10, "resolution": "2K", "image_url": "https://x/a.png", "end_image_url": "https://x/b.png"},
|
||||
"minimax/h3/image-to-video",
|
||||
)
|
||||
assert payload["image"] == "https://x/a.png"
|
||||
assert payload["end_image"] == "https://x/b.png"
|
||||
|
||||
def test_invalid_route_and_enum_fail_loudly(self):
|
||||
tool = AtlasVideo()
|
||||
with pytest.raises(ValueError, match="does not expose"):
|
||||
tool._resolve_model("minimax/h3/text-to-video", "video_edit")
|
||||
with pytest.raises(ValueError, match="duration"):
|
||||
tool._build_payload({"prompt": "p", "duration": 30}, "minimax/h3/text-to-video")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected",
|
||||
[
|
||||
("bytedance/seedance-2.5/text-to-video", 1.34),
|
||||
("bytedance/seedance-2.0/text-to-video", 1.12),
|
||||
("google/gemini-omni-flash/text-to-video", 1.25),
|
||||
("google/gemini-omni-flash/text-to-video-developer", 1.12),
|
||||
("minimax/h3/text-to-video", 1.00),
|
||||
],
|
||||
)
|
||||
def test_verified_costs(self, model, expected):
|
||||
assert AtlasVideo().estimate_cost({"model": model, "duration": 10}) == pytest.approx(expected)
|
||||
|
||||
|
||||
class TestImageRoutes:
|
||||
def test_exact_live_catalog(self):
|
||||
assert set(IMAGE_MODELS) == {
|
||||
"bytedance/seedream-v5.0-pro/text-to-image", "bytedance/seedream-v5.0-pro/edit",
|
||||
"bytedance/seedream-v5.0-pro/layer-decomposition", "bytedance/seedream-v5.0-lite/edit",
|
||||
"openai/gpt-image-2/text-to-image", "openai/gpt-image-2/edit",
|
||||
"google/nano-banana-2/text-to-image", "google/nano-banana-2/edit",
|
||||
}
|
||||
|
||||
def test_seedream_uses_star_size_and_images(self):
|
||||
tool = AtlasImage()
|
||||
generated = tool._build_payload({"prompt": "p", "width": 2048, "height": 1152}, "bytedance/seedream-v5.0-pro/text-to-image")
|
||||
edited = tool._build_payload({"prompt": "p", "width": 2048, "height": 1152, "image_urls": ["https://x/a.png"]}, "bytedance/seedream-v5.0-pro/edit")
|
||||
assert generated["size"] == "2048*1152"
|
||||
assert edited["images"] == ["https://x/a.png"]
|
||||
|
||||
def test_gpt_image_uses_x_size(self):
|
||||
payload = AtlasImage()._build_payload({"prompt": "p", "width": 1536, "height": 1024}, "openai/gpt-image-2/text-to-image")
|
||||
assert payload["size"] == "1536x1024"
|
||||
assert "aspect_ratio" not in payload
|
||||
|
||||
def test_nano_banana_uses_ratio_and_resolution(self):
|
||||
payload = AtlasImage()._build_payload(
|
||||
{"prompt": "p", "width": 1920, "height": 1080, "resolution": "2k", "thinking_level": "high"},
|
||||
"google/nano-banana-2/text-to-image",
|
||||
)
|
||||
assert payload["aspect_ratio"] == "16:9"
|
||||
assert payload["resolution"] == "2k"
|
||||
assert payload["thinking_level"] == "high"
|
||||
|
||||
def test_decomposition_requires_one_image(self):
|
||||
with pytest.raises(ValueError, match="exactly one"):
|
||||
AtlasImage()._build_payload({"prompt": ""}, "bytedance/seedream-v5.0-pro/layer-decomposition")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected",
|
||||
[
|
||||
("bytedance/seedream-v5.0-pro/text-to-image", 0.045),
|
||||
("openai/gpt-image-2/text-to-image", 0.009),
|
||||
("google/nano-banana-2/text-to-image", 0.080),
|
||||
],
|
||||
)
|
||||
def test_verified_costs(self, model, expected):
|
||||
assert AtlasImage().estimate_cost({"model": model}) == pytest.approx(expected)
|
||||
76
tests/test_network_guard.py
Normal file
76
tests/test_network_guard.py
Normal file
@@ -0,0 +1,76 @@
|
||||
"""Meta-tests: prove the session network guard actually blocks paid calls.
|
||||
|
||||
If these fail, every other test in the suite is one bug away from spending money.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import socket
|
||||
|
||||
import pytest
|
||||
|
||||
from tools.graphics.atlas_image import AtlasImage
|
||||
from tools.video.atlas_video import AtlasVideo
|
||||
|
||||
|
||||
class TestGuardBlocksOutbound:
|
||||
|
||||
def test_raw_socket_connect_is_blocked(self):
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
with pytest.raises(Exception) as exc:
|
||||
sock.connect(("api.atlascloud.ai", 443))
|
||||
assert "Blocked a network connection" in str(exc.value)
|
||||
|
||||
def test_create_connection_is_blocked(self):
|
||||
with pytest.raises(Exception) as exc:
|
||||
socket.create_connection(("api.atlascloud.ai", 443), timeout=5)
|
||||
assert "Blocked a network connection" in str(exc.value)
|
||||
|
||||
def test_requests_cannot_reach_a_provider(self):
|
||||
requests = pytest.importorskip("requests")
|
||||
with pytest.raises(Exception):
|
||||
requests.get("https://api.atlascloud.ai/api/v1/model/prediction/x", timeout=5)
|
||||
|
||||
def test_loopback_still_permitted(self):
|
||||
"""Local servers and fixtures must keep working."""
|
||||
server = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
server.bind(("127.0.0.1", 0))
|
||||
server.listen(1)
|
||||
port = server.getsockname()[1]
|
||||
try:
|
||||
client = socket.create_connection(("127.0.0.1", port), timeout=5)
|
||||
client.close()
|
||||
finally:
|
||||
server.close()
|
||||
|
||||
|
||||
class TestPaidToolsCannotSpend:
|
||||
"""The guard must hold even with a real key present in the environment."""
|
||||
|
||||
def test_atlas_image_fails_instead_of_billing(self, monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("ATLASCLOUD_API_KEY", "sk-looks-real-but-must-not-be-used")
|
||||
result = AtlasImage().execute({
|
||||
"prompt": "this must never reach the API",
|
||||
"output_path": str(tmp_path / "nope.png"),
|
||||
})
|
||||
assert result.success is False
|
||||
assert result.cost_usd == 0.0
|
||||
|
||||
def test_atlas_video_fails_instead_of_billing(self, monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("ATLASCLOUD_API_KEY", "sk-looks-real-but-must-not-be-used")
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "this must never reach the API",
|
||||
"output_path": str(tmp_path / "nope.mp4"),
|
||||
})
|
||||
assert result.success is False
|
||||
assert result.cost_usd == 0.0
|
||||
|
||||
|
||||
class TestLiveApiMarkerIsSkipped:
|
||||
|
||||
@pytest.mark.live_api
|
||||
def test_this_should_never_run_by_default(self):
|
||||
raise AssertionError(
|
||||
"A @live_api test executed without OPENMONTAGE_ALLOW_NETWORK=1 — "
|
||||
"the opt-in gate is broken and real spending is possible."
|
||||
)
|
||||
339
tests/tools/test_atlas_video.py
Normal file
339
tests/tools/test_atlas_video.py
Normal file
@@ -0,0 +1,339 @@
|
||||
"""Behavioral tests for the Atlas Cloud tools with a faked `requests` module.
|
||||
|
||||
Covers the submit -> poll -> download cycle, the media-upload path for
|
||||
image-to-video, and the error paths (failed prediction, HTTP error, timeout).
|
||||
No network access and no API key required.
|
||||
|
||||
Run: pytest tests/tools/test_atlas_video.py -v
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from tools import atlas_client
|
||||
from tools.graphics.atlas_image import AtlasImage
|
||||
from tools.video.atlas_video import AtlasVideo
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, payload=None, content=b"", status_code=200, text=""):
|
||||
self._payload = payload if payload is not None else {}
|
||||
self.content = content
|
||||
self.status_code = status_code
|
||||
self.text = text or str(self._payload)
|
||||
|
||||
def json(self):
|
||||
if isinstance(self._payload, Exception):
|
||||
raise self._payload
|
||||
return self._payload
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_requests(monkeypatch):
|
||||
"""Install a fake `requests` module and record every call made through it."""
|
||||
calls = {"post": [], "get": []}
|
||||
queues = {"post": [], "get": []}
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
calls["post"].append({"url": url, **kwargs})
|
||||
if not queues["post"]:
|
||||
raise AssertionError(f"Unexpected POST to {url}")
|
||||
return queues["post"].pop(0)
|
||||
|
||||
def fake_get(url, **kwargs):
|
||||
calls["get"].append({"url": url, **kwargs})
|
||||
if not queues["get"]:
|
||||
raise AssertionError(f"Unexpected GET to {url}")
|
||||
return queues["get"].pop(0)
|
||||
|
||||
module = types.ModuleType("requests")
|
||||
module.post = fake_post
|
||||
module.get = fake_get
|
||||
monkeypatch.setitem(sys.modules, "requests", module)
|
||||
monkeypatch.setenv("ATLASCLOUD_API_KEY", "test-key")
|
||||
# Keep polling instantaneous.
|
||||
monkeypatch.setattr(atlas_client.time, "sleep", lambda _s: None)
|
||||
|
||||
return types.SimpleNamespace(calls=calls, queues=queues)
|
||||
|
||||
|
||||
SUBMITTED = FakeResponse({"code": 200, "data": {"id": "pred_abc123", "status": "processing"}})
|
||||
|
||||
|
||||
def completed(url: str) -> FakeResponse:
|
||||
return FakeResponse({"code": 200, "data": {"id": "pred_abc123", "status": "completed", "outputs": [url]}})
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Happy paths
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
class TestAtlasVideoExecute:
|
||||
|
||||
def test_text_to_video_round_trip(self, fake_requests, tmp_path):
|
||||
out = tmp_path / "clip.mp4"
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [
|
||||
FakeResponse({"code": 200, "data": {"status": "processing"}}),
|
||||
completed("https://storage.atlascloud.ai/outputs/clip.mp4"),
|
||||
FakeResponse(content=b"MP4DATA"), # download
|
||||
]
|
||||
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "a rocket launching",
|
||||
"duration": 5,
|
||||
"output_path": str(out),
|
||||
})
|
||||
|
||||
assert result.success is True, result.error
|
||||
assert out.read_bytes() == b"MP4DATA"
|
||||
assert result.data["prediction_id"] == "pred_abc123"
|
||||
assert result.data["provider"] == "atlascloud"
|
||||
assert result.cost_usd == pytest.approx(0.67) # Seedance 2.5, 5s
|
||||
|
||||
body = fake_requests.calls["post"][0]["json"]
|
||||
assert body["model"] == "bytedance/seedance-2.5/text-to-video"
|
||||
assert body["prompt"] == "a rocket launching"
|
||||
assert body["duration"] == 5
|
||||
assert body["ratio"] == "adaptive"
|
||||
assert body["resolution"] == "720p"
|
||||
|
||||
def test_succeeded_status_is_treated_as_success(self, fake_requests, tmp_path):
|
||||
"""Atlas documents both `completed` and `succeeded` as terminal success."""
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [
|
||||
FakeResponse({"code": 200, "data": {"status": "succeeded", "outputs": ["https://x/c.mp4"]}}),
|
||||
FakeResponse(content=b"MP4"),
|
||||
]
|
||||
result = AtlasVideo().execute({"prompt": "p", "output_path": str(tmp_path / "c.mp4")})
|
||||
assert result.success is True, result.error
|
||||
|
||||
def test_image_to_video_uploads_local_file(self, fake_requests, tmp_path):
|
||||
source = tmp_path / "frame.png"
|
||||
source.write_bytes(b"PNGDATA")
|
||||
|
||||
fake_requests.queues["post"] = [
|
||||
FakeResponse({"data": {"download_url": "https://storage.atlascloud.ai/uploads/frame.png"}}),
|
||||
SUBMITTED,
|
||||
]
|
||||
fake_requests.queues["get"] = [
|
||||
completed("https://storage.atlascloud.ai/outputs/clip.mp4"),
|
||||
FakeResponse(content=b"MP4"),
|
||||
]
|
||||
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "the scene comes alive",
|
||||
"operation": "image_to_video",
|
||||
"image_path": str(source),
|
||||
"output_path": str(tmp_path / "out.mp4"),
|
||||
})
|
||||
|
||||
assert result.success is True, result.error
|
||||
upload_call, generate_call = fake_requests.calls["post"]
|
||||
assert upload_call["url"].endswith("/model/uploadMedia")
|
||||
assert "files" in upload_call
|
||||
# The task suffix must have been rewritten to match the operation.
|
||||
assert generate_call["json"]["model"] == "bytedance/seedance-2.5/image-to-video"
|
||||
assert generate_call["json"]["image"] == "https://storage.atlascloud.ai/uploads/frame.png"
|
||||
|
||||
def test_upload_accepts_bare_url_response_shape(self, fake_requests, tmp_path):
|
||||
"""Older Atlas docs return {"url": ...} instead of {"data": {"download_url": ...}}."""
|
||||
source = tmp_path / "frame.png"
|
||||
source.write_bytes(b"PNG")
|
||||
|
||||
fake_requests.queues["post"] = [
|
||||
FakeResponse({"url": "https://storage.atlascloud.ai/uploads/legacy.png"}),
|
||||
SUBMITTED,
|
||||
]
|
||||
fake_requests.queues["get"] = [completed("https://x/c.mp4"), FakeResponse(content=b"MP4")]
|
||||
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "p",
|
||||
"operation": "image_to_video",
|
||||
"reference_image_path": str(source),
|
||||
"output_path": str(tmp_path / "out.mp4"),
|
||||
})
|
||||
assert result.success is True, result.error
|
||||
assert fake_requests.calls["post"][1]["json"]["image"].endswith("legacy.png")
|
||||
|
||||
|
||||
class TestReferenceToVideo:
|
||||
"""The 12-asset multimodal path: @image1 / @audio1 style references."""
|
||||
|
||||
def test_local_references_are_uploaded_and_remote_ones_pass_through(self, fake_requests, tmp_path):
|
||||
portrait = tmp_path / "cowgirl.png"
|
||||
portrait.write_bytes(b"PNG")
|
||||
track = tmp_path / "beat.mp3"
|
||||
track.write_bytes(b"MP3")
|
||||
|
||||
fake_requests.queues["post"] = [
|
||||
FakeResponse({"data": {"download_url": "https://storage.atlascloud.ai/u/cowgirl.png"}}),
|
||||
FakeResponse({"data": {"download_url": "https://storage.atlascloud.ai/u/beat.mp3"}}),
|
||||
SUBMITTED,
|
||||
]
|
||||
fake_requests.queues["get"] = [completed("https://x/c.mp4"), FakeResponse(content=b"MP4")]
|
||||
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "@image1 raps the vocal in @audio1",
|
||||
"model": "bytedance/seedance-2.0/reference-to-video",
|
||||
"operation": "reference_to_video",
|
||||
"reference_images": [str(portrait), "https://cdn.example.com/already-hosted.png"],
|
||||
"reference_audios": [str(track)],
|
||||
"duration": 15,
|
||||
"output_path": str(tmp_path / "out.mp4"),
|
||||
})
|
||||
|
||||
assert result.success is True, result.error
|
||||
body = fake_requests.calls["post"][-1]["json"]
|
||||
# Local file uploaded, hosted URL untouched, order preserved (@image1 = first).
|
||||
assert body["reference_images"] == [
|
||||
"https://storage.atlascloud.ai/u/cowgirl.png",
|
||||
"https://cdn.example.com/already-hosted.png",
|
||||
]
|
||||
assert body["reference_audios"] == ["https://storage.atlascloud.ai/u/beat.mp3"]
|
||||
assert body["model"] == "bytedance/seedance-2.0/reference-to-video"
|
||||
# Only the two local files were uploaded, not the hosted one.
|
||||
uploads = [c for c in fake_requests.calls["post"] if c["url"].endswith("/uploadMedia")]
|
||||
assert len(uploads) == 2
|
||||
|
||||
def test_reference_to_video_without_any_asset_is_rejected(self, fake_requests, tmp_path):
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "p",
|
||||
"operation": "reference_to_video",
|
||||
"output_path": str(tmp_path / "o.mp4"),
|
||||
})
|
||||
assert result.success is False
|
||||
assert "requires supported reference media" in result.error
|
||||
|
||||
def test_reference_image_alone_satisfies_reference_to_video(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [completed("https://x/c.mp4"), FakeResponse(content=b"MP4")]
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "@image1 walks",
|
||||
"operation": "reference_to_video",
|
||||
"reference_images": ["https://cdn.example.com/a.png"],
|
||||
"output_path": str(tmp_path / "o.mp4"),
|
||||
})
|
||||
assert result.success is True, result.error
|
||||
|
||||
def test_too_many_references_rejected_before_calling_api(self, fake_requests, tmp_path):
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "p",
|
||||
"model": "bytedance/seedance-2.0/reference-to-video",
|
||||
"operation": "reference_to_video",
|
||||
"reference_images": [f"https://x/{i}.png" for i in range(10)],
|
||||
"output_path": str(tmp_path / "o.mp4"),
|
||||
})
|
||||
assert result.success is False
|
||||
assert "at most 9" in result.error
|
||||
assert not fake_requests.calls["post"], "must fail before spending an API call"
|
||||
|
||||
|
||||
class TestAtlasImageExecute:
|
||||
|
||||
def test_text_to_image_round_trip(self, fake_requests, tmp_path):
|
||||
out = tmp_path / "img.png"
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [
|
||||
completed("https://storage.atlascloud.ai/outputs/img.png"),
|
||||
FakeResponse(content=b"PNGDATA"),
|
||||
]
|
||||
|
||||
result = AtlasImage().execute({
|
||||
"prompt": "a japanese garden",
|
||||
"width": 2048,
|
||||
"height": 1152,
|
||||
"output_path": str(out),
|
||||
})
|
||||
|
||||
assert result.success is True, result.error
|
||||
assert out.read_bytes() == b"PNGDATA"
|
||||
assert result.cost_usd == pytest.approx(0.045)
|
||||
|
||||
body = fake_requests.calls["post"][0]["json"]
|
||||
assert body["model"] == "bytedance/seedream-v5.0-pro/text-to-image"
|
||||
assert body["size"] == "2048*1152"
|
||||
|
||||
def test_auth_header_is_bearer(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [completed("https://x/i.png"), FakeResponse(content=b"P")]
|
||||
AtlasImage().execute({"prompt": "p", "output_path": str(tmp_path / "i.png")})
|
||||
assert fake_requests.calls["post"][0]["headers"]["Authorization"] == "Bearer test-key"
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Error paths
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
class TestErrorHandling:
|
||||
|
||||
def test_failed_prediction_surfaces_error_message(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [
|
||||
FakeResponse({"code": 200, "data": {
|
||||
"status": "failed",
|
||||
"error": "Invalid parameter: resolution not supported by this model",
|
||||
}}),
|
||||
]
|
||||
result = AtlasVideo().execute({"prompt": "p", "output_path": str(tmp_path / "o.mp4")})
|
||||
assert result.success is False
|
||||
assert "resolution not supported" in result.error
|
||||
|
||||
def test_http_error_includes_body(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [
|
||||
FakeResponse({"error": "insufficient balance"}, status_code=402, text="insufficient balance")
|
||||
]
|
||||
result = AtlasVideo().execute({"prompt": "p", "output_path": str(tmp_path / "o.mp4")})
|
||||
assert result.success is False
|
||||
assert "402" in result.error
|
||||
assert "insufficient balance" in result.error
|
||||
|
||||
def test_non_200_envelope_code_is_an_error(self, fake_requests, tmp_path):
|
||||
"""Atlas can return a failure code inside an HTTP 200 envelope."""
|
||||
fake_requests.queues["post"] = [FakeResponse({"code": 400, "message": "unknown model"})]
|
||||
result = AtlasVideo().execute({"prompt": "p", "output_path": str(tmp_path / "o.mp4")})
|
||||
assert result.success is False
|
||||
assert "unknown model" in result.error
|
||||
|
||||
def test_missing_prediction_id_is_an_error(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [FakeResponse({"code": 200, "data": {"status": "processing"}})]
|
||||
result = AtlasImage().execute({"prompt": "p", "output_path": str(tmp_path / "i.png")})
|
||||
assert result.success is False
|
||||
assert "prediction id" in result.error
|
||||
|
||||
def test_poll_timeout_reports_prediction_id(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [
|
||||
FakeResponse({"code": 200, "data": {"status": "processing"}}) for _ in range(5)
|
||||
]
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "p",
|
||||
"poll_interval": 1,
|
||||
"poll_timeout": 3,
|
||||
"output_path": str(tmp_path / "o.mp4"),
|
||||
})
|
||||
assert result.success is False
|
||||
assert "pred_abc123" in result.error
|
||||
|
||||
def test_completed_with_no_outputs_is_an_error(self, fake_requests, tmp_path):
|
||||
fake_requests.queues["post"] = [SUBMITTED]
|
||||
fake_requests.queues["get"] = [
|
||||
FakeResponse({"code": 200, "data": {"status": "completed", "outputs": []}}),
|
||||
]
|
||||
result = AtlasImage().execute({"prompt": "p", "output_path": str(tmp_path / "i.png")})
|
||||
assert result.success is False
|
||||
assert "no outputs" in result.error
|
||||
|
||||
def test_missing_upload_file_is_an_error(self, fake_requests, tmp_path):
|
||||
result = AtlasVideo().execute({
|
||||
"prompt": "p",
|
||||
"operation": "image_to_video",
|
||||
"image_path": str(tmp_path / "nope.png"),
|
||||
"output_path": str(tmp_path / "o.mp4"),
|
||||
})
|
||||
assert result.success is False
|
||||
assert "file not found" in result.error
|
||||
264
tools/atlas_client.py
Normal file
264
tools/atlas_client.py
Normal file
@@ -0,0 +1,264 @@
|
||||
"""Shared Atlas Cloud API plumbing for the image and video provider tools.
|
||||
|
||||
Atlas Cloud (https://www.atlascloud.ai) is a multi-model gateway: one key and one
|
||||
async prediction contract in front of 400+ third-party models (Seedream, FLUX,
|
||||
Nano Banana, Kling, Seedance, Hailuo, ...).
|
||||
|
||||
Every non-LLM generation follows the same three-step shape:
|
||||
|
||||
POST /api/v1/model/generateImage -> {"data": {"id": ...}}
|
||||
POST /api/v1/model/generateVideo -> {"data": {"id": ...}}
|
||||
GET /api/v1/model/prediction/{id} -> {"data": {"status", "outputs", "error"}}
|
||||
|
||||
`requests` is imported lazily inside functions so registry discovery stays fast
|
||||
(see tests/contracts — the lazy-import convention is enforced by the suite).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
BASE_URL = "https://api.atlascloud.ai/api/v1"
|
||||
|
||||
GENERATE_IMAGE_ENDPOINT = f"{BASE_URL}/model/generateImage"
|
||||
GENERATE_VIDEO_ENDPOINT = f"{BASE_URL}/model/generateVideo"
|
||||
PREDICTION_ENDPOINT = f"{BASE_URL}/model/prediction"
|
||||
UPLOAD_MEDIA_ENDPOINT = f"{BASE_URL}/model/uploadMedia"
|
||||
|
||||
# Atlas documents `created`/`processing` as in-flight and both `completed` and
|
||||
# `succeeded` as terminal success. Treat any unrecognised status as in-flight so a
|
||||
# newly introduced intermediate state can't be mistaken for a failure.
|
||||
TERMINAL_SUCCESS = {"completed", "succeeded"}
|
||||
TERMINAL_FAILURE = {"failed", "canceled", "cancelled"}
|
||||
|
||||
ENV_KEYS = ("ATLASCLOUD_API_KEY", "ATLAS_CLOUD_API_KEY", "ATLAS_API_KEY")
|
||||
|
||||
INSTALL_INSTRUCTIONS = (
|
||||
"Set ATLASCLOUD_API_KEY to your Atlas Cloud API key.\n"
|
||||
" Get one at https://www.atlascloud.ai/ -> Dashboard -> API Keys"
|
||||
)
|
||||
|
||||
|
||||
class AtlasError(RuntimeError):
|
||||
"""Raised for any Atlas Cloud API or transport failure."""
|
||||
|
||||
|
||||
def get_api_key() -> str | None:
|
||||
"""Return the first configured Atlas Cloud key.
|
||||
|
||||
ATLASCLOUD_API_KEY is the name Atlas uses in its own docs and CLI; the other
|
||||
two are accepted because they are the names people reach for by habit.
|
||||
"""
|
||||
for key in ENV_KEYS:
|
||||
value = os.environ.get(key)
|
||||
if value:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _headers(api_key: str, json_body: bool = True) -> dict[str, str]:
|
||||
headers = {"Authorization": f"Bearer {api_key}"}
|
||||
if json_body:
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
|
||||
def _payload_of(response: Any) -> dict[str, Any]:
|
||||
"""Parse an Atlas envelope, raising AtlasError with the body on any problem.
|
||||
|
||||
Atlas wraps results as {"code": 200, "data": {...}}. A non-200 `code` can ride
|
||||
along with HTTP 200, so the envelope is checked even on a successful request.
|
||||
"""
|
||||
try:
|
||||
body = response.json()
|
||||
except Exception as exc: # noqa: BLE001 - surface the raw text, not a parse trace
|
||||
text = getattr(response, "text", "")
|
||||
raise AtlasError(f"Atlas Cloud returned a non-JSON response: {text[:500]}") from exc
|
||||
|
||||
if not isinstance(body, dict):
|
||||
raise AtlasError(f"Atlas Cloud returned an unexpected payload: {str(body)[:500]}")
|
||||
|
||||
code = body.get("code")
|
||||
if code is not None and int(code) != 200:
|
||||
message = body.get("message") or body.get("error") or str(body)[:500]
|
||||
raise AtlasError(f"Atlas Cloud error (code {code}): {message}")
|
||||
|
||||
data = body.get("data")
|
||||
if data is None:
|
||||
# uploadMedia historically answered with a bare {"url": ...}.
|
||||
return body
|
||||
if not isinstance(data, dict):
|
||||
raise AtlasError(f"Atlas Cloud returned an unexpected 'data' field: {str(data)[:500]}")
|
||||
return data
|
||||
|
||||
|
||||
def _raise_for_status(response: Any, context: str) -> None:
|
||||
status = getattr(response, "status_code", 200)
|
||||
if status >= 400:
|
||||
text = getattr(response, "text", "")
|
||||
raise AtlasError(f"{context} failed with HTTP {status}: {text[:500]}")
|
||||
|
||||
|
||||
def submit(endpoint: str, payload: dict[str, Any], api_key: str, timeout: int = 60) -> str:
|
||||
"""Submit a generation request and return its prediction id."""
|
||||
import requests
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
endpoint, headers=_headers(api_key), json=payload, timeout=timeout
|
||||
)
|
||||
except AtlasError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise AtlasError(f"Could not reach Atlas Cloud at {endpoint}: {exc}") from exc
|
||||
|
||||
_raise_for_status(response, "Atlas Cloud submission")
|
||||
data = _payload_of(response)
|
||||
|
||||
prediction_id = data.get("id")
|
||||
if not prediction_id:
|
||||
raise AtlasError(f"Atlas Cloud did not return a prediction id: {str(data)[:500]}")
|
||||
return str(prediction_id)
|
||||
|
||||
|
||||
def poll(
|
||||
prediction_id: str,
|
||||
api_key: str,
|
||||
interval: float = 3.0,
|
||||
timeout: float = 600.0,
|
||||
request_timeout: int = 30,
|
||||
) -> dict[str, Any]:
|
||||
"""Poll a prediction until it terminates. Returns the final `data` object.
|
||||
|
||||
Raises AtlasError on reported failure or when `timeout` seconds elapse.
|
||||
"""
|
||||
import requests
|
||||
|
||||
url = f"{PREDICTION_ENDPOINT}/{prediction_id}"
|
||||
elapsed = 0.0
|
||||
last_status = "unknown"
|
||||
consecutive_transport_errors = 0
|
||||
|
||||
while elapsed < timeout:
|
||||
try:
|
||||
response = requests.get(
|
||||
url, headers=_headers(api_key, json_body=False), timeout=request_timeout
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
consecutive_transport_errors += 1
|
||||
if consecutive_transport_errors >= 5:
|
||||
raise AtlasError(
|
||||
f"Polling prediction {prediction_id} failed after "
|
||||
f"{consecutive_transport_errors} consecutive transport errors: {exc}"
|
||||
) from exc
|
||||
time.sleep(interval)
|
||||
elapsed += interval
|
||||
continue
|
||||
|
||||
_raise_for_status(response, f"Atlas Cloud poll for {prediction_id}")
|
||||
consecutive_transport_errors = 0
|
||||
data = _payload_of(response)
|
||||
last_status = str(data.get("status", "unknown")).lower()
|
||||
|
||||
if last_status in TERMINAL_SUCCESS:
|
||||
outputs = data.get("outputs") or []
|
||||
if not outputs:
|
||||
raise AtlasError(
|
||||
f"Prediction {prediction_id} reported '{last_status}' but returned no outputs."
|
||||
)
|
||||
return data
|
||||
if last_status in TERMINAL_FAILURE:
|
||||
error = data.get("error") or "no error detail provided"
|
||||
raise AtlasError(f"Atlas Cloud generation failed ({last_status}): {error}")
|
||||
|
||||
time.sleep(interval)
|
||||
elapsed += interval
|
||||
|
||||
raise AtlasError(
|
||||
f"Prediction {prediction_id} did not finish within {timeout:.0f}s "
|
||||
f"(last status: {last_status}). The job may still complete — "
|
||||
f"check {PREDICTION_ENDPOINT}/{prediction_id}"
|
||||
)
|
||||
|
||||
|
||||
def upload_media(file_path: str | Path, api_key: str, timeout: int = 120) -> str:
|
||||
"""Upload a local file and return the hosted URL Atlas assigns to it.
|
||||
|
||||
Used to turn a local reference image into the `image_url` that image-to-video
|
||||
models expect. Atlas has answered this endpoint with both {"data":
|
||||
{"download_url": ...}} and a bare {"url": ...}, so both shapes are accepted.
|
||||
"""
|
||||
import requests
|
||||
|
||||
path = Path(file_path)
|
||||
if not path.exists():
|
||||
raise AtlasError(f"Cannot upload — file not found: {path}")
|
||||
|
||||
try:
|
||||
with path.open("rb") as handle:
|
||||
response = requests.post(
|
||||
UPLOAD_MEDIA_ENDPOINT,
|
||||
headers=_headers(api_key, json_body=False),
|
||||
files={"file": (path.name, handle)},
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise AtlasError(f"Uploading {path.name} to Atlas Cloud failed: {exc}") from exc
|
||||
|
||||
_raise_for_status(response, "Atlas Cloud upload")
|
||||
data = _payload_of(response)
|
||||
|
||||
url = data.get("download_url") or data.get("url")
|
||||
if not url:
|
||||
raise AtlasError(f"Atlas Cloud upload returned no URL: {str(data)[:500]}")
|
||||
return str(url)
|
||||
|
||||
|
||||
def download(url: str, output_path: str | Path, timeout: int = 300) -> Path:
|
||||
"""Download a generated asset to disk and return the written path."""
|
||||
import requests
|
||||
|
||||
try:
|
||||
response = requests.get(url, timeout=timeout)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise AtlasError(f"Downloading Atlas Cloud output failed: {exc}") from exc
|
||||
|
||||
_raise_for_status(response, "Atlas Cloud output download")
|
||||
|
||||
path = Path(output_path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(response.content)
|
||||
return path
|
||||
|
||||
|
||||
def aspect_ratio_from_size(width: int, height: int, allowed: list[str]) -> str:
|
||||
"""Pick the closest ratio in `allowed` to width/height.
|
||||
|
||||
OpenMontage's canonical params are width/height, but many Atlas models only
|
||||
accept a ratio enum. Snapping to the nearest supported ratio beats sending a
|
||||
value the model will reject.
|
||||
"""
|
||||
if not allowed:
|
||||
return "16:9"
|
||||
if height <= 0 or width <= 0:
|
||||
return allowed[0]
|
||||
|
||||
target = width / height
|
||||
best = allowed[0]
|
||||
best_delta = float("inf")
|
||||
for ratio in allowed:
|
||||
if ratio == "auto" or ":" not in ratio:
|
||||
continue
|
||||
left, _, right = ratio.partition(":")
|
||||
try:
|
||||
candidate = float(left) / float(right)
|
||||
except (ValueError, ZeroDivisionError):
|
||||
continue
|
||||
delta = abs(candidate - target)
|
||||
if delta < best_delta:
|
||||
best_delta = delta
|
||||
best = ratio
|
||||
return best
|
||||
220
tools/atlas_models.py
Normal file
220
tools/atlas_models.py
Normal file
@@ -0,0 +1,220 @@
|
||||
"""Authoritative Atlas Cloud model catalog used by the media gateway tools.
|
||||
|
||||
The entries mirror the machine-readable schemas served by each live model page.
|
||||
Keeping task ids explicit prevents a valid text-to-video id from being rewritten
|
||||
to a sibling route that does not actually exist.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
SEEDANCE_RATIOS = ("16:9", "4:3", "1:1", "3:4", "9:16", "21:9", "adaptive")
|
||||
SEEDANCE_25_RESOLUTIONS = (
|
||||
"480p", "720p", "720p-esr", "1080p-esr", "1440p-esr", "4k-esr",
|
||||
"1080p-esr & 60fps",
|
||||
)
|
||||
SEEDANCE_20_RESOLUTIONS = (
|
||||
"480p", "720p", "720p-SR", "1080p", "1080p-SR", "1440p-SR", "4k",
|
||||
)
|
||||
H3_RATIOS = ("21:9", "16:9", "4:3", "1:1", "3:4", "9:16")
|
||||
|
||||
|
||||
def _video_spec(
|
||||
family: str,
|
||||
operation: str,
|
||||
rate: float,
|
||||
*,
|
||||
duration: tuple[int, ...] | None,
|
||||
resolution: tuple[str, ...],
|
||||
default_resolution: str,
|
||||
ratio_key: str,
|
||||
ratios: tuple[str, ...],
|
||||
default_ratio: str,
|
||||
media_style: str = "none",
|
||||
variant: str = "standard",
|
||||
optional_fields: tuple[str, ...] = (),
|
||||
media_limits: dict[str, int] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"family": family,
|
||||
"operation": operation,
|
||||
"variant": variant,
|
||||
"cost_per_second": rate,
|
||||
"durations": duration,
|
||||
"resolutions": resolution,
|
||||
"default_resolution": default_resolution,
|
||||
"ratio_key": ratio_key,
|
||||
"ratios": ratios,
|
||||
"default_ratio": default_ratio,
|
||||
"media_style": media_style,
|
||||
"optional_fields": optional_fields,
|
||||
"media_limits": media_limits or {},
|
||||
}
|
||||
|
||||
|
||||
VIDEO_MODELS: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for version, rate, durations, resolutions in (
|
||||
("2.5", 0.134, tuple(range(4, 31)) + (-1,), SEEDANCE_25_RESOLUTIONS),
|
||||
("2.0", 0.112, tuple(range(4, 16)) + (-1,), SEEDANCE_20_RESOLUTIONS),
|
||||
):
|
||||
family = f"bytedance/seedance-{version}"
|
||||
common = ("generate_audio", "watermark", "return_last_frame")
|
||||
common += ("output_format",) if version == "2.5" else ("bitrate_mode",)
|
||||
for operation, suffix, media_style in (
|
||||
("text_to_video", "text-to-video", "none"),
|
||||
("image_to_video", "image-to-video", "seedance_image"),
|
||||
("reference_to_video", "reference-to-video", "seedance_references"),
|
||||
):
|
||||
model_id = f"{family}/{suffix}"
|
||||
limits = {}
|
||||
if operation == "reference_to_video":
|
||||
limits = {"images": 30, "videos": 10, "audios": 10} if version == "2.5" else {"images": 9, "videos": 3, "audios": 3}
|
||||
VIDEO_MODELS[model_id] = _video_spec(
|
||||
family,
|
||||
operation,
|
||||
rate,
|
||||
duration=durations,
|
||||
resolution=resolutions,
|
||||
default_resolution="720p",
|
||||
ratio_key="ratio",
|
||||
ratios=("adaptive",) if operation == "image_to_video" and version == "2.5" else SEEDANCE_RATIOS,
|
||||
default_ratio="adaptive",
|
||||
media_style=media_style,
|
||||
optional_fields=common,
|
||||
media_limits=limits,
|
||||
)
|
||||
|
||||
for operation, suffix, rate, media_style in (
|
||||
("text_to_video", "text-to-video", 0.125, "none"),
|
||||
("image_to_video", "image-to-video", 0.130, "gemini_image"),
|
||||
("reference_to_video", "reference-to-video", 0.135, "gemini_images"),
|
||||
("video_edit", "video-edit", 0.140, "gemini_video_edit"),
|
||||
):
|
||||
model_id = f"google/gemini-omni-flash/{suffix}"
|
||||
VIDEO_MODELS[model_id] = _video_spec(
|
||||
"google/gemini-omni-flash",
|
||||
operation,
|
||||
rate,
|
||||
duration=None if operation == "video_edit" else tuple(range(3, 11)),
|
||||
resolution=("720p",),
|
||||
default_resolution="720p",
|
||||
ratio_key="aspect_ratio",
|
||||
ratios=("16:9", "9:16") if operation != "video_edit" else (),
|
||||
default_ratio="16:9",
|
||||
media_style=media_style,
|
||||
optional_fields=("thinking_level", "seed"),
|
||||
media_limits={"images": 10} if operation in {"reference_to_video", "video_edit"} else {},
|
||||
)
|
||||
|
||||
for operation, suffix, rate, media_style in (
|
||||
("text_to_video", "text-to-video-developer", 0.112, "none"),
|
||||
("image_to_video", "image-to-video-developer", 0.112, "gemini_images"),
|
||||
("reference_to_video", "reference-to-video-developer", 0.120, "gemini_video_clips"),
|
||||
):
|
||||
model_id = f"google/gemini-omni-flash/{suffix}"
|
||||
VIDEO_MODELS[model_id] = _video_spec(
|
||||
"google/gemini-omni-flash",
|
||||
operation,
|
||||
rate,
|
||||
duration=(4, 6, 8, 10),
|
||||
resolution=("720p", "1080p", "4k"),
|
||||
default_resolution="720p",
|
||||
ratio_key="aspect_ratio",
|
||||
ratios=("16:9", "9:16"),
|
||||
default_ratio="16:9",
|
||||
media_style=media_style,
|
||||
variant="developer",
|
||||
optional_fields=("seed",),
|
||||
)
|
||||
|
||||
for operation, suffix, media_style in (
|
||||
("text_to_video", "text-to-video", "none"),
|
||||
("image_to_video", "image-to-video", "h3_image"),
|
||||
("reference_to_video", "reference-to-video", "h3_refers"),
|
||||
):
|
||||
model_id = f"minimax/h3/{suffix}"
|
||||
VIDEO_MODELS[model_id] = _video_spec(
|
||||
"minimax/h3",
|
||||
operation,
|
||||
0.100,
|
||||
duration=tuple(range(4, 16)),
|
||||
resolution=("768P", "2K"),
|
||||
default_resolution="2K",
|
||||
ratio_key="ratio",
|
||||
ratios=("adaptive", *H3_RATIOS) if operation != "text_to_video" else H3_RATIOS,
|
||||
default_ratio="adaptive" if operation != "text_to_video" else "1:1",
|
||||
media_style=media_style,
|
||||
)
|
||||
|
||||
|
||||
IMAGE_MODELS: dict[str, dict[str, Any]] = {
|
||||
"bytedance/seedream-v5.0-pro/text-to-image": {
|
||||
"family": "bytedance/seedream-v5.0-pro", "operation": "generate",
|
||||
"cost_per_image": 0.045, "size_style": "star", "media_style": "none",
|
||||
"optional_fields": ("thinking", "prompt_optimization_mode", "enable_base64_output"),
|
||||
},
|
||||
"bytedance/seedream-v5.0-pro/edit": {
|
||||
"family": "bytedance/seedream-v5.0-pro", "operation": "edit",
|
||||
"cost_per_image": 0.045, "size_style": "star", "media_style": "images",
|
||||
"max_images": 10,
|
||||
"optional_fields": ("thinking", "prompt_optimization_mode", "enable_base64_output"),
|
||||
},
|
||||
"bytedance/seedream-v5.0-pro/layer-decomposition": {
|
||||
"family": "bytedance/seedream-v5.0-pro", "operation": "decompose",
|
||||
"cost_per_image": 0.022, "size_style": "tier", "media_style": "image",
|
||||
"max_images": 1,
|
||||
"optional_fields": ("optimize_prompt_options", "enable_sync_mode", "enable_base64_output"),
|
||||
},
|
||||
"bytedance/seedream-v5.0-lite/edit": {
|
||||
"family": "bytedance/seedream-v5.0-lite", "operation": "edit",
|
||||
"cost_per_image": 0.032, "size_style": "star", "media_style": "images",
|
||||
"max_images": 14, "optional_fields": ("enable_base64_output",),
|
||||
},
|
||||
"openai/gpt-image-2/text-to-image": {
|
||||
"family": "openai/gpt-image-2", "operation": "generate",
|
||||
"cost_per_image": 0.009, "size_style": "x", "media_style": "none",
|
||||
"optional_fields": ("quality", "enable_sync_mode", "enable_base64_output"),
|
||||
},
|
||||
"openai/gpt-image-2/edit": {
|
||||
"family": "openai/gpt-image-2", "operation": "edit",
|
||||
"cost_per_image": 0.010, "size_style": "x", "media_style": "images",
|
||||
"max_images": 10,
|
||||
"optional_fields": ("quality", "enable_sync_mode", "enable_base64_output"),
|
||||
},
|
||||
"google/nano-banana-2/text-to-image": {
|
||||
"family": "google/nano-banana-2", "operation": "generate",
|
||||
"cost_per_image": 0.080, "size_style": "ratio", "media_style": "none",
|
||||
"optional_fields": (
|
||||
"resolution", "thinking_level", "media_resolution", "enable_web_search",
|
||||
"enable_image_search", "enable_sync_mode", "enable_base64_output",
|
||||
),
|
||||
},
|
||||
"google/nano-banana-2/edit": {
|
||||
"family": "google/nano-banana-2", "operation": "edit",
|
||||
"cost_per_image": 0.080, "size_style": "ratio", "media_style": "images",
|
||||
"max_images": 14,
|
||||
"optional_fields": (
|
||||
"resolution", "thinking_level", "media_resolution", "enable_web_search",
|
||||
"enable_image_search", "enable_sync_mode", "enable_base64_output",
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def operation_routes(catalog: dict[str, dict[str, Any]]) -> dict[str, dict[str, str]]:
|
||||
"""Return family -> operation -> exact live model id."""
|
||||
routes: dict[str, dict[str, str]] = {}
|
||||
for model_id, spec in catalog.items():
|
||||
family = spec["family"]
|
||||
operation = spec["operation"]
|
||||
variant = spec.get("variant", "standard")
|
||||
route_key = operation if variant == "standard" else f"{operation}_{variant}"
|
||||
routes.setdefault(family, {})[route_key] = model_id
|
||||
return routes
|
||||
|
||||
|
||||
VIDEO_ROUTES = operation_routes(VIDEO_MODELS)
|
||||
IMAGE_ROUTES = operation_routes(IMAGE_MODELS)
|
||||
243
tools/graphics/atlas_image.py
Normal file
243
tools/graphics/atlas_image.py
Normal file
@@ -0,0 +1,243 @@
|
||||
"""Atlas Cloud image generation and editing with exact model schemas."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from tools import atlas_client
|
||||
from tools.atlas_models import IMAGE_MODELS, IMAGE_ROUTES
|
||||
from tools.base_tool import (
|
||||
BaseTool,
|
||||
Determinism,
|
||||
ExecutionMode,
|
||||
ResourceProfile,
|
||||
RetryPolicy,
|
||||
ToolResult,
|
||||
ToolRuntime,
|
||||
ToolStability,
|
||||
ToolStatus,
|
||||
ToolTier,
|
||||
)
|
||||
|
||||
_DEFAULT_MODEL = "bytedance/seedream-v5.0-pro/text-to-image"
|
||||
_DEFAULT_COST = 0.04
|
||||
_COMMON_RATIOS = ["1:1", "3:2", "2:3", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9", "21:9"]
|
||||
|
||||
|
||||
def _is_remote(entry: str) -> bool:
|
||||
return str(entry).strip().lower().startswith(("http://", "https://", "data:", "asset://"))
|
||||
|
||||
|
||||
class AtlasImage(BaseTool):
|
||||
name = "atlas_image"
|
||||
version = "0.2.0"
|
||||
tier = ToolTier.GENERATE
|
||||
capability = "image_generation"
|
||||
provider = "atlascloud"
|
||||
stability = ToolStability.BETA
|
||||
execution_mode = ExecutionMode.SYNC
|
||||
determinism = Determinism.STOCHASTIC
|
||||
runtime = ToolRuntime.API
|
||||
|
||||
dependencies = ["env:ATLASCLOUD_API_KEY"]
|
||||
install_instructions = atlas_client.INSTALL_INSTRUCTIONS
|
||||
agent_skills = ["atlas-cloud", "flux-best-practices"]
|
||||
|
||||
capabilities = ["generate_image", "text_to_image", "image_edit", "layer_decomposition"]
|
||||
supports = {
|
||||
"custom_size": True, "aspect_ratio": True, "image_edit": True,
|
||||
"multiple_reference_images": True, "layer_decomposition": True,
|
||||
"multi_model_gateway": True,
|
||||
}
|
||||
provider_matrix = {
|
||||
family: {key: value for key, value in routes.items()}
|
||||
for family, routes in IMAGE_ROUTES.items()
|
||||
}
|
||||
best_for = [
|
||||
"Seedream 5.0 Pro generation, multi-image editing, and layer decomposition",
|
||||
"GPT Image 2 generation and editing with arbitrary supported dimensions",
|
||||
"Nano Banana 2 generation and editing up to 4K with as many as 14 references",
|
||||
]
|
||||
not_good_for = ["offline generation", "assuming one parameter schema fits every Atlas model"]
|
||||
fallback_tools = ["flux_image", "google_imagen", "openai_image", "recraft_image"]
|
||||
quality_score = 0.86
|
||||
|
||||
input_schema = {
|
||||
"type": "object",
|
||||
"required": ["prompt"],
|
||||
"properties": {
|
||||
"prompt": {"type": "string"},
|
||||
"model": {"type": "string", "default": _DEFAULT_MODEL, "enum": sorted(IMAGE_MODELS)},
|
||||
"generation_mode": {"type": "string", "enum": ["generate", "edit", "decompose"], "default": "generate"},
|
||||
"width": {"type": "integer", "default": 2048},
|
||||
"height": {"type": "integer", "default": 2048},
|
||||
"aspect_ratio": {"type": "string"},
|
||||
"resolution": {"type": "string"},
|
||||
"quality": {"type": "string", "enum": ["low", "medium", "high"]},
|
||||
"thinking": {"type": "string", "enum": ["enabled", "disabled"]},
|
||||
"prompt_optimization_mode": {"type": "string", "enum": ["standard", "fast"]},
|
||||
"thinking_level": {"type": "string", "enum": ["default", "high", "minimal"]},
|
||||
"media_resolution": {"type": "string", "enum": ["default", "low", "medium", "high"]},
|
||||
"enable_web_search": {"type": "boolean"},
|
||||
"enable_image_search": {"type": "boolean"},
|
||||
"image_url": {"type": "string"},
|
||||
"image_path": {"type": "string"},
|
||||
"image_urls": {"type": "array", "items": {"type": "string"}},
|
||||
"image_paths": {"type": "array", "items": {"type": "string"}},
|
||||
"output_format": {"type": "string", "enum": ["default", "jpeg", "png"], "default": "png"},
|
||||
"extra_params": {"type": "object"},
|
||||
"poll_interval": {"type": "number", "default": 2.0},
|
||||
"poll_timeout": {"type": "number", "default": 600.0},
|
||||
"output_path": {"type": "string"},
|
||||
},
|
||||
}
|
||||
|
||||
resource_profile = ResourceProfile(cpu_cores=1, ram_mb=512, disk_mb=250, network_required=True)
|
||||
retry_policy = RetryPolicy(max_retries=2, retryable_errors=["rate_limit", "timeout"])
|
||||
idempotency_key_fields = ["prompt", "model", "generation_mode", "width", "height", "aspect_ratio"]
|
||||
side_effects = ["writes image files to output_path", "calls Atlas Cloud API"]
|
||||
user_visible_verification = ["Inspect generated images for prompt fidelity and edit consistency"]
|
||||
|
||||
def get_status(self) -> ToolStatus:
|
||||
return ToolStatus.AVAILABLE if atlas_client.get_api_key() else ToolStatus.UNAVAILABLE
|
||||
|
||||
def get_info(self) -> dict[str, Any]:
|
||||
info = super().get_info()
|
||||
info["model_catalog"] = {
|
||||
model_id: {
|
||||
"family": spec["family"], "operation": spec["operation"],
|
||||
"cost_per_image": spec["cost_per_image"], "size_style": spec["size_style"],
|
||||
"media_style": spec["media_style"], "max_images": spec.get("max_images"),
|
||||
}
|
||||
for model_id, spec in IMAGE_MODELS.items()
|
||||
}
|
||||
return info
|
||||
|
||||
def estimate_cost(self, inputs: dict[str, Any]) -> float:
|
||||
try:
|
||||
model = self._resolve_model(
|
||||
str(inputs.get("model", _DEFAULT_MODEL)),
|
||||
str(inputs.get("generation_mode", "generate")),
|
||||
)
|
||||
except ValueError:
|
||||
return _DEFAULT_COST
|
||||
return float(IMAGE_MODELS[model]["cost_per_image"])
|
||||
|
||||
def estimate_runtime(self, inputs: dict[str, Any]) -> float:
|
||||
return 30.0
|
||||
|
||||
def _resolve_model(self, model: str, operation: str) -> str:
|
||||
if model not in IMAGE_MODELS:
|
||||
raise ValueError(
|
||||
f"Unsupported Atlas image model id {model!r}. Use get_info()['model_catalog'] for live routes."
|
||||
)
|
||||
family = IMAGE_MODELS[model]["family"]
|
||||
route = IMAGE_ROUTES.get(family, {}).get(operation)
|
||||
if not route:
|
||||
raise ValueError(f"{family} does not expose generation_mode={operation!r} on Atlas Cloud")
|
||||
return route
|
||||
|
||||
def _build_payload(self, inputs: dict[str, Any], model: str) -> dict[str, Any]:
|
||||
spec = IMAGE_MODELS[model]
|
||||
payload: dict[str, Any] = {"model": model, "prompt": inputs.get("prompt", "")}
|
||||
width = int(inputs.get("width", 2048))
|
||||
height = int(inputs.get("height", 2048))
|
||||
style = spec["size_style"]
|
||||
|
||||
if style == "star":
|
||||
payload["size"] = f"{width}*{height}"
|
||||
elif style == "x":
|
||||
payload["size"] = f"{width}x{height}"
|
||||
elif style == "ratio":
|
||||
payload["aspect_ratio"] = inputs.get("aspect_ratio") or atlas_client.aspect_ratio_from_size(
|
||||
width, height, _COMMON_RATIOS
|
||||
)
|
||||
elif style == "tier":
|
||||
payload["size"] = inputs.get("resolution", "auto")
|
||||
|
||||
images = list(inputs.get("image_urls") or [])
|
||||
if inputs.get("image_url"):
|
||||
images.insert(0, inputs["image_url"])
|
||||
media_style = spec["media_style"]
|
||||
if media_style == "images":
|
||||
maximum = int(spec["max_images"])
|
||||
if not images:
|
||||
raise ValueError(f"{spec['operation']} requires at least one source image")
|
||||
if len(images) > maximum:
|
||||
raise ValueError(f"{model} accepts at most {maximum} source images")
|
||||
payload["images"] = images
|
||||
elif media_style == "image":
|
||||
if len(images) != 1:
|
||||
raise ValueError("layer decomposition requires exactly one source image")
|
||||
payload["image"] = images[0]
|
||||
|
||||
for field in spec.get("optional_fields", ()):
|
||||
if inputs.get(field) is not None:
|
||||
payload[field] = inputs[field]
|
||||
if inputs.get("output_format") and inputs["output_format"] != "default":
|
||||
payload["output_format"] = inputs["output_format"]
|
||||
|
||||
extra = inputs.get("extra_params")
|
||||
if isinstance(extra, dict):
|
||||
payload.update(extra)
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _upload_value(value: str, api_key: str) -> str:
|
||||
return value if _is_remote(value) else atlas_client.upload_media(value, api_key)
|
||||
|
||||
def _resolve_media(self, inputs: dict[str, Any], api_key: str) -> dict[str, Any]:
|
||||
resolved = dict(inputs)
|
||||
urls = [str(value) for value in resolved.get("image_urls", [])]
|
||||
paths = [str(value) for value in resolved.get("image_paths", [])]
|
||||
if resolved.get("image_url"):
|
||||
urls.insert(0, str(resolved["image_url"]))
|
||||
if resolved.get("image_path"):
|
||||
paths.insert(0, str(resolved["image_path"]))
|
||||
resolved["image_urls"] = [self._upload_value(value, api_key) for value in [*urls, *paths]]
|
||||
resolved.pop("image_url", None)
|
||||
return resolved
|
||||
|
||||
def execute(self, inputs: dict[str, Any]) -> ToolResult:
|
||||
api_key = atlas_client.get_api_key()
|
||||
if not api_key:
|
||||
return ToolResult(success=False, error="ATLASCLOUD_API_KEY not set. " + self.install_instructions)
|
||||
|
||||
started = time.time()
|
||||
operation = str(inputs.get("generation_mode", "generate"))
|
||||
try:
|
||||
model = self._resolve_model(str(inputs.get("model", _DEFAULT_MODEL)), operation)
|
||||
resolved = self._resolve_media(inputs, api_key)
|
||||
payload = self._build_payload(resolved, model)
|
||||
prediction_id = atlas_client.submit(atlas_client.GENERATE_IMAGE_ENDPOINT, payload, api_key)
|
||||
data = atlas_client.poll(
|
||||
prediction_id, api_key,
|
||||
interval=float(inputs.get("poll_interval", 2.0)),
|
||||
timeout=float(inputs.get("poll_timeout", 600.0)),
|
||||
)
|
||||
outputs = list(data["outputs"])
|
||||
requested = Path(inputs.get("output_path") or "atlas_image.png")
|
||||
output_paths: list[Path] = []
|
||||
for index, url in enumerate(outputs):
|
||||
output_path = requested if index == 0 else requested.with_name(f"{requested.stem}_{index + 1}{requested.suffix}")
|
||||
atlas_client.download(url, output_path)
|
||||
output_paths.append(output_path)
|
||||
except (atlas_client.AtlasError, ValueError, KeyError) as exc:
|
||||
return ToolResult(success=False, error=f"Atlas Cloud image generation failed: {exc}")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return ToolResult(success=False, error=f"Atlas Cloud image generation failed: {exc}")
|
||||
|
||||
return ToolResult(
|
||||
success=True,
|
||||
data={
|
||||
"provider": "atlascloud", "model": model, "prompt": inputs.get("prompt", ""),
|
||||
"generation_mode": operation, "output": str(output_paths[0]),
|
||||
"output_path": str(output_paths[0]), "outputs": [str(path) for path in output_paths],
|
||||
"prediction_id": prediction_id, "source_url": outputs[0], "source_urls": outputs,
|
||||
"request_params": payload,
|
||||
},
|
||||
artifacts=[str(path) for path in output_paths], cost_usd=self.estimate_cost({**inputs, "model": model}),
|
||||
duration_seconds=round(time.time() - started, 2), model=model,
|
||||
)
|
||||
@@ -20,7 +20,7 @@ class ImageSelector(BaseTool):
|
||||
provider = "selector"
|
||||
stability = ToolStability.BETA
|
||||
runtime = ToolRuntime.HYBRID
|
||||
agent_skills = ["flux-best-practices", "bfl-api"]
|
||||
agent_skills = ["flux-best-practices", "bfl-api", "atlas-cloud"]
|
||||
|
||||
capabilities = [
|
||||
"generate_image", "search_image", "download_image",
|
||||
@@ -69,6 +69,10 @@ class ImageSelector(BaseTool):
|
||||
"type": "string",
|
||||
"description": "Provider-specific model name passed through when supported.",
|
||||
},
|
||||
"model": {
|
||||
"type": "string",
|
||||
"description": "Exact provider model id, e.g. an Atlas Cloud live model route.",
|
||||
},
|
||||
"generation_mode": {
|
||||
"type": "string",
|
||||
"enum": ["generate", "edit"],
|
||||
@@ -289,6 +293,7 @@ class ImageSelector(BaseTool):
|
||||
"element_list",
|
||||
"api_family",
|
||||
"model_name",
|
||||
"model",
|
||||
"image_reference",
|
||||
"image_fidelity",
|
||||
"human_fidelity",
|
||||
@@ -396,6 +401,16 @@ class ImageSelector(BaseTool):
|
||||
return serialized
|
||||
|
||||
def _filter_candidates(self, inputs: dict[str, Any], candidates: list[BaseTool]) -> list[BaseTool]:
|
||||
exact_model = inputs.get("model")
|
||||
if exact_model:
|
||||
model_matches = [
|
||||
tool for tool in candidates
|
||||
if exact_model in getattr(tool, "input_schema", {}).get("properties", {}).get("model", {}).get("enum", [])
|
||||
or exact_model in tool.get_info().get("model_catalog", {})
|
||||
]
|
||||
if model_matches:
|
||||
candidates = model_matches
|
||||
|
||||
# A caller-supplied custom workflow is provider-specific (ComfyUI graph
|
||||
# JSON). Route it only to custom-workflow-capable providers whose server
|
||||
# is reachable — bundled-model readiness is irrelevant in that case.
|
||||
|
||||
396
tools/video/atlas_video.py
Normal file
396
tools/video/atlas_video.py
Normal file
@@ -0,0 +1,396 @@
|
||||
"""Atlas Cloud video generation with exact, discoverable live-model routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import mimetypes
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from tools import atlas_client
|
||||
from tools.atlas_models import VIDEO_MODELS, VIDEO_ROUTES
|
||||
from tools.base_tool import (
|
||||
BaseTool,
|
||||
Determinism,
|
||||
ExecutionMode,
|
||||
ResourceProfile,
|
||||
RetryPolicy,
|
||||
ToolResult,
|
||||
ToolRuntime,
|
||||
ToolStability,
|
||||
ToolStatus,
|
||||
ToolTier,
|
||||
)
|
||||
|
||||
_DEFAULT_MODEL = "bytedance/seedance-2.5/text-to-video"
|
||||
_DEFAULT_COST_PER_SECOND = 0.10
|
||||
_OPERATIONS = ("text_to_video", "image_to_video", "reference_to_video", "video_edit")
|
||||
|
||||
|
||||
def _is_remote(entry: str) -> bool:
|
||||
lowered = str(entry).strip().lower()
|
||||
return lowered.startswith(("http://", "https://", "data:", "asset://"))
|
||||
|
||||
|
||||
def _media_type(entry: str) -> str:
|
||||
mime, _ = mimetypes.guess_type(str(entry).split("?", 1)[0])
|
||||
if mime:
|
||||
return mime.split("/", 1)[0]
|
||||
suffix = Path(str(entry).split("?", 1)[0]).suffix.lower()
|
||||
if suffix in {".png", ".jpg", ".jpeg", ".webp", ".bmp", ".gif", ".heic", ".heif"}:
|
||||
return "image"
|
||||
if suffix in {".mp4", ".mov", ".webm", ".mkv"}:
|
||||
return "video"
|
||||
if suffix in {".mp3", ".wav", ".m4a", ".aac", ".flac"}:
|
||||
return "audio"
|
||||
raise ValueError(f"Cannot infer Atlas reference media type from {entry!r}; pass `refers` with an explicit type")
|
||||
|
||||
|
||||
class AtlasVideo(BaseTool):
|
||||
name = "atlas_video"
|
||||
version = "0.2.0"
|
||||
tier = ToolTier.GENERATE
|
||||
capability = "video_generation"
|
||||
provider = "atlascloud"
|
||||
stability = ToolStability.BETA
|
||||
execution_mode = ExecutionMode.SYNC
|
||||
determinism = Determinism.STOCHASTIC
|
||||
runtime = ToolRuntime.API
|
||||
|
||||
dependencies = ["env:ATLASCLOUD_API_KEY"]
|
||||
install_instructions = atlas_client.INSTALL_INSTRUCTIONS
|
||||
agent_skills = ["atlas-cloud", "ai-video-gen", "seedance-2-0", "gemini-omni"]
|
||||
|
||||
capabilities = list(_OPERATIONS)
|
||||
supports = {
|
||||
"text_to_video": True,
|
||||
"image_to_video": True,
|
||||
"reference_to_video": True,
|
||||
"video_edit": True,
|
||||
"first_last_frame": True,
|
||||
"mixed_media_references": True,
|
||||
"native_audio": True,
|
||||
"custom_duration": True,
|
||||
"aspect_ratio": True,
|
||||
"multi_model_gateway": True,
|
||||
}
|
||||
provider_matrix = {
|
||||
family: {key: value for key, value in routes.items()}
|
||||
for family, routes in VIDEO_ROUTES.items()
|
||||
}
|
||||
best_for = [
|
||||
"Seedance 2.5/2.0 text, image, and mixed-reference video through Atlas Cloud",
|
||||
"Gemini Omni Flash text, image, reference, and video-edit workflows",
|
||||
"MiniMax H3 text, start/end image, and mixed-media reference generation up to 2K",
|
||||
"one Atlas key with exact per-model request validation and cost estimates",
|
||||
]
|
||||
not_good_for = ["offline generation", "clips longer than the selected model permits"]
|
||||
fallback_tools = ["seedance_video", "kling_video", "minimax_video", "veo_video"]
|
||||
quality_score = 0.86
|
||||
|
||||
input_schema = {
|
||||
"type": "object",
|
||||
"required": ["prompt"],
|
||||
"properties": {
|
||||
"prompt": {"type": "string"},
|
||||
"model": {
|
||||
"type": "string",
|
||||
"default": _DEFAULT_MODEL,
|
||||
"enum": sorted(VIDEO_MODELS),
|
||||
"description": "Exact live Atlas model id. See provider_matrix in get_info() for supported sibling routes.",
|
||||
},
|
||||
"operation": {"type": "string", "enum": list(_OPERATIONS), "default": "text_to_video"},
|
||||
"model_variant": {
|
||||
"type": "string", "enum": ["standard", "developer"], "default": "standard",
|
||||
"description": "Select Gemini Omni's developer route when it exists.",
|
||||
},
|
||||
"duration": {"type": "integer", "default": 10},
|
||||
"aspect_ratio": {"type": "string", "default": "16:9"},
|
||||
"resolution": {"type": "string"},
|
||||
"seed": {"type": "integer"},
|
||||
"thinking_level": {"type": "string", "enum": ["default", "high", "low"]},
|
||||
"generate_audio": {"type": "boolean"},
|
||||
"watermark": {"type": "boolean"},
|
||||
"return_last_frame": {"type": "boolean"},
|
||||
"output_format": {"type": "string", "enum": ["mp4", "mov"]},
|
||||
"bitrate_mode": {"type": "string", "enum": ["standard", "high"]},
|
||||
"image_url": {"type": "string"},
|
||||
"image_path": {"type": "string"},
|
||||
"reference_image_url": {"type": "string"},
|
||||
"reference_image_path": {"type": "string"},
|
||||
"last_image_url": {"type": "string"},
|
||||
"last_image_path": {"type": "string"},
|
||||
"end_image_url": {"type": "string"},
|
||||
"end_image_path": {"type": "string"},
|
||||
"reference_images": {"type": "array", "items": {"type": "string"}},
|
||||
"reference_videos": {"type": "array", "items": {"type": "string"}},
|
||||
"reference_audios": {"type": "array", "items": {"type": "string"}},
|
||||
"video_url": {"type": "string"},
|
||||
"video_path": {"type": "string"},
|
||||
"reference_video_url": {"type": "string"},
|
||||
"reference_video_path": {"type": "string"},
|
||||
"video_clips": {"type": "array", "items": {"type": "object"}},
|
||||
"refers": {"type": "array", "items": {"type": "object"}},
|
||||
"extra_params": {"type": "object"},
|
||||
"poll_interval": {"type": "number", "default": 5.0},
|
||||
"poll_timeout": {"type": "number", "default": 1200.0},
|
||||
"output_path": {"type": "string"},
|
||||
},
|
||||
}
|
||||
|
||||
resource_profile = ResourceProfile(cpu_cores=1, ram_mb=512, disk_mb=500, network_required=True)
|
||||
retry_policy = RetryPolicy(max_retries=2, retryable_errors=["rate_limit", "timeout"])
|
||||
idempotency_key_fields = ["prompt", "model", "operation", "duration", "aspect_ratio"]
|
||||
side_effects = ["writes video file to output_path", "calls Atlas Cloud API"]
|
||||
user_visible_verification = ["Watch the generated clip for prompt fidelity, motion coherence, and audio quality"]
|
||||
|
||||
def get_status(self) -> ToolStatus:
|
||||
return ToolStatus.AVAILABLE if atlas_client.get_api_key() else ToolStatus.UNAVAILABLE
|
||||
|
||||
def get_info(self) -> dict[str, Any]:
|
||||
info = super().get_info()
|
||||
info["model_catalog"] = {
|
||||
model_id: {
|
||||
"family": spec["family"],
|
||||
"operation": spec["operation"],
|
||||
"variant": spec.get("variant", "standard"),
|
||||
"cost_per_second": spec["cost_per_second"],
|
||||
"durations": list(spec["durations"]) if spec["durations"] else None,
|
||||
"resolutions": list(spec["resolutions"]),
|
||||
"media_style": spec["media_style"],
|
||||
}
|
||||
for model_id, spec in VIDEO_MODELS.items()
|
||||
}
|
||||
return info
|
||||
|
||||
@staticmethod
|
||||
def _family(model: str) -> str:
|
||||
spec = VIDEO_MODELS.get(model)
|
||||
return spec["family"] if spec else "/".join(model.split("/")[:2])
|
||||
|
||||
def estimate_cost(self, inputs: dict[str, Any]) -> float:
|
||||
model = self._resolve_model(
|
||||
str(inputs.get("model", _DEFAULT_MODEL)),
|
||||
str(inputs.get("operation", "text_to_video")),
|
||||
str(inputs["model_variant"]) if inputs.get("model_variant") else None,
|
||||
)
|
||||
rate = VIDEO_MODELS.get(model, {}).get("cost_per_second", _DEFAULT_COST_PER_SECOND)
|
||||
duration = int(inputs.get("duration", 10))
|
||||
return round(rate * max(duration, 0), 4)
|
||||
|
||||
def estimate_runtime(self, inputs: dict[str, Any]) -> float:
|
||||
return 180.0
|
||||
|
||||
def is_operation_available(self, operation: str) -> bool:
|
||||
return operation in _OPERATIONS
|
||||
|
||||
def _resolve_model(self, model: str, operation: str, variant: str | None = None) -> str:
|
||||
if model not in VIDEO_MODELS:
|
||||
raise ValueError(
|
||||
f"Unsupported Atlas video model id {model!r}. Use get_info()['model_catalog'] for live routes."
|
||||
)
|
||||
spec = VIDEO_MODELS[model]
|
||||
variant = variant or str(spec.get("variant", "standard"))
|
||||
route_key = operation if variant == "standard" else f"{operation}_{variant}"
|
||||
resolved = VIDEO_ROUTES.get(spec["family"], {}).get(route_key)
|
||||
if not resolved:
|
||||
raise ValueError(
|
||||
f"{spec['family']} does not expose operation={operation!r}, variant={variant!r} on Atlas Cloud"
|
||||
)
|
||||
return resolved
|
||||
|
||||
@staticmethod
|
||||
def _validate_choice(name: str, value: Any, allowed: tuple[Any, ...] | None) -> Any:
|
||||
if allowed and value not in allowed:
|
||||
raise ValueError(f"{name}={value!r} is not supported; choose one of {list(allowed)}")
|
||||
return value
|
||||
|
||||
def _build_payload(self, inputs: dict[str, Any], model: str) -> dict[str, Any]:
|
||||
spec = VIDEO_MODELS[model]
|
||||
payload: dict[str, Any] = {"model": model, "prompt": inputs.get("prompt", "")}
|
||||
|
||||
if spec["operation"] != "video_edit":
|
||||
duration = int(inputs.get("duration", 10))
|
||||
payload["duration"] = self._validate_choice("duration", duration, spec["durations"])
|
||||
ratio = inputs.get("aspect_ratio", spec["default_ratio"])
|
||||
if ratio == "16:9" and spec["default_ratio"] == "adaptive" and spec["ratios"] == ("adaptive",):
|
||||
ratio = "adaptive"
|
||||
payload[spec["ratio_key"]] = self._validate_choice("aspect_ratio", ratio, spec["ratios"])
|
||||
|
||||
resolution = inputs.get("resolution", spec["default_resolution"])
|
||||
payload["resolution"] = self._validate_choice("resolution", resolution, spec["resolutions"])
|
||||
|
||||
for field in spec.get("optional_fields", ()):
|
||||
if inputs.get(field) is not None:
|
||||
payload[field] = inputs[field]
|
||||
|
||||
style = spec["media_style"]
|
||||
image = inputs.get("image_url") or inputs.get("reference_image_url")
|
||||
last_image = inputs.get("last_image_url") or inputs.get("end_image_url")
|
||||
images = list(inputs.get("reference_images") or [])
|
||||
videos = list(inputs.get("reference_videos") or [])
|
||||
audios = list(inputs.get("reference_audios") or [])
|
||||
video = inputs.get("video_url") or inputs.get("reference_video_url")
|
||||
|
||||
if style in {"seedance_image", "gemini_image", "h3_image"}:
|
||||
if not image:
|
||||
raise ValueError("image_to_video requires image_url, image_path, or reference_image_path")
|
||||
payload["image"] = image
|
||||
if last_image:
|
||||
payload["last_image" if style == "seedance_image" else "end_image"] = last_image
|
||||
elif style == "gemini_images":
|
||||
if image and not images:
|
||||
images = [image]
|
||||
if not images:
|
||||
raise ValueError("This Gemini route requires at least one reference image")
|
||||
payload["images"] = images
|
||||
elif style == "seedance_references":
|
||||
if image and not images:
|
||||
images = [image]
|
||||
if not (images or videos or (audios and spec["family"] == "bytedance/seedance-2.5")):
|
||||
raise ValueError("reference_to_video requires supported reference media for the selected model")
|
||||
limits = spec["media_limits"]
|
||||
if len(images) > limits["images"] or len(videos) > limits["videos"] or len(audios) > limits["audios"]:
|
||||
raise ValueError(
|
||||
f"{model} accepts at most {limits['images']} images, {limits['videos']} videos, "
|
||||
f"and {limits['audios']} audio references"
|
||||
)
|
||||
if images:
|
||||
payload["reference_images"] = images
|
||||
if videos:
|
||||
payload["reference_videos"] = videos
|
||||
if audios:
|
||||
payload["reference_audios"] = audios
|
||||
elif style == "h3_refers":
|
||||
refers = list(inputs.get("refers") or [])
|
||||
if not refers:
|
||||
refers = [
|
||||
*({"url": value, "type": "image"} for value in images),
|
||||
*({"url": value, "type": "video"} for value in videos),
|
||||
*({"url": value, "type": "audio"} for value in audios),
|
||||
]
|
||||
if image:
|
||||
refers.insert(0, {"url": image, "type": "image"})
|
||||
if not refers or not any(item.get("type") in {"image", "video"} for item in refers):
|
||||
raise ValueError("MiniMax H3 reference_to_video requires at least one image or video in refers")
|
||||
payload["refers"] = refers
|
||||
elif style == "gemini_video_clips":
|
||||
clips = list(inputs.get("video_clips") or [])
|
||||
if not clips and video:
|
||||
clips = [{"url": video, "start": 0, "ends": min(int(inputs.get("duration", 10)), 10)}]
|
||||
if len(clips) != 1:
|
||||
raise ValueError("Gemini Omni developer reference_to_video requires exactly one video_clip")
|
||||
payload["video_clips"] = clips
|
||||
if images:
|
||||
payload["images"] = images
|
||||
elif style == "gemini_video_edit":
|
||||
if not video:
|
||||
raise ValueError("video_edit requires video_url, video_path, or reference_video_path")
|
||||
payload["video"] = video
|
||||
if images:
|
||||
if len(images) > spec["media_limits"].get("images", 10):
|
||||
raise ValueError("Gemini Omni video_edit accepts at most 10 reference images")
|
||||
payload["images"] = images
|
||||
|
||||
extra = inputs.get("extra_params")
|
||||
if isinstance(extra, dict):
|
||||
payload.update(extra)
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _upload_value(value: str | None, api_key: str) -> str | None:
|
||||
if not value or _is_remote(value):
|
||||
return value
|
||||
return atlas_client.upload_media(value, api_key)
|
||||
|
||||
def _resolve_media(self, inputs: dict[str, Any], api_key: str) -> dict[str, Any]:
|
||||
resolved = dict(inputs)
|
||||
resolved["reference_images"] = [
|
||||
*(resolved.get("reference_images") or []),
|
||||
*(resolved.get("reference_image_urls") or []),
|
||||
*(resolved.get("reference_image_paths") or []),
|
||||
]
|
||||
resolved["reference_videos"] = [
|
||||
*(resolved.get("reference_videos") or []),
|
||||
*(resolved.get("reference_video_urls") or []),
|
||||
*(resolved.get("reference_video_paths") or []),
|
||||
]
|
||||
resolved["reference_audios"] = [
|
||||
*(resolved.get("reference_audios") or []),
|
||||
*(resolved.get("reference_audio_urls") or []),
|
||||
*(resolved.get("reference_audio_paths") or []),
|
||||
]
|
||||
aliases = {
|
||||
"image_url": ("image_url", "reference_image_url", "image_path", "reference_image_path"),
|
||||
"last_image_url": ("last_image_url", "end_image_url", "last_image_path", "end_image_path"),
|
||||
"video_url": ("video_url", "reference_video_url", "video_path", "reference_video_path"),
|
||||
}
|
||||
for target, sources in aliases.items():
|
||||
value = next((resolved.get(key) for key in sources if resolved.get(key)), None)
|
||||
if value:
|
||||
resolved[target] = self._upload_value(str(value), api_key)
|
||||
|
||||
for key in ("reference_images", "reference_videos", "reference_audios"):
|
||||
resolved[key] = [self._upload_value(str(value), api_key) for value in resolved.get(key, [])]
|
||||
|
||||
if resolved.get("refers"):
|
||||
normalized = []
|
||||
for item in resolved["refers"]:
|
||||
entry = dict(item)
|
||||
entry["url"] = self._upload_value(str(entry["url"]), api_key)
|
||||
entry.setdefault("type", _media_type(str(item["url"])))
|
||||
normalized.append(entry)
|
||||
resolved["refers"] = normalized
|
||||
|
||||
if resolved.get("video_clips"):
|
||||
clips = []
|
||||
for item in resolved["video_clips"]:
|
||||
clip = dict(item)
|
||||
clip["url"] = self._upload_value(str(clip["url"]), api_key)
|
||||
clips.append(clip)
|
||||
resolved["video_clips"] = clips
|
||||
return resolved
|
||||
|
||||
def execute(self, inputs: dict[str, Any]) -> ToolResult:
|
||||
api_key = atlas_client.get_api_key()
|
||||
if not api_key:
|
||||
return ToolResult(success=False, error="ATLASCLOUD_API_KEY not set. " + self.install_instructions)
|
||||
|
||||
started = time.time()
|
||||
operation = str(inputs.get("operation", "text_to_video"))
|
||||
try:
|
||||
model = self._resolve_model(
|
||||
str(inputs.get("model", _DEFAULT_MODEL)), operation,
|
||||
str(inputs["model_variant"]) if inputs.get("model_variant") else None,
|
||||
)
|
||||
resolved = self._resolve_media(inputs, api_key)
|
||||
payload = self._build_payload(resolved, model)
|
||||
prediction_id = atlas_client.submit(atlas_client.GENERATE_VIDEO_ENDPOINT, payload, api_key)
|
||||
data = atlas_client.poll(
|
||||
prediction_id, api_key,
|
||||
interval=float(inputs.get("poll_interval", 5.0)),
|
||||
timeout=float(inputs.get("poll_timeout", 1200.0)),
|
||||
)
|
||||
suffix = str(inputs.get("output_format", "mp4"))
|
||||
output_path = Path(inputs.get("output_path") or f"atlas_video.{suffix}")
|
||||
atlas_client.download(data["outputs"][0], output_path)
|
||||
except (atlas_client.AtlasError, ValueError, KeyError) as exc:
|
||||
return ToolResult(success=False, error=f"Atlas Cloud video generation failed: {exc}")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return ToolResult(success=False, error=f"Atlas Cloud video generation failed: {exc}")
|
||||
|
||||
from tools.video._shared import probe_output
|
||||
|
||||
probed = probe_output(output_path)
|
||||
cost_inputs = {**inputs, "model": model, "operation": operation}
|
||||
return ToolResult(
|
||||
success=True,
|
||||
data={
|
||||
"provider": "atlascloud", "model": model, "prompt": inputs["prompt"],
|
||||
"operation": operation, "output": str(output_path), "output_path": str(output_path),
|
||||
"prediction_id": prediction_id, "source_url": data["outputs"][0],
|
||||
"format": output_path.suffix.lstrip("."), "request_params": payload, **probed,
|
||||
},
|
||||
artifacts=[str(output_path)], cost_usd=self.estimate_cost(cost_inputs),
|
||||
duration_seconds=round(time.time() - started, 2), model=model,
|
||||
)
|
||||
@@ -20,16 +20,16 @@ class VideoSelector(BaseTool):
|
||||
provider = "selector"
|
||||
stability = ToolStability.BETA
|
||||
runtime = ToolRuntime.HYBRID
|
||||
agent_skills = ["ai-video-gen", "create-video", "ltx2", "gemini-omni"]
|
||||
agent_skills = ["ai-video-gen", "create-video", "ltx2", "gemini-omni", "atlas-cloud"]
|
||||
|
||||
# Operations that REQUIRE motion: an image-only tool (image_selector) is not
|
||||
# an acceptable last-resort fallback for these, so fallback_tools_for() drops it.
|
||||
MOTION_REQUIRED_OPERATIONS = frozenset({"image_to_video", "reference_to_video"})
|
||||
MOTION_REQUIRED_OPERATIONS = frozenset({"image_to_video", "reference_to_video", "video_edit"})
|
||||
# Default score gap for the preferred_provider override (see input_schema).
|
||||
PREFERRED_PROVIDER_GAP = 0.15
|
||||
|
||||
capabilities = [
|
||||
"text_to_video", "image_to_video", "stock_video",
|
||||
"text_to_video", "image_to_video", "reference_to_video", "video_edit", "stock_video",
|
||||
"provider_selection", "search_video", "download_video",
|
||||
]
|
||||
supports = {
|
||||
@@ -70,12 +70,12 @@ class VideoSelector(BaseTool):
|
||||
"allowed_providers": {"type": "array", "items": {"type": "string"}},
|
||||
"operation": {
|
||||
"type": "string",
|
||||
"enum": ["text_to_video", "image_to_video", "reference_to_video", "rank"],
|
||||
"enum": ["text_to_video", "image_to_video", "reference_to_video", "video_edit", "rank"],
|
||||
"default": "text_to_video",
|
||||
},
|
||||
"target_operation": {
|
||||
"type": "string",
|
||||
"enum": ["text_to_video", "image_to_video", "reference_to_video"],
|
||||
"enum": ["text_to_video", "image_to_video", "reference_to_video", "video_edit"],
|
||||
"description": "Operation to score when operation='rank'.",
|
||||
"default": "text_to_video",
|
||||
},
|
||||
@@ -115,6 +115,32 @@ class VideoSelector(BaseTool):
|
||||
"type": "string",
|
||||
"description": "Local reference video path. Providers that require URLs should reject this clearly.",
|
||||
},
|
||||
"reference_video_urls": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Reference video URLs for mixed-media generation.",
|
||||
},
|
||||
"reference_video_paths": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Local reference video paths for mixed-media generation.",
|
||||
},
|
||||
"reference_audio_urls": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Reference audio URLs for mixed-media generation.",
|
||||
},
|
||||
"reference_audio_paths": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Local reference audio paths for mixed-media generation.",
|
||||
},
|
||||
"last_image_url": {"type": "string", "description": "Optional final frame for first/last-frame generation."},
|
||||
"last_image_path": {"type": "string", "description": "Optional local final frame."},
|
||||
"video_url": {"type": "string", "description": "Source video URL for video editing."},
|
||||
"video_path": {"type": "string", "description": "Local source video for video editing."},
|
||||
"video_clips": {"type": "array", "items": {"type": "object"}},
|
||||
"refers": {"type": "array", "items": {"type": "object"}},
|
||||
"image_list": {
|
||||
"type": "array",
|
||||
"description": "Provider-specific list of image references, e.g. Kling Official Video Omni.",
|
||||
@@ -155,6 +181,14 @@ class VideoSelector(BaseTool):
|
||||
"type": "string",
|
||||
"description": "Provider-specific model name passed through when supported.",
|
||||
},
|
||||
"model": {
|
||||
"type": "string",
|
||||
"description": "Exact provider model id, e.g. an Atlas Cloud live model route.",
|
||||
},
|
||||
"model_variant": {
|
||||
"type": "string",
|
||||
"description": "Provider route variant, e.g. standard or developer.",
|
||||
},
|
||||
"mode": {
|
||||
"type": "string",
|
||||
"description": "Provider-specific quality mode passed through when supported.",
|
||||
@@ -449,6 +483,16 @@ class VideoSelector(BaseTool):
|
||||
inputs: dict[str, object],
|
||||
candidates: list[BaseTool],
|
||||
) -> list[BaseTool]:
|
||||
exact_model = inputs.get("model")
|
||||
if exact_model:
|
||||
model_matches = [
|
||||
tool for tool in candidates
|
||||
if exact_model in getattr(tool, "input_schema", {}).get("properties", {}).get("model", {}).get("enum", [])
|
||||
or exact_model in tool.get_info().get("model_catalog", {})
|
||||
]
|
||||
if model_matches:
|
||||
candidates = model_matches
|
||||
|
||||
# A caller-supplied custom workflow is provider-specific (ComfyUI graph
|
||||
# JSON). Route it only to custom-workflow-capable providers whose server
|
||||
# is reachable — bundled-model readiness is irrelevant in that case.
|
||||
|
||||
Reference in New Issue
Block a user