diff --git a/.env.example b/.env.example index e22ca63..7fe11da 100644 --- a/.env.example +++ b/.env.example @@ -17,3 +17,8 @@ # Example of a credential a live/network-gated test would need. Tests that # require one must skip cleanly when it is unset. # EXAMPLE_API_KEY=replace-me + +# OpenRouter API key for the live media examples +# (`live_openrouter_image`, `live_openrouter_video`). They skip when unset. +# Live runs are billed to this key. +# OPENROUTER_API_KEY=sk-or-replace-me diff --git a/CHANGELOG.md b/CHANGELOG.md index b51be38..d647661 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,20 @@ # Changelog +## Unreleased + +### Added + +- `tinyinference-image`: the `ImageGenerator` trait, `OpenRouterImageGenerator` + (`POST /images`), `MockImageGenerator`, media-reference standards + (URL, `data:` URL, bytes, local path → OpenRouter content parts), aspect-ratio, + resolution and size normalization, per-model capability pre-flight checks, and + a billing-aware OpenRouter media transport usable directly or through a + proxying backend. +- `tinyinference-video`: the `VideoGenerator` trait, `OpenRouterVideoGenerator` + (`POST /videos`, `GET /videos/{id}`, `GET /videos/{id}/content`), + `wait_for_job` (resume by job id), and `MockVideoGenerator`. A `completed` + job with no outputs keeps polling instead of failing. + ## 0.3.0 ### Breaking changes diff --git a/Cargo.lock b/Cargo.lock index da64b9b..2e78e3f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1352,6 +1352,24 @@ dependencies = [ "url", ] +[[package]] +name = "tinyinference-image" +version = "0.3.0" +dependencies = [ + "async-trait", + "axum", + "base64 0.23.1", + "bytes", + "reqwest", + "serde", + "serde_json", + "tempfile", + "thiserror", + "tinyinference-core", + "tokio", + "tracing", +] + [[package]] name = "tinyinference-llm" version = "0.3.0" @@ -1416,6 +1434,21 @@ dependencies = [ "url", ] +[[package]] +name = "tinyinference-video" +version = "0.3.0" +dependencies = [ + "async-trait", + "axum", + "serde", + "serde_json", + "thiserror", + "tinyinference-image", + "tokio", + "tokio-test", + "tracing", +] + [[package]] name = "tinyinference-voice" version = "0.3.0" @@ -1515,6 +1548,28 @@ dependencies = [ "tokio", ] +[[package]] +name = "tokio-stream" +version = "0.1.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b" +dependencies = [ + "futures-core", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "tokio-test" +version = "0.4.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f6d24790a10a7af737693a3e8f1d03faef7e6ca0cc99aae5066f533766de545" +dependencies = [ + "futures-core", + "tokio", + "tokio-stream", +] + [[package]] name = "tokio-util" version = "0.7.19" diff --git a/README.md b/README.md index 786205d..3c9bcb5 100644 --- a/README.md +++ b/README.md @@ -50,8 +50,11 @@ can depend on `tinyinference-llm` for language models, `tinyinference-embeddings` for vector generation and retrieval, `tinyinference-local` for local runtimes and installers, `tinyinference-providers` for provider authentication and routing primitives, -`tinyinference-voice` for speech inference and streaming-audio mechanics, and -`tinyinference-core` only for shared infrastructure. +`tinyinference-voice` for speech inference and streaming-audio mechanics, +`tinyinference-image` for image generation and the shared media-reference +standards and OpenRouter media transport, +`tinyinference-video` for asynchronous video generation (submit, poll, +download, resume), and `tinyinference-core` only for shared infrastructure. ## Layout @@ -80,8 +83,26 @@ crates/tinyinference-providers/ └── src/ OAuth/PKCE flows and provider error classification crates/tinyinference-voice/ └── src/ hosted STT, Piper TTS, cleanup, and PCM streaming helpers +crates/tinyinference-image/ +└── src/ ImageGenerator, media references and output-shape + normalization, OpenRouter media transport, capabilities +crates/tinyinference-video/ +└── src/ VideoGenerator, submit/poll/download job loop, resume by id ``` +### Media generation + +`tinyinference-image` and `tinyinference-video` speak OpenRouter's media wire +format (`POST /images`, `POST /videos`, `GET /videos/{id}`, +`GET /videos/{id}/content`). The same generators run against OpenRouter +directly (`MediaAuth::ApiKey`) or against a host backend that proxies those +routes verbatim (`MediaAuth::Bearer` with `MediaTransport::with_base_url`). +A generator returns delivered media or an error — never an empty success — +and every error after a billed submit names the job and says not to resubmit. +Live smoke tests: `cargo run -p tinyinference-image --example +live_openrouter_image` and `cargo run -p tinyinference-video --example +live_openrouter_video` (skip without `OPENROUTER_API_KEY`). + ## Development ```sh diff --git a/crates/tinyinference-image/Cargo.toml b/crates/tinyinference-image/Cargo.toml new file mode 100644 index 0000000..18c8e41 --- /dev/null +++ b/crates/tinyinference-image/Cargo.toml @@ -0,0 +1,32 @@ +[package] +name = "tinyinference-image" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +repository.workspace = true +description = "Provider-neutral image generation, media references, and OpenRouter transport for Rust." +documentation = "https://docs.rs/tinyinference-image" +readme = "../../README.md" +keywords = ["image-generation", "inference", "openrouter", "media"] +categories = ["api-bindings", "asynchronous", "multimedia::images"] + +[dependencies] +async-trait = { workspace = true } +base64 = "0.23" +bytes = { workspace = true } +reqwest = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } +thiserror = { workspace = true } +tinyinference-core = { version = "0.3.0", path = "../tinyinference-core" } +tokio = { workspace = true, features = ["fs"] } +tracing = { workspace = true } + +[dev-dependencies] +axum = { workspace = true } +tempfile = { workspace = true } +tokio = { workspace = true, features = ["fs", "net"] } + +[lints] +workspace = true diff --git a/crates/tinyinference-image/examples/live_openrouter_image.rs b/crates/tinyinference-image/examples/live_openrouter_image.rs new file mode 100644 index 0000000..a6e00e7 --- /dev/null +++ b/crates/tinyinference-image/examples/live_openrouter_image.rs @@ -0,0 +1,95 @@ +//! Live smoke test: generate an image through OpenRouter and save it. +//! +//! Network- and credential-gated. Reads `OPENROUTER_API_KEY` from the +//! environment, falling back to the workspace `.env` file; exits cleanly +//! (status 0, "skipped") when neither provides one. +//! +//! ```sh +//! cargo run -p tinyinference-image --example live_openrouter_image +//! # image-to-image: pass a reference image (path or URL) +//! LIVE_REFERENCE=path/to/ref.png cargo run -p tinyinference-image --example live_openrouter_image +//! # optional overrides +//! LIVE_IMAGE_MODEL=google/gemini-3.1-flash-lite-image LIVE_PROMPT="…" cargo run … +//! ``` +//! +//! Output lands in `target/live-media/`. + +use std::path::PathBuf; +use std::time::Instant; + +use tinyinference_image::{ + ImageGenerator, ImageRequest, MediaAuth, MediaReference, OpenRouterImageGenerator, +}; + +fn api_key() -> Option { + if let Ok(key) = std::env::var("OPENROUTER_API_KEY") + && !key.trim().is_empty() + { + return Some(key.trim().to_owned()); + } + let env_file = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../../.env"); + std::fs::read_to_string(env_file) + .ok()? + .lines() + .find_map(|line| { + let value = line.trim().strip_prefix("OPENROUTER_API_KEY=")?; + let value = value.trim().trim_matches('"').trim_matches('\''); + (!value.is_empty()).then(|| value.to_owned()) + }) +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let Some(key) = api_key() else { + println!("skipped: OPENROUTER_API_KEY is not set (env or workspace .env)"); + return Ok(()); + }; + let generator = OpenRouterImageGenerator::new(MediaAuth::ApiKey(key)); + let model = + std::env::var("LIVE_IMAGE_MODEL").unwrap_or_else(|_| generator.default_model().to_owned()); + let prompt = std::env::var("LIVE_PROMPT").unwrap_or_else(|_| { + "A four-panel anime comic of two cheerful engineers shaking hands in front of a glowing \ + server, bright colors, thick ink outlines" + .to_owned() + }); + + let mut request = ImageRequest::new(prompt) + .with_model(&model) + .with_aspect_ratio("landscape") + .with_seed(42); + if let Ok(reference) = std::env::var("LIVE_REFERENCE") { + println!("image-to-image with reference {reference}"); + request = request.with_reference(MediaReference::parse(&reference)); + } + + let started = Instant::now(); + let response = generator.generate(request).await?; + let out_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../../target/live-media"); + for (index, image) in response.images.iter().enumerate() { + let path = image + .persist( + &out_dir, + &format!( + "image-{}-{}-{index}", + model.replace('/', "_"), + response.created.unwrap_or_default() + ), + "png", + ) + .await?; + println!( + "saved {} ({} bytes, {})", + path.display(), + image.data.len(), + image.media_type + ); + } + println!( + "model={} images={} cost_usd={:?} elapsed={:.1}s", + response.model, + response.images.len(), + response.cost_usd, + started.elapsed().as_secs_f64() + ); + Ok(()) +} diff --git a/crates/tinyinference-image/src/capabilities.rs b/crates/tinyinference-image/src/capabilities.rs new file mode 100644 index 0000000..2664bef --- /dev/null +++ b/crates/tinyinference-image/src/capabilities.rs @@ -0,0 +1,176 @@ +//! Per-model capability records and pre-flight request validation. +//! +//! Media generation is billed on submit, so a request the model cannot honor +//! should fail locally before it costs anything. Capabilities come from the +//! provider's model listings (`GET /images/models`, `GET /videos/models`). +//! Every field is optional: `None` means "not advertised", and an unadvertised +//! field is never used to reject a request. + +use serde_json::Value; + +use crate::{Error, Result}; + +/// What a model advertises it accepts. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct ModelCapabilities { + /// Accepted resolution tiers (`1K`, `720p`, …). + pub resolutions: Option>, + /// Accepted aspect ratios (`16:9`, …). + pub aspect_ratios: Option>, + /// Accepted clip durations in seconds (video only). + pub durations: Option>, + /// Accepted frame-image roles (`first_frame`, `last_frame`; video only). + pub frame_images: Option>, + /// Inclusive range of images per request. + pub n_range: Option<(u32, u32)>, + /// Maximum number of reference assets. + pub max_references: Option, + /// Whether a deterministic seed is accepted. + pub seed: Option, + /// Whether audio generation is available (video only). + pub generate_audio: Option, +} + +impl ModelCapabilities { + /// Reads an image-model record from `GET /images/models`. + /// + /// The record's `supported_parameters` map uses typed descriptors + /// (`{"type":"enum","values":[…]}`, `{"type":"range","min":…,"max":…}`, + /// `{"type":"boolean"}`); a key absent from a present map means the + /// parameter is unsupported. + #[must_use] + pub fn from_image_model(record: &Value) -> Self { + let Some(params) = record + .get("supported_parameters") + .and_then(Value::as_object) + else { + return Self::default(); + }; + // OpenRouter's contract: within a present `supported_parameters` map, + // an absent key means the endpoint does not support that parameter. + // So an omitted enum is "supports nothing" (`Some(vec![])`), which + // rejects the field before a billed call instead of letting the + // provider silently ignore it. + let enum_values = |key: &str| { + Some( + params + .get(key) + .and_then(|descriptor| descriptor.get("values")) + .and_then(Value::as_array) + .map(|values| string_list(values)) + .unwrap_or_default(), + ) + }; + let range = |key: &str| { + params.get(key).map(|descriptor| { + let bound = |name: &str| { + descriptor + .get(name) + .and_then(Value::as_u64) + .and_then(|n| u32::try_from(n).ok()) + }; + (bound("min").unwrap_or(0), bound("max").unwrap_or(u32::MAX)) + }) + }; + Self { + resolutions: enum_values("resolution"), + aspect_ratios: enum_values("aspect_ratio"), + durations: None, + frame_images: None, + n_range: Some(range("n").unwrap_or((1, 1))), + max_references: Some(range("input_references").map_or(0, |(_, max)| max)), + seed: Some(params.contains_key("seed")), + generate_audio: None, + } + } + + /// Reads a video-model record from `GET /videos/models` + /// (`supported_resolutions`, `supported_aspect_ratios`, + /// `supported_durations`, `supported_frame_images`, `seed`, + /// `generate_audio`). A `null` list means "not advertised". + #[must_use] + pub fn from_video_model(record: &Value) -> Self { + let list = |key: &str| { + record + .get(key) + .and_then(Value::as_array) + .map(|values| string_list(values)) + }; + Self { + resolutions: list("supported_resolutions"), + aspect_ratios: list("supported_aspect_ratios"), + durations: record + .get("supported_durations") + .and_then(Value::as_array) + .map(|values| { + values + .iter() + .filter_map(Value::as_u64) + .filter_map(|n| u32::try_from(n).ok()) + .collect() + }), + frame_images: list("supported_frame_images"), + n_range: None, + max_references: None, + seed: record.get("seed").and_then(Value::as_bool), + generate_audio: record.get("generate_audio").and_then(Value::as_bool), + } + } + + /// Rejects `value` for `field` when the model advertises a list that does + /// not contain it. + /// + /// # Errors + /// + /// [`Error::Unsupported`] naming the field and the advertised values. + pub fn check_one_of( + model: &str, + field: &str, + value: &str, + allowed: Option<&[String]>, + ) -> Result<()> { + let Some(allowed) = allowed else { + return Ok(()); + }; + if allowed + .iter() + .any(|candidate| candidate.eq_ignore_ascii_case(value)) + { + return Ok(()); + } + Err(Error::Unsupported { + model: model.to_owned(), + field: field.to_owned(), + value: value.to_owned(), + allowed: allowed.to_vec(), + }) + } + + /// Rejects a feature the model advertises as unavailable. + /// + /// # Errors + /// + /// [`Error::Unsupported`] when `supported` is `Some(false)`. + pub fn check_flag(model: &str, field: &str, supported: Option) -> Result<()> { + if supported == Some(false) { + return Err(Error::Unsupported { + model: model.to_owned(), + field: field.to_owned(), + value: "true".to_owned(), + allowed: Vec::new(), + }); + } + Ok(()) + } +} + +fn string_list(values: &[Value]) -> Vec { + values + .iter() + .filter_map(|value| match value { + Value::String(text) => Some(text.clone()), + Value::Number(number) => Some(number.to_string()), + _ => None, + }) + .collect() +} diff --git a/crates/tinyinference-image/src/error.rs b/crates/tinyinference-image/src/error.rs new file mode 100644 index 0000000..d6a7028 --- /dev/null +++ b/crates/tinyinference-image/src/error.rs @@ -0,0 +1,89 @@ +//! Error type shared by image generation and the OpenRouter media transport. + +use thiserror::Error; + +/// Result returned by TinyInference image APIs. +pub type Result = std::result::Result; + +/// A normalized media-generation failure. +/// +/// Variants are split by what a caller can do about them: fix the request +/// ([`Error::Validation`], [`Error::Unsupported`]), fix the credential +/// ([`Error::Auth`]), retry later ([`Error::Http`] with a retryable status, +/// [`Error::Transport`]), or report a billed non-delivery +/// ([`Error::NoMedia`]) without retrying. +#[derive(Debug, Error)] +pub enum Error { + /// Caller input or configuration was invalid before any request was sent. + #[error("validation error: {0}")] + Validation(String), + /// A request parameter is not supported by the selected model. + #[error("model '{model}' does not support {field}={value}; supported: {}", allowed.join(", "))] + Unsupported { + /// Model id the capability check ran against. + model: String, + /// Request field that failed the check (for example `aspect_ratio`). + field: String, + /// The rejected value. + value: String, + /// Values the model advertises for `field` (empty when unsupported). + allowed: Vec, + }, + /// No usable credential was available, or the provider rejected it. + #[error("authentication error: {0}")] + Auth(String), + /// The provider answered with a non-success HTTP status. + #[error("provider returned HTTP {status}: {message}")] + Http { + /// HTTP status code. + status: u16, + /// Sanitized provider error message. + message: String, + }, + /// The request never produced an HTTP response (DNS, TLS, timeout, reset). + #[error("transport error: {0}")] + Transport(String), + /// A provider payload or media body could not be decoded. + #[error("decode error: {0}")] + Decode(String), + /// The provider accepted (and billed) the request but returned no media. + /// + /// Retrying submits and bills a new generation, so callers should report + /// this rather than retry automatically. + #[error( + "generation was accepted and billed but returned no media{}; do not retry automatically — report this to the user", + request_id.as_deref().map(|id| format!(" (request_id: {id})")).unwrap_or_default() + )] + NoMedia { + /// Provider request or job id, when one was returned. + request_id: Option, + }, + /// A media body exceeded the configured size cap. + #[error("media body exceeds the {limit}-byte limit")] + TooLarge { + /// The cap that was exceeded, in bytes. + limit: usize, + }, + /// Reading a local reference or writing an artifact failed. + #[error("io error: {0}")] + Io(#[from] std::io::Error), + /// A JSON payload could not be encoded or decoded. + #[error("serialization error: {0}")] + Serialization(#[from] serde_json::Error), +} + +impl Error { + /// Whether the failure is transient and the same call may succeed later. + /// + /// Only rate limits, upstream 5xx responses and transport failures are + /// retryable. A billed non-delivery ([`Error::NoMedia`]) is deliberately + /// not: a retry is a new, separately billed generation. + #[must_use] + pub fn is_retryable(&self) -> bool { + match self { + Self::Http { status, .. } => *status == 429 || (500..=599).contains(status), + Self::Transport(_) => true, + _ => false, + } + } +} diff --git a/crates/tinyinference-image/src/lib.rs b/crates/tinyinference-image/src/lib.rs new file mode 100644 index 0000000..efd9fe5 --- /dev/null +++ b/crates/tinyinference-image/src/lib.rs @@ -0,0 +1,93 @@ +//! Provider-neutral image generation for TinyInference. +//! +//! This crate owns three things: +//! +//! - **Media standards** ([`reference`](mod@crate::reference)) — how a reference asset is described +//! (URL, `data:` URL, bytes, local path) and inlined, and how loose +//! output-shape spellings (`"16x9"`, `"landscape"`, `"full hd"`) normalize to +//! canonical wire values. The video crate reuses these, so image and video +//! generation agree on one vocabulary. +//! - **The OpenRouter media transport** ([`transport`]) — authenticated HTTP +//! with billing-aware retries, shared by image and video generation. The same +//! transport serves OpenRouter directly (API key) and any host backend that +//! proxies OpenRouter's media routes verbatim (host-resolved bearer). +//! - **Image generation** — the [`ImageGenerator`] trait, +//! [`OpenRouterImageGenerator`], and [`MockImageGenerator`] for offline +//! tests. +//! +//! A generator either returns at least one image or fails: an accepted request +//! that yields nothing is [`Error::NoMedia`], never an empty success, because +//! the request was already billed and a caller told "success" would report an +//! image that does not exist. +//! +//! # Example +//! ``` +//! use tinyinference_image::{ImageGenerator, ImageRequest, MockImageGenerator}; +//! +//! # tokio::runtime::Runtime::new().unwrap().block_on(async { +//! let generator = MockImageGenerator::new(); +//! let response = generator +//! .generate(ImageRequest::new("a red panda astronaut").with_aspect_ratio("landscape")) +//! .await +//! .unwrap(); +//! assert_eq!(response.images.len(), 1); +//! # }); +//! ``` + +pub mod capabilities; +mod error; +pub mod media; +mod mock; +pub mod openrouter; +pub mod reference; +pub mod transport; +mod types; + +pub use capabilities::ModelCapabilities; +pub use error::{Error, Result}; +pub use media::GeneratedMedia; +pub use mock::MockImageGenerator; +pub use openrouter::{DEFAULT_IMAGE_MODEL, OpenRouterImageGenerator}; +pub use reference::{MediaReference, ReferenceKind}; +pub use transport::{BearerResolver, MediaAuth, MediaTransport}; +pub use types::{ImageRequest, ImageResponse, MAX_IMAGES_PER_REQUEST, MediaModel}; + +use async_trait::async_trait; + +/// An image generation provider. +/// +/// Implementations must be `Send + Sync` so hosts can share one generator +/// across concurrent tool calls. +#[async_trait] +pub trait ImageGenerator: Send + Sync { + /// Short provider name for logs and diagnostics (`"openrouter"`). + fn name(&self) -> &str; + + /// Model used when a request names none. + fn default_model(&self) -> &str; + + /// Generates one or more images. + /// + /// # Errors + /// + /// [`Error::Validation`] / [`Error::Unsupported`] before any request is + /// sent; [`Error::Auth`], [`Error::Http`], [`Error::Transport`] from the + /// provider; [`Error::NoMedia`] when the provider accepted and billed the + /// request but returned no images. + async fn generate(&self, request: ImageRequest) -> Result; + + /// Lists the models this provider can generate with. + /// + /// # Errors + /// + /// Provider or decode errors from the listing endpoint. + async fn list_models(&self) -> Result>; +} + +#[cfg(test)] +#[path = "reference_test.rs"] +mod reference_test; + +#[cfg(test)] +#[path = "openrouter_test.rs"] +mod openrouter_test; diff --git a/crates/tinyinference-image/src/media.rs b/crates/tinyinference-image/src/media.rs new file mode 100644 index 0000000..a217945 --- /dev/null +++ b/crates/tinyinference-image/src/media.rs @@ -0,0 +1,115 @@ +//! Generated media artifacts and their persistence. + +use std::path::{Path, PathBuf}; + +use bytes::Bytes; + +use crate::reference::extension_for_media_type; +use crate::{Error, Result}; + +/// One generated artifact, held in memory. +/// +/// Providers return media either inline (OpenRouter images arrive as base64) +/// or behind an authenticated content endpoint (OpenRouter videos); generators +/// always download before returning, so a `GeneratedMedia` is self-contained +/// and never carries a URL that expires or needs a credential. +#[derive(Clone, PartialEq, Eq)] +pub struct GeneratedMedia { + /// MIME type, for example `image/png` or `video/mp4`. + pub media_type: String, + /// The artifact bytes. + pub data: Bytes, +} + +impl GeneratedMedia { + /// Creates an artifact from its media type and bytes. + #[must_use] + pub fn new(media_type: impl Into, data: impl Into) -> Self { + Self { + media_type: media_type.into(), + data: data.into(), + } + } + + /// File extension for this artifact's media type, or `fallback`. + #[must_use] + pub fn extension<'a>(&self, fallback: &'a str) -> &'a str { + extension_for_media_type(&self.media_type, fallback) + } + + /// Writes the artifact to `dir/.`, creating `dir` if needed, + /// and returns the written path. `fallback_extension` is used when the + /// media type is unknown. + /// + /// The stem is sanitized to `[A-Za-z0-9_-]` so a provider id can never + /// traverse out of `dir`. + /// + /// # Errors + /// + /// [`Error::Validation`] for an empty artifact, [`Error::Io`] when the + /// directory or file cannot be written. + pub async fn persist( + &self, + dir: &Path, + stem: &str, + fallback_extension: &str, + ) -> Result { + if self.data.is_empty() { + return Err(Error::Validation( + "refusing to persist an empty artifact".into(), + )); + } + tokio::fs::create_dir_all(dir).await?; + let stem: String = stem + .chars() + .map(|c| { + if c.is_ascii_alphanumeric() || c == '-' || c == '_' { + c + } else { + '_' + } + }) + .collect(); + let stem = if stem.is_empty() { + "media".to_owned() + } else { + stem + }; + // Sanitize the fallback extension to a safe filename component. + let safe_fallback: String = fallback_extension + .chars() + .map(|c| { + if c.is_ascii_alphanumeric() || c == '-' || c == '_' { + c + } else { + '_' + } + }) + .collect(); + let safe_fallback = if safe_fallback.is_empty() { + "bin" + } else { + &safe_fallback + }; + let extension = extension_for_media_type(&self.media_type, safe_fallback); + let path = dir.join(format!("{stem}.{extension}")); + tokio::fs::write(&path, &self.data).await?; + tracing::debug!( + path = %path.display(), + bytes = self.data.len(), + media_type = %self.media_type, + "[tinyinference-image] persisted generated media" + ); + Ok(path) + } +} + +impl std::fmt::Debug for GeneratedMedia { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("GeneratedMedia") + .field("media_type", &self.media_type) + .field("len", &self.data.len()) + .finish() + } +} diff --git a/crates/tinyinference-image/src/mock.rs b/crates/tinyinference-image/src/mock.rs new file mode 100644 index 0000000..e113821 --- /dev/null +++ b/crates/tinyinference-image/src/mock.rs @@ -0,0 +1,100 @@ +//! Deterministic, offline image generator for tests. + +use std::sync::Mutex; + +use async_trait::async_trait; + +use crate::media::GeneratedMedia; +use crate::types::{ImageRequest, ImageResponse, MediaModel}; +use crate::{Error, ImageGenerator, ModelCapabilities, Result}; + +/// The smallest valid PNG (1×1, transparent), returned by the mock. +pub(crate) const TINY_PNG: &[u8] = &[ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52, + 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x06, 0x00, 0x00, 0x00, 0x1F, 0x15, 0xC4, + 0x89, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x44, 0x41, 0x54, 0x78, 0x9C, 0x63, 0x00, 0x01, 0x00, 0x00, + 0x05, 0x00, 0x01, 0x0D, 0x0A, 0x2D, 0xB4, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45, 0x4E, 0x44, 0xAE, + 0x42, 0x60, 0x82, +]; + +/// Returns one tiny PNG per requested image, records every request, and can +/// be told to simulate a billed non-delivery. +#[derive(Debug, Default)] +pub struct MockImageGenerator { + requests: Mutex>, + return_no_media: bool, +} + +impl MockImageGenerator { + /// Creates a mock that always succeeds. + #[must_use] + pub fn new() -> Self { + Self::default() + } + + /// Creates a mock whose every call fails with [`Error::NoMedia`]. + #[must_use] + pub fn returning_no_media() -> Self { + Self { + return_no_media: true, + ..Self::default() + } + } + + /// Requests received so far, in order. + /// + /// # Panics + /// + /// If a previous holder of the internal lock panicked. + #[must_use] + pub fn requests(&self) -> Vec { + self.requests.lock().expect("mock lock poisoned").clone() + } +} + +#[async_trait] +impl ImageGenerator for MockImageGenerator { + fn name(&self) -> &str { + "mock" + } + + fn default_model(&self) -> &str { + "mock/image" + } + + async fn generate(&self, request: ImageRequest) -> Result { + request.validate()?; + let n = request.n.unwrap_or(1); + let model = request + .model + .clone() + .unwrap_or_else(|| self.default_model().to_owned()); + self.requests + .lock() + .expect("mock lock poisoned") + .push(request); + if self.return_no_media { + return Err(Error::NoMedia { + request_id: Some("mock-request".into()), + }); + } + Ok(ImageResponse { + model, + images: (0..n) + .map(|_| GeneratedMedia::new("image/png", TINY_PNG)) + .collect(), + cost_usd: Some(0.0), + created: None, + }) + } + + async fn list_models(&self) -> Result> { + Ok(vec![MediaModel { + id: self.default_model().to_owned(), + name: Some("Mock image".into()), + description: None, + capabilities: ModelCapabilities::default(), + raw: serde_json::Value::Null, + }]) + } +} diff --git a/crates/tinyinference-image/src/openrouter.rs b/crates/tinyinference-image/src/openrouter.rs new file mode 100644 index 0000000..d9eac69 --- /dev/null +++ b/crates/tinyinference-image/src/openrouter.rs @@ -0,0 +1,385 @@ +//! OpenRouter image generation (`POST /images`). + +use std::collections::HashMap; + +use async_trait::async_trait; +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD as BASE64; +use serde::Deserialize; +use serde_json::{Map, Value, json}; +use tokio::sync::Mutex; + +use crate::capabilities::ModelCapabilities; +use crate::media::GeneratedMedia; +use crate::reference::{ + DEFAULT_MAX_REFERENCE_BYTES, normalize_aspect_ratio, normalize_image_resolution, normalize_size, +}; +use crate::transport::{MediaAuth, MediaTransport, wire_model_id}; +use crate::types::{ImageRequest, ImageResponse, MediaModel}; +use crate::{Error, ImageGenerator, Result}; + +/// Default image model: Seedream 5.0 Lite (flat per-image price, +/// image-to-image with up to 14 references, deterministic seed). +pub const DEFAULT_IMAGE_MODEL: &str = "bytedance-seed/seedream-5-0-lite"; + +/// Image generator for OpenRouter's image API, or any backend that proxies it +/// verbatim. +#[derive(Debug)] +pub struct OpenRouterImageGenerator { + transport: MediaTransport, + default_model: String, + check_capabilities: bool, + max_reference_bytes: usize, + // Some(None) = listing unavailable; skip checks without refetching. + capabilities: Mutex>>>, +} + +#[derive(Deserialize)] +struct WireResponse { + #[serde(default)] + created: Option, + #[serde(default)] + data: Vec, + #[serde(default)] + usage: Option, +} + +#[derive(Deserialize)] +struct WireImage { + #[serde(default)] + b64_json: Option, + #[serde(default)] + media_type: Option, +} + +#[derive(Deserialize)] +struct WireUsage { + #[serde(default)] + cost: Option, +} + +impl OpenRouterImageGenerator { + /// Creates a generator against OpenRouter's public API. + #[must_use] + pub fn new(auth: MediaAuth) -> Self { + Self::with_transport(MediaTransport::new(auth)) + } + + /// Creates a generator from a pre-configured transport (base URL, client, + /// headers, retry and size limits). + #[must_use] + pub fn with_transport(transport: MediaTransport) -> Self { + Self { + transport, + default_model: DEFAULT_IMAGE_MODEL.to_owned(), + check_capabilities: true, + max_reference_bytes: DEFAULT_MAX_REFERENCE_BYTES, + capabilities: Mutex::new(None), + } + } + + /// Creates a generator from `OPENROUTER_API_KEY`. + /// + /// # Errors + /// + /// [`Error::Auth`] when the variable is unset. + pub fn from_env() -> Result { + Ok(Self::new(MediaAuth::from_env()?)) + } + + /// Sets the model used when a request names none. + #[must_use] + pub fn with_default_model(mut self, model: impl Into) -> Self { + self.default_model = model.into(); + self + } + + /// Enables or disables pre-flight validation against the model listing + /// (enabled by default; skipped silently when the listing is unavailable). + #[must_use] + pub fn with_capability_check(mut self, enabled: bool) -> Self { + self.check_capabilities = enabled; + self + } + + /// Caps the size of each inlined reference (default 20 MiB). + #[must_use] + pub fn with_max_reference_bytes(mut self, max_reference_bytes: usize) -> Self { + self.max_reference_bytes = max_reference_bytes; + self + } + + /// The underlying transport. + #[must_use] + pub fn transport(&self) -> &MediaTransport { + &self.transport + } + + async fn capabilities_for(&self, model: &str) -> Option { + let mut cache = self.capabilities.lock().await; + if cache.is_none() { + match self.list_models().await { + Ok(models) => { + *cache = Some(Some( + models + .into_iter() + .map(|model| (wire_model_id(&model.id).to_owned(), model.capabilities)) + .collect(), + )); + } + Err(error) => { + tracing::debug!( + %error, + "[tinyinference-image] model listing unavailable; skipping capability check" + ); + *cache = Some(None); + return None; + } + } + } + cache.as_ref()?.as_ref()?.get(model).cloned() + } + + fn validate_against(model: &str, request: &WireFields, caps: &ModelCapabilities) -> Result<()> { + if let Some(value) = &request.aspect_ratio + && value != "auto" + { + ModelCapabilities::check_one_of( + model, + "aspect_ratio", + value, + caps.aspect_ratios.as_deref(), + )?; + } + if let Some(value) = &request.resolution { + ModelCapabilities::check_one_of( + model, + "resolution", + value, + caps.resolutions.as_deref(), + )?; + } + if let (Some(n), Some((min, max))) = (request.n, caps.n_range) + && n > 1 + && !(min..=max).contains(&n) + { + return Err(Error::Unsupported { + model: model.to_owned(), + field: "n".into(), + value: n.to_string(), + allowed: vec![format!("{min}..={max}")], + }); + } + if let Some(max) = caps.max_references + && request.references > max as usize + { + return Err(Error::Unsupported { + model: model.to_owned(), + field: "input_references".into(), + value: request.references.to_string(), + allowed: vec![format!("at most {max}")], + }); + } + if request.seed { + ModelCapabilities::check_flag(model, "seed", caps.seed)?; + } + Ok(()) + } +} + +/// Normalized fields the capability check reads. +struct WireFields { + aspect_ratio: Option, + resolution: Option, + n: Option, + references: usize, + seed: bool, +} + +/// Builds the wire body for `request`. Exposed for tests and for hosts that +/// need to audit exactly what leaves the process. +/// +/// # Errors +/// +/// Reference resolution errors from [`crate::MediaReference::resolve`]. +pub async fn build_image_body( + model: &str, + request: &ImageRequest, + max_reference_bytes: usize, +) -> Result { + let mut body = Map::new(); + body.insert("model".into(), json!(model)); + body.insert("prompt".into(), json!(request.prompt)); + if let Some(n) = request.n { + body.insert("n".into(), json!(n)); + } + let normalized = |value: &Option, normalize: fn(&str) -> Option| { + value + .as_deref() + .map(|raw| normalize(raw).unwrap_or_else(|| raw.trim().to_owned())) + }; + if let Some(size) = normalized(&request.size, normalize_size) { + body.insert("size".into(), json!(size)); + } + if let Some(resolution) = normalized(&request.resolution, normalize_image_resolution) { + body.insert("resolution".into(), json!(resolution)); + } + if let Some(aspect_ratio) = normalized(&request.aspect_ratio, normalize_aspect_ratio) { + body.insert("aspect_ratio".into(), json!(aspect_ratio)); + } + for (key, value) in [ + ("quality", &request.quality), + ("output_format", &request.output_format), + ("background", &request.background), + ("user", &request.user), + ("session_id", &request.session_id), + ] { + if let Some(value) = value { + body.insert(key.into(), json!(value.trim())); + } + } + if let Some(seed) = request.seed { + body.insert("seed".into(), json!(seed)); + } + if !request.references.is_empty() { + let mut parts = Vec::with_capacity(request.references.len()); + for reference in &request.references { + parts.push(reference.to_content_part(max_reference_bytes).await?); + } + body.insert("input_references".into(), Value::Array(parts)); + } + for (key, value) in &request.extra { + body.entry(key.clone()).or_insert_with(|| value.clone()); + } + Ok(Value::Object(body)) +} + +/// Sniffs a media type from magic bytes, for responses that omit it. +fn sniff_image_type(data: &[u8]) -> &'static str { + match data { + [0x89, b'P', b'N', b'G', ..] => "image/png", + [0xFF, 0xD8, 0xFF, ..] => "image/jpeg", + [ + b'R', + b'I', + b'F', + b'F', + _, + _, + _, + _, + b'W', + b'E', + b'B', + b'P', + .., + ] => "image/webp", + [b'G', b'I', b'F', b'8', ..] => "image/gif", + [b'<', ..] => "image/svg+xml", + _ => "image/png", + } +} + +#[async_trait] +impl ImageGenerator for OpenRouterImageGenerator { + fn name(&self) -> &str { + "openrouter" + } + + fn default_model(&self) -> &str { + &self.default_model + } + + async fn generate(&self, request: ImageRequest) -> Result { + request.validate()?; + let model = + wire_model_id(request.model.as_deref().unwrap_or(&self.default_model)).to_owned(); + let body = build_image_body(&model, &request, self.max_reference_bytes).await?; + + if self.check_capabilities + && let Some(caps) = self.capabilities_for(&model).await + { + let field = |key: &str| body.get(key).and_then(Value::as_str).map(str::to_owned); + Self::validate_against( + &model, + &WireFields { + aspect_ratio: field("aspect_ratio"), + resolution: field("resolution"), + n: request.n, + references: request.references.len(), + seed: request.seed.is_some(), + }, + &caps, + )?; + } + + tracing::info!( + model = %model, + n = request.n.unwrap_or(1), + references = request.references.len(), + base_url = %self.transport.redacted_base_url(), + "[tinyinference-image] generating image" + ); + let response: WireResponse = self.transport.post_json("images", &body).await?; + let cost_usd = response.usage.as_ref().and_then(|usage| usage.cost); + + let mut images = Vec::with_capacity(response.data.len()); + for (index, image) in response.data.into_iter().enumerate() { + let Some(encoded) = image.b64_json.filter(|value| !value.is_empty()) else { + tracing::warn!(index, "[tinyinference-image] image entry without b64_json"); + continue; + }; + if encoded.len() / 4 * 3 > self.transport.max_media_bytes() { + return Err(Error::TooLarge { + limit: self.transport.max_media_bytes(), + }); + } + // The request was already billed: an undecodable entry is a + // non-delivery for that image, not a transient decode error the + // caller might retry. Skip it; if none decode, `NoMedia` below. + let data = match BASE64.decode(encoded.as_bytes()) { + Ok(data) if !data.is_empty() => data, + Ok(_) | Err(_) => { + tracing::warn!( + index, + "[tinyinference-image] undecodable image entry skipped" + ); + continue; + } + }; + let media_type = image + .media_type + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| sniff_image_type(&data).to_owned()); + images.push(GeneratedMedia::new(media_type, data)); + } + if images.is_empty() { + tracing::warn!( + model = %model, + cost_usd, + "[tinyinference-image] provider accepted the request but returned no images" + ); + return Err(Error::NoMedia { request_id: None }); + } + tracing::info!( + model = %model, + images = images.len(), + cost_usd, + "[tinyinference-image] image generation complete" + ); + Ok(ImageResponse { + model, + images, + cost_usd, + created: response.created, + }) + } + + async fn list_models(&self) -> Result> { + let body: Value = self.transport.get_json("images/models").await?; + Ok(MediaModel::parse_listing( + &body, + ModelCapabilities::from_image_model, + )) + } +} diff --git a/crates/tinyinference-image/src/openrouter_test.rs b/crates/tinyinference-image/src/openrouter_test.rs new file mode 100644 index 0000000..c2d4a13 --- /dev/null +++ b/crates/tinyinference-image/src/openrouter_test.rs @@ -0,0 +1,450 @@ +//! Offline tests for the OpenRouter image generator and media transport. + +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +use axum::Router; +use axum::extract::{Request, State}; +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD as BASE64; +use bytes::Bytes; +use serde_json::{Value, json}; + +use crate::mock::TINY_PNG; +use crate::{ + Error, ImageGenerator, ImageRequest, MediaAuth, MediaReference, MediaTransport, + OpenRouterImageGenerator, +}; + +const KEY: &str = "sk-or-test-secret-key-0123456789"; + +#[derive(Clone, Default)] +struct Captured { + bodies: Arc>>, + headers: Arc>>, + image_calls: Arc, +} + +struct Fixture { + base_url: String, + captured: Captured, +} + +async fn serve(router: Router) -> String { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + tokio::spawn(async move { + axum::serve(listener, router).await.unwrap(); + }); + format!("http://{address}") +} + +/// Serves `images` and `images/models` under `prefix` with canned handlers. +async fn fixture( + prefix: &str, + image_reply: fn(usize) -> Response, + listing: Option, +) -> Fixture { + let captured = Captured::default(); + let images_state = captured.clone(); + let images = post( + move |State(state): State, request: Request| async move { + let headers = request.headers().clone(); + let body = axum::body::to_bytes(request.into_body(), usize::MAX) + .await + .unwrap(); + state.headers.lock().unwrap().push(headers); + state + .bodies + .lock() + .unwrap() + .push(serde_json::from_slice(&body).unwrap()); + let call = state.image_calls.fetch_add(1, Ordering::SeqCst); + image_reply(call) + }, + ); + let models = get(move || async move { + match listing { + Some(listing) => axum::Json(listing).into_response(), + None => StatusCode::NOT_FOUND.into_response(), + } + }); + let router = Router::new() + .route(&format!("{prefix}/images"), images) + .route(&format!("{prefix}/images/models"), models) + .with_state(images_state); + Fixture { + base_url: format!("{}{prefix}", serve(router).await), + captured, + } +} + +fn png_reply(_call: usize) -> Response { + axum::Json(json!({ + "created": 1_748_372_400, + "data": [{ "b64_json": BASE64.encode(TINY_PNG), "media_type": "image/png" }], + "usage": { "cost": 0.035 } + })) + .into_response() +} + +fn generator(base_url: &str) -> OpenRouterImageGenerator { + OpenRouterImageGenerator::with_transport( + MediaTransport::new(MediaAuth::ApiKey(KEY.into())) + .with_base_url(base_url) + .with_header("x-title", "tinyinference-tests") + .with_max_retries(2), + ) +} + +#[tokio::test] +async fn generates_decodes_and_reports_cost() { + let fixture = fixture("/api/v1", png_reply, None).await; + let response = generator(&fixture.base_url) + .generate( + ImageRequest::new("a red panda astronaut") + .with_model("openrouter/bytedance-seed/seedream-5-0-lite") + .with_aspect_ratio("landscape") + .with_resolution("2k") + .with_seed(7) + .with_reference(MediaReference::Url("https://example.com/ref.png".into())) + .with_reference(MediaReference::Bytes { + media_type: "image/png".into(), + data: Bytes::from_static(TINY_PNG), + }), + ) + .await + .unwrap(); + + assert_eq!(response.model, "bytedance-seed/seedream-5-0-lite"); + assert_eq!(response.images.len(), 1); + assert_eq!(response.images[0].media_type, "image/png"); + assert_eq!(response.images[0].data.as_ref(), TINY_PNG); + assert_eq!(response.cost_usd, Some(0.035)); + + let body = fixture.captured.bodies.lock().unwrap()[0].clone(); + assert_eq!( + body["model"], "bytedance-seed/seedream-5-0-lite", + "openrouter/ prefix stripped" + ); + assert_eq!(body["aspect_ratio"], "16:9"); + assert_eq!(body["resolution"], "2K"); + assert_eq!(body["seed"], 7); + assert_eq!(body["input_references"][0]["type"], "image_url"); + assert_eq!( + body["input_references"][0]["image_url"]["url"], + "https://example.com/ref.png" + ); + assert!( + body["input_references"][1]["image_url"]["url"] + .as_str() + .unwrap() + .starts_with("data:image/png;base64,") + ); + + let headers = fixture.captured.headers.lock().unwrap()[0].clone(); + assert_eq!(headers["authorization"], format!("Bearer {KEY}")); + assert_eq!(headers["x-title"], "tinyinference-tests"); +} + +/// Regression (R1): an accepted, billed request that returns no images must be +/// an error that tells the caller not to retry — never an empty success. +#[tokio::test] +async fn accepted_request_without_images_is_no_media_error() { + fn empty(_call: usize) -> Response { + axum::Json(json!({ "created": 1, "data": [], "usage": { "cost": 0.035 } })).into_response() + } + let fixture = fixture("/api/v1", empty, None).await; + let error = generator(&fixture.base_url) + .generate(ImageRequest::new("anything")) + .await + .unwrap_err(); + assert!(matches!(error, Error::NoMedia { .. }), "{error:?}"); + assert!(error.to_string().contains("do not retry"), "{error}"); + assert!(!error.is_retryable()); + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn entries_without_payload_are_not_counted_as_images() { + fn blank(_call: usize) -> Response { + axum::Json(json!({ "created": 1, "data": [{ "b64_json": "" }] })).into_response() + } + let fixture = fixture("/api/v1", blank, None).await; + let error = generator(&fixture.base_url) + .generate(ImageRequest::new("anything")) + .await + .unwrap_err(); + assert!(matches!(error, Error::NoMedia { .. }), "{error:?}"); +} + +#[tokio::test] +async fn unsupported_aspect_ratio_fails_before_the_paid_call() { + let listing = json!({ "data": [{ + "id": "bytedance-seed/seedream-5-0-lite", + "supported_parameters": { + "aspect_ratio": { "type": "enum", "values": ["1:1", "3:4"] }, + "n": { "type": "range", "min": 1, "max": 4 } + } + }]}); + let fixture = fixture("/api/v1", png_reply, Some(listing)).await; + let error = generator(&fixture.base_url) + .generate(ImageRequest::new("x").with_aspect_ratio("16:9")) + .await + .unwrap_err(); + match error { + Error::Unsupported { field, allowed, .. } => { + assert_eq!(field, "aspect_ratio"); + assert_eq!(allowed, vec!["1:1", "3:4"]); + } + other => panic!("expected Unsupported, got {other:?}"), + } + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn unsupported_seed_is_rejected_when_the_listing_omits_it() { + let listing = json!({ "data": [{ + "id": "google/gemini-3.1-flash-lite-image", + "supported_parameters": { "aspect_ratio": { "type": "enum", "values": ["1:1"] } } + }]}); + let fixture = fixture("/api/v1", png_reply, Some(listing)).await; + let error = generator(&fixture.base_url) + .generate( + ImageRequest::new("x") + .with_model("google/gemini-3.1-flash-lite-image") + .with_seed(1), + ) + .await + .unwrap_err(); + assert!( + matches!(error, Error::Unsupported { ref field, .. } if field == "seed"), + "{error:?}" + ); +} + +#[tokio::test] +async fn listing_without_capabilities_does_not_block_generation() { + // A proxying backend may list ids only; that must never reject a request. + let listing = json!({ "object": "list", "data": [{ + "id": "bytedance-seed/seedream-5-0-lite", "display_name": "Seedream 5.0 Lite" + }]}); + let fixture = fixture("/api/v1", png_reply, Some(listing)).await; + generator(&fixture.base_url) + .generate( + ImageRequest::new("x") + .with_aspect_ratio("21:9") + .with_seed(3), + ) + .await + .unwrap(); +} + +/// A proxying backend wraps OpenRouter's body in `{success, data}`; the +/// transport unwraps it transparently. +fn enveloped_png_reply(_call: usize) -> Response { + axum::Json(json!({ "success": true, "data": { + "created": 1, + "data": [{ "b64_json": BASE64.encode(TINY_PNG), "media_type": "image/png" }], + "usage": { "cost": 0.035 } + }})) + .into_response() +} + +#[tokio::test] +async fn proxied_backend_base_url_and_bearer_resolver() { + let fixture = fixture("/agent-integrations/openrouter", enveloped_png_reply, None).await; + let resolver: crate::BearerResolver = Arc::new(|| Ok("session-jwt-abcdefgh".to_owned())); + let generator = OpenRouterImageGenerator::with_transport( + MediaTransport::new(MediaAuth::Bearer(resolver)).with_base_url(&fixture.base_url), + ); + let response = generator.generate(ImageRequest::new("x")).await.unwrap(); + assert_eq!(response.images.len(), 1); + assert_eq!(response.cost_usd, Some(0.035)); + let headers = fixture.captured.headers.lock().unwrap()[0].clone(); + assert_eq!(headers["authorization"], "Bearer session-jwt-abcdefgh"); +} + +#[tokio::test] +async fn failed_envelope_is_an_error_even_with_a_2xx_status() { + fn failed(_call: usize) -> Response { + axum::Json(json!({ "success": false, "error": "Insufficient balance" })).into_response() + } + let fixture = fixture("/agent-integrations/openrouter", failed, None).await; + let error = generator(&fixture.base_url) + .with_capability_check(false) + .generate(ImageRequest::new("x")) + .await + .unwrap_err(); + assert!( + error.to_string().contains("Insufficient balance"), + "{error}" + ); +} + +#[tokio::test] +async fn blank_bearer_fails_without_a_request() { + let fixture = fixture("/api/v1", png_reply, None).await; + let resolver: crate::BearerResolver = Arc::new(|| Ok(" ".to_owned())); + let generator = OpenRouterImageGenerator::with_transport( + MediaTransport::new(MediaAuth::Bearer(resolver)).with_base_url(&fixture.base_url), + ) + .with_capability_check(false); + let error = generator + .generate(ImageRequest::new("x")) + .await + .unwrap_err(); + assert!(matches!(error, Error::Auth(_)), "{error:?}"); + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 0); +} + +/// A 5xx after a submit may already have started (and billed) a generation, so +/// the billable POST must not be retried. +#[tokio::test] +async fn billable_post_is_not_retried_on_server_error() { + fn boom(_call: usize) -> Response { + ( + StatusCode::BAD_GATEWAY, + axum::Json(json!({"error": {"code": 502, "message": "upstream"}})), + ) + .into_response() + } + let fixture = fixture("/api/v1", boom, None).await; + let error = generator(&fixture.base_url) + .with_capability_check(false) + .generate(ImageRequest::new("x")) + .await + .unwrap_err(); + assert!( + matches!(error, Error::Http { status: 502, .. }), + "{error:?}" + ); + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn billable_post_is_retried_on_rate_limit() { + fn limited_once(call: usize) -> Response { + if call == 0 { + ( + StatusCode::TOO_MANY_REQUESTS, + [("retry-after", "0")], + "slow down", + ) + .into_response() + } else { + png_reply(call) + } + } + let fixture = fixture("/api/v1", limited_once, None).await; + generator(&fixture.base_url) + .with_capability_check(false) + .generate(ImageRequest::new("x")) + .await + .unwrap(); + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn provider_errors_and_debug_never_leak_the_key() { + fn echo_key(_call: usize) -> Response { + ( + StatusCode::UNAUTHORIZED, + axum::Json(json!({"error": {"code": 401, "message": format!("bad key {KEY}")}})), + ) + .into_response() + } + let fixture = fixture("/api/v1", echo_key, None).await; + let generator = generator(&fixture.base_url).with_capability_check(false); + let error = generator + .generate(ImageRequest::new("x")) + .await + .unwrap_err(); + assert!(matches!(error, Error::Auth(_)), "{error:?}"); + assert!(!error.to_string().contains(KEY), "{error}"); + assert!(!format!("{generator:?}").contains(KEY)); +} + +#[tokio::test] +async fn validation_errors_are_local() { + let fixture = fixture("/api/v1", png_reply, None).await; + let generator = generator(&fixture.base_url); + assert!(matches!( + generator.generate(ImageRequest::new(" ")).await, + Err(Error::Validation(_)) + )); + assert!(matches!( + generator.generate(ImageRequest::new("x").with_n(11)).await, + Err(Error::Validation(_)) + )); + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn lists_models_from_both_listing_shapes() { + let listing = json!({ "object": "list", "data": [ + { "id": "a/one", "display_name": "One" }, + { "id": "b/two", "name": "Two", "supported_parameters": { "seed": { "type": "boolean" } } } + ]}); + let fixture = fixture("/api/v1", png_reply, Some(listing)).await; + let models = generator(&fixture.base_url).list_models().await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].name.as_deref(), Some("One")); + assert_eq!(models[0].capabilities.seed, None); + assert_eq!(models[1].capabilities.seed, Some(true)); +} + +/// A billed response whose image payloads cannot be decoded is a billed +/// non-delivery (`NoMedia`, do not retry), not a retryable decode error. +#[tokio::test] +async fn undecodable_images_are_a_billed_non_delivery() { + fn garbage(_call: usize) -> Response { + axum::Json(json!({ "created": 1, "data": [{ "b64_json": "!!!not base64!!!" }] })) + .into_response() + } + let fixture = fixture("/api/v1", garbage, None).await; + let error = generator(&fixture.base_url) + .with_capability_check(false) + .generate(ImageRequest::new("x")) + .await + .unwrap_err(); + assert!(matches!(error, Error::NoMedia { .. }), "{error:?}"); + assert!(!error.is_retryable()); +} + +#[test] +fn transport_debug_redacts_base_url_userinfo() { + let transport = MediaTransport::new(MediaAuth::ApiKey(KEY.into())) + .with_base_url("https://user:hunter2secret@proxy.example/agent-integrations/openrouter"); + let debug = format!("{transport:?}"); + assert!(!debug.contains("hunter2secret"), "{debug}"); + assert!(!transport.redacted_base_url().contains("hunter2secret")); +} + +/// Per OpenRouter's listing contract, a key absent from a present +/// `supported_parameters` map is unsupported: the request fails before the +/// billed call instead of the provider silently ignoring the field. +#[tokio::test] +async fn omitted_enum_capability_is_unsupported() { + let listing = json!({ "data": [{ + "id": "google/gemini-3.1-flash-lite-image", + "supported_parameters": { "n": { "type": "range", "min": 1, "max": 1 } } + }]}); + let fixture = fixture("/api/v1", png_reply, Some(listing)).await; + let error = generator(&fixture.base_url) + .generate( + ImageRequest::new("x") + .with_model("google/gemini-3.1-flash-lite-image") + .with_resolution("2K"), + ) + .await + .unwrap_err(); + assert!( + matches!(&error, Error::Unsupported { field, allowed, .. } if field == "resolution" && allowed.is_empty()), + "{error:?}" + ); + assert_eq!(fixture.captured.image_calls.load(Ordering::SeqCst), 0); +} diff --git a/crates/tinyinference-image/src/reference.rs b/crates/tinyinference-image/src/reference.rs new file mode 100644 index 0000000..37c83e9 --- /dev/null +++ b/crates/tinyinference-image/src/reference.rs @@ -0,0 +1,418 @@ +//! Media reference and output-shape standards shared by image and video +//! generation. +//! +//! Callers hand references to a generator in whatever form they hold them — +//! an HTTP(S) URL, a `data:` URL, raw bytes, or a local file path — and this +//! module normalizes each one into the content-part shape OpenRouter's media +//! APIs accept (`{"type": "image_url", "image_url": {"url": …}}`, and the +//! `video_url` / `audio_url` equivalents). Local files and bytes are inlined as +//! base64 `data:` URLs, bounded by a size cap, so a reference never depends on +//! the provider being able to reach the caller's filesystem. +//! +//! It also normalizes the loose spellings people and models use for output +//! shape — `"16x9"`, `"landscape"`, `"1080"`, `"full hd"`, `"1024×1024"` — into +//! the canonical values the wire format expects. + +use std::path::{Path, PathBuf}; + +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD as BASE64; +use bytes::Bytes; +use serde::Serialize; + +use crate::{Error, Result}; + +/// Default cap on one inlined reference (20 MiB of raw bytes). +pub const DEFAULT_MAX_REFERENCE_BYTES: usize = 20 * 1024 * 1024; + +/// The modality a reference carries, which selects its wire content-part type. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum ReferenceKind { + /// A still image (`image_url`). + Image, + /// A video clip (`video_url`). + Video, + /// An audio clip (`audio_url`). + Audio, +} + +impl ReferenceKind { + /// Infers the kind from a MIME type, defaulting to [`ReferenceKind::Image`]. + #[must_use] + pub fn from_media_type(media_type: &str) -> Self { + let media_type = media_type.trim().to_ascii_lowercase(); + if media_type.starts_with("video/") { + Self::Video + } else if media_type.starts_with("audio/") { + Self::Audio + } else { + Self::Image + } + } +} + +/// A caller-supplied reference asset. +#[derive(Clone, PartialEq, Eq)] +pub enum MediaReference { + /// A publicly reachable HTTP(S) URL, forwarded as-is. + Url(String), + /// An inline `data:;base64,` URL, forwarded as-is. + DataUrl(String), + /// Raw bytes with their media type, inlined as a `data:` URL. + Bytes { + /// MIME type such as `image/png` or `video/mp4`. + media_type: String, + /// The asset bytes. + data: Bytes, + }, + /// A local file, read and inlined as a `data:` URL when the request is built. + Path(PathBuf), + /// A URL with an explicitly-specified modality, useful for extensionless CDN URLs. + Typed { + /// The modality of this reference (image, video, or audio). + kind: ReferenceKind, + /// The URL, either HTTP(S) or a local file path. + url: String, + }, +} + +impl MediaReference { + /// Classifies a free-form string: `http(s)://` is a URL, `data:` is a data + /// URL, and anything else is treated as a local path. + #[must_use] + pub fn parse(value: &str) -> Self { + let value = value.trim(); + let lower = value.to_ascii_lowercase(); + if lower.starts_with("http://") || lower.starts_with("https://") { + Self::Url(value.to_owned()) + } else if lower.starts_with("data:") { + Self::DataUrl(value.to_owned()) + } else { + Self::Path(PathBuf::from(value)) + } + } + + /// The reference's modality, from its media type, data-URL header, or + /// file/URL extension. + #[must_use] + pub fn kind(&self) -> ReferenceKind { + match self { + Self::Typed { kind, .. } => *kind, + Self::Bytes { media_type, .. } => ReferenceKind::from_media_type(media_type), + Self::DataUrl(url) => ReferenceKind::from_media_type( + url.get(5..) + .and_then(|rest| rest.split([';', ',']).next()) + .unwrap_or_default(), + ), + Self::Url(url) => { + let path = url.split(['?', '#']).next().unwrap_or(url); + ReferenceKind::from_media_type(media_type_for_path(Path::new(path))) + } + Self::Path(path) => ReferenceKind::from_media_type(media_type_for_path(path)), + } + } + + /// Resolves the reference into the URL string sent on the wire, inlining + /// bytes and local files as base64 `data:` URLs. + /// + /// # Errors + /// + /// [`Error::Validation`] for an empty or malformed reference, + /// [`Error::TooLarge`] when the asset exceeds `max_bytes`, and + /// [`Error::Io`] when a local file cannot be read. + pub async fn resolve(&self, max_bytes: usize) -> Result { + match self { + Self::Typed { kind, url } => { + let trimmed = url.trim(); + if trimmed.is_empty() { + return Err(Error::Validation("reference URL is empty".into())); + } + match Self::parse(trimmed) { + // A local path is inlined like `Path`, but keeps the + // caller-stated kind when the extension says nothing. + Self::Path(path) => { + let data = read_local(&path, max_bytes).await?; + let guessed = media_type_for_path(&path); + let media_type = if guessed == "application/octet-stream" { + default_media_type(*kind) + } else { + guessed + }; + Ok(data_url(media_type, &data)) + } + other => Box::pin(other.resolve(max_bytes)).await, + } + } + Self::Url(url) => { + if url.trim().is_empty() { + return Err(Error::Validation("reference URL is empty".into())); + } + Ok(url.clone()) + } + Self::DataUrl(url) => { + let Some((header, payload)) = url.split_once(',') else { + return Err(Error::Validation("malformed data: URL reference".into())); + }; + if !header.to_ascii_lowercase().starts_with("data:") || payload.is_empty() { + return Err(Error::Validation("malformed data: URL reference".into())); + } + // Reject on the encoded length first (base64 inflates by + // 4/3), so an oversized payload is never decoded. + if payload.len() / 4 * 3 > max_bytes.saturating_add(3) { + return Err(Error::TooLarge { limit: max_bytes }); + } + // Then decode, which validates the payload, and check the + // exact size. + let decoded = BASE64 + .decode(payload) + .map_err(|_| Error::Validation("malformed base64 in data: URL".into()))?; + if decoded.len() > max_bytes { + return Err(Error::TooLarge { limit: max_bytes }); + } + Ok(url.clone()) + } + Self::Bytes { media_type, data } => { + if data.is_empty() { + return Err(Error::Validation("reference bytes are empty".into())); + } + if data.len() > max_bytes { + return Err(Error::TooLarge { limit: max_bytes }); + } + Ok(data_url(media_type, data)) + } + Self::Path(path) => { + let data = read_local(path, max_bytes).await?; + Ok(data_url(media_type_for_path(path), &data)) + } + } + } + + /// Resolves the reference into an OpenRouter content part. + /// + /// # Errors + /// + /// As for [`MediaReference::resolve`]. + pub async fn to_content_part(&self, max_bytes: usize) -> Result { + let url = self.resolve(max_bytes).await?; + Ok(content_part(self.kind(), &url)) + } +} + +impl std::fmt::Debug for MediaReference { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Never print inline payloads or signed-URL query strings. + match self { + Self::Typed { kind, url } => { + let sanitized = url.split(['?', '#']).next().unwrap_or(url); + let sanitized = sanitized.split('@').next_back().unwrap_or(sanitized); + formatter + .debug_struct("Typed") + .field("kind", kind) + .field("url", &sanitized) + .finish() + } + Self::Url(url) => { + let sanitized = url.split(['?', '#']).next().unwrap_or(url); + let sanitized = sanitized.split('@').next_back().unwrap_or(sanitized); + formatter.debug_tuple("Url").field(&sanitized).finish() + } + Self::DataUrl(_) => formatter + .debug_struct("DataUrl") + .field("data", &"") + .finish(), + Self::Bytes { media_type, data } => formatter + .debug_struct("Bytes") + .field("media_type", media_type) + .field("len", &data.len()) + .finish(), + Self::Path(path) => formatter.debug_tuple("Path").field(path).finish(), + } + } +} + +/// Builds an OpenRouter content part for a resolved URL. +#[must_use] +pub fn content_part(kind: ReferenceKind, url: &str) -> serde_json::Value { + let key = match kind { + ReferenceKind::Image => "image_url", + ReferenceKind::Video => "video_url", + ReferenceKind::Audio => "audio_url", + }; + serde_json::json!({ "type": key, key: { "url": url } }) +} + +/// Reads a local reference, enforcing the size cap before reading. +async fn read_local(path: &Path, max_bytes: usize) -> Result> { + let metadata = tokio::fs::metadata(path).await?; + if metadata.len() > max_bytes as u64 { + return Err(Error::TooLarge { limit: max_bytes }); + } + let data = tokio::fs::read(path).await?; + if data.is_empty() { + return Err(Error::Validation(format!( + "reference file {} is empty", + path.display() + ))); + } + Ok(data) +} + +/// A generic media type for a kind, used when a typed local file's extension +/// identifies nothing; providers sniff the actual format from the bytes. +fn default_media_type(kind: ReferenceKind) -> &'static str { + match kind { + ReferenceKind::Image => "image/png", + ReferenceKind::Video => "video/mp4", + ReferenceKind::Audio => "audio/mpeg", + } +} + +fn data_url(media_type: &str, data: &[u8]) -> String { + format!("data:{media_type};base64,{}", BASE64.encode(data)) +} + +/// Guesses a MIME type from a file extension, defaulting to +/// `application/octet-stream`. +#[must_use] +pub fn media_type_for_path(path: &Path) -> &'static str { + let extension = path + .extension() + .and_then(|extension| extension.to_str()) + .unwrap_or_default() + .to_ascii_lowercase(); + match extension.as_str() { + "png" => "image/png", + "jpg" | "jpeg" => "image/jpeg", + "webp" => "image/webp", + "gif" => "image/gif", + "bmp" => "image/bmp", + "svg" => "image/svg+xml", + "heic" => "image/heic", + "avif" => "image/avif", + "mp4" | "m4v" => "video/mp4", + "webm" => "video/webm", + "mov" => "video/quicktime", + "mp3" => "audio/mpeg", + "wav" => "audio/wav", + "m4a" => "audio/mp4", + "ogg" => "audio/ogg", + "flac" => "audio/flac", + _ => "application/octet-stream", + } +} + +/// Picks a file extension for a MIME type, falling back to `fallback`. +#[must_use] +pub fn extension_for_media_type<'a>(media_type: &str, fallback: &'a str) -> &'a str { + let media_type = media_type + .split(';') + .next() + .unwrap_or_default() + .trim() + .to_ascii_lowercase(); + match media_type.as_str() { + "image/png" => "png", + "image/jpeg" | "image/jpg" => "jpg", + "image/webp" => "webp", + "image/gif" => "gif", + "image/svg+xml" => "svg", + "image/avif" => "avif", + "video/mp4" => "mp4", + "video/webm" => "webm", + "video/quicktime" => "mov", + "audio/mpeg" => "mp3", + "audio/wav" | "audio/x-wav" => "wav", + _ => fallback, + } +} + +/// Normalizes an aspect-ratio spelling to the canonical `W:H` form. +/// +/// Accepts `16:9`, `16x9`, `16/9`, `16 × 9`, and the names `square`, +/// `landscape`, `portrait`, `widescreen`, `vertical`, `ultrawide`, `auto`. +/// Returns `None` when the value is not recognizable, in which case callers +/// should forward it unchanged and let the provider decide. +#[must_use] +pub fn normalize_aspect_ratio(value: &str) -> Option { + let value = value.trim().to_ascii_lowercase(); + let named = match value.as_str() { + "auto" => Some("auto"), + "square" => Some("1:1"), + "landscape" | "widescreen" | "horizontal" => Some("16:9"), + "portrait" | "vertical" | "story" | "reel" => Some("9:16"), + "ultrawide" | "cinematic" => Some("21:9"), + _ => None, + }; + if let Some(named) = named { + return Some(named.to_owned()); + } + let separators: &[char] = &[':', 'x', '/', '×', '*']; + let mut parts = value.split(separators).map(str::trim); + let (width, height) = (parts.next()?, parts.next()?); + if parts.next().is_some() { + return None; + } + let valid = |part: &str| !part.is_empty() && part.parse::().is_ok_and(|n| n > 0.0); + (valid(width) && valid(height)).then(|| format!("{width}:{height}")) +} + +/// Normalizes an image resolution tier (`512`, `1K`, `2K`, `4K`). +/// +/// Accepts case and spelling variants (`1k`, `1024`, `2048`, `4096`, `hd`, +/// `4k uhd`). Returns `None` when unrecognized. +#[must_use] +pub fn normalize_image_resolution(value: &str) -> Option { + let value = value.trim().to_ascii_lowercase().replace(' ', ""); + let tier = match value.as_str() { + "512" | "0.5k" | "sd" => "512", + "1k" | "1024" | "hd" => "1K", + "2k" | "2048" | "qhd" => "2K", + "4k" | "4096" | "uhd" | "4kuhd" => "4K", + _ => return None, + }; + Some(tier.to_owned()) +} + +/// Normalizes a video resolution (`360p` … `1080p`, `1K`, `2K`, `4K`). +/// +/// Accepts `720`, `720P`, `hd`, `full hd`, `fhd`, `1080`, `4k`, `uhd`. +/// Returns `None` when unrecognized. +#[must_use] +pub fn normalize_video_resolution(value: &str) -> Option { + let value = value.trim().to_ascii_lowercase().replace([' ', '-'], ""); + let resolution = match value.as_str() { + "360" | "360p" => "360p", + "480" | "480p" | "sd" => "480p", + "720" | "720p" | "hd" => "720p", + "768" | "768p" => "768p", + "1080" | "1080p" | "fullhd" | "fhd" => "1080p", + "1k" => "1K", + "2k" | "1440p" | "qhd" => "2K", + "4k" | "2160p" | "uhd" => "4K", + _ => return None, + }; + Some(resolution.to_owned()) +} + +/// Normalizes an explicit pixel size to `WIDTHxHEIGHT`, accepting `×`, `*`, +/// and surrounding whitespace. Tier sizes (`2K`) are returned uppercased. +/// Returns `None` when the value is neither. +#[must_use] +pub fn normalize_size(value: &str) -> Option { + let trimmed = value.trim(); + if let Some(tier) = + normalize_image_resolution(trimmed).filter(|_| trimmed.ends_with(['k', 'K'])) + { + return Some(tier); + } + let lower = trimmed.to_ascii_lowercase(); + let mut parts = lower.split(['x', '×', '*']).map(str::trim); + let (width, height) = (parts.next()?, parts.next()?); + if parts.next().is_some() { + return None; + } + let width: u32 = width.parse().ok().filter(|n| *n > 0)?; + let height: u32 = height.parse().ok().filter(|n| *n > 0)?; + Some(format!("{width}x{height}")) +} diff --git a/crates/tinyinference-image/src/reference_test.rs b/crates/tinyinference-image/src/reference_test.rs new file mode 100644 index 0000000..343cd9e --- /dev/null +++ b/crates/tinyinference-image/src/reference_test.rs @@ -0,0 +1,237 @@ +//! Tests for media reference and output-shape standards. + +use std::path::Path; + +use bytes::Bytes; + +use crate::media::GeneratedMedia; +use crate::mock::TINY_PNG; +use crate::reference::{ + MediaReference, ReferenceKind, content_part, extension_for_media_type, media_type_for_path, + normalize_aspect_ratio, normalize_image_resolution, normalize_size, normalize_video_resolution, +}; +use crate::{Error, ImageGenerator, ImageRequest, MockImageGenerator}; + +#[test] +fn aspect_ratio_spellings_normalize() { + for (input, expected) in [ + ("16:9", "16:9"), + ("16x9", "16:9"), + ("16/9", "16:9"), + (" 9 × 16 ", "9:16"), + ("Landscape", "16:9"), + ("portrait", "9:16"), + ("square", "1:1"), + ("auto", "auto"), + ("2.35:1", "2.35:1"), + ] { + assert_eq!( + normalize_aspect_ratio(input).as_deref(), + Some(expected), + "{input}" + ); + } + for input in ["wide-ish", "16:0", "1:2:3", ""] { + assert_eq!(normalize_aspect_ratio(input), None, "{input}"); + } +} + +#[test] +fn resolution_spellings_normalize() { + assert_eq!(normalize_image_resolution("2k").as_deref(), Some("2K")); + assert_eq!(normalize_image_resolution("1024").as_deref(), Some("1K")); + assert_eq!(normalize_image_resolution("4K UHD").as_deref(), Some("4K")); + assert_eq!(normalize_image_resolution("720p"), None); + assert_eq!(normalize_video_resolution("720").as_deref(), Some("720p")); + assert_eq!( + normalize_video_resolution("Full HD").as_deref(), + Some("1080p") + ); + assert_eq!(normalize_video_resolution("4k").as_deref(), Some("4K")); + assert_eq!(normalize_video_resolution("hd").as_deref(), Some("720p")); + assert_eq!(normalize_video_resolution("8k"), None); +} + +#[test] +fn sizes_normalize() { + assert_eq!(normalize_size("1536x1024").as_deref(), Some("1536x1024")); + assert_eq!(normalize_size(" 1024 × 768 ").as_deref(), Some("1024x768")); + assert_eq!(normalize_size("2k").as_deref(), Some("2K")); + assert_eq!(normalize_size("0x10"), None); + assert_eq!(normalize_size("big"), None); +} + +#[test] +fn references_classify_and_infer_kind() { + assert!(matches!( + MediaReference::parse("https://x.test/a.png"), + MediaReference::Url(_) + )); + assert!(matches!( + MediaReference::parse("data:image/png;base64,AA=="), + MediaReference::DataUrl(_) + )); + assert!(matches!( + MediaReference::parse("./frames/first.jpg"), + MediaReference::Path(_) + )); + + assert_eq!( + MediaReference::parse("https://x.test/clip.mp4?sig=1").kind(), + ReferenceKind::Video + ); + assert_eq!( + MediaReference::parse("data:audio/wav;base64,AA==").kind(), + ReferenceKind::Audio + ); + assert_eq!( + MediaReference::parse("photo.jpeg").kind(), + ReferenceKind::Image + ); +} + +#[test] +fn content_parts_match_the_wire_shape() { + let image = content_part(ReferenceKind::Image, "https://x.test/a.png"); + assert_eq!(image["type"], "image_url"); + assert_eq!(image["image_url"]["url"], "https://x.test/a.png"); + let video = content_part(ReferenceKind::Video, "https://x.test/a.mp4"); + assert_eq!(video["type"], "video_url"); + assert_eq!(video["video_url"]["url"], "https://x.test/a.mp4"); +} + +#[tokio::test] +async fn local_files_inline_as_data_urls_within_the_cap() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("ref.png"); + std::fs::write(&path, TINY_PNG).unwrap(); + + let reference = MediaReference::Path(path.clone()); + let url = reference.resolve(1024).await.unwrap(); + assert!(url.starts_with("data:image/png;base64,"), "{url}"); + + let error = reference.resolve(8).await.unwrap_err(); + assert!(matches!(error, Error::TooLarge { limit: 8 }), "{error:?}"); + + let missing = MediaReference::Path(dir.path().join("missing.png")); + assert!(matches!(missing.resolve(1024).await, Err(Error::Io(_)))); +} + +#[tokio::test] +async fn malformed_and_empty_references_are_rejected() { + assert!(matches!( + MediaReference::DataUrl("data:image/png;base64".into()) + .resolve(1024) + .await, + Err(Error::Validation(_)) + )); + assert!(matches!( + MediaReference::Bytes { + media_type: "image/png".into(), + data: Bytes::new() + } + .resolve(1024) + .await, + Err(Error::Validation(_)) + )); +} + +#[test] +fn debug_output_hides_payloads_and_signatures() { + let debug = format!( + "{:?} {:?}", + MediaReference::Url("https://x.test/a.png?X-Amz-Signature=secret".into()), + MediaReference::DataUrl("data:image/png;base64,SECRETPAYLOAD".into()) + ); + assert!( + !debug.contains("secret") && !debug.contains("SECRETPAYLOAD"), + "{debug}" + ); +} + +#[test] +fn media_types_and_extensions_round_trip() { + assert_eq!(media_type_for_path(Path::new("a.WEBP")), "image/webp"); + assert_eq!(media_type_for_path(Path::new("a.mov")), "video/quicktime"); + assert_eq!(extension_for_media_type("image/jpeg", "bin"), "jpg"); + assert_eq!( + extension_for_media_type("video/mp4; codecs=avc1", "bin"), + "mp4" + ); + assert_eq!( + extension_for_media_type("application/x-unknown", "bin"), + "bin" + ); +} + +#[tokio::test] +async fn persist_sanitizes_the_stem_and_refuses_empty_artifacts() { + let dir = tempfile::tempdir().unwrap(); + let media = GeneratedMedia::new("image/png", TINY_PNG); + let path = media + .persist(dir.path(), "../../etc/passwd", "bin") + .await + .unwrap(); + assert_eq!(path.parent().unwrap(), dir.path()); + assert_eq!(path.file_name().unwrap(), "______etc_passwd.png"); + assert_eq!(std::fs::read(&path).unwrap(), TINY_PNG); + + let empty = GeneratedMedia::new("image/png", Bytes::new()); + assert!(matches!( + empty.persist(dir.path(), "x", "png").await, + Err(Error::Validation(_)) + )); +} + +#[tokio::test] +async fn mock_generator_records_requests_and_simulates_no_media() { + let mock = MockImageGenerator::new(); + let response = mock + .generate(ImageRequest::new("x").with_n(2)) + .await + .unwrap(); + assert_eq!(response.images.len(), 2); + assert_eq!(mock.requests().len(), 1); + + let failing = MockImageGenerator::returning_no_media(); + assert!(matches!( + failing.generate(ImageRequest::new("x")).await, + Err(Error::NoMedia { .. }) + )); +} + +/// A `Typed` reference that names a local file is inlined (never sent as a +/// raw path), and an extensionless file takes its media type from the stated +/// kind. +#[tokio::test] +async fn typed_local_paths_are_inlined_with_the_stated_kind() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("frame-without-extension"); + std::fs::write(&path, TINY_PNG).unwrap(); + let reference = MediaReference::Typed { + kind: ReferenceKind::Video, + url: path.display().to_string(), + }; + let url = reference.resolve(1024).await.unwrap(); + assert!(url.starts_with("data:video/mp4;base64,"), "{url}"); + assert_eq!(reference.kind(), ReferenceKind::Video); + + let remote = MediaReference::Typed { + kind: ReferenceKind::Audio, + url: "https://cdn.test/opaque-id".into(), + }; + assert_eq!( + remote.resolve(1024).await.unwrap(), + "https://cdn.test/opaque-id" + ); +} + +#[tokio::test] +async fn oversized_data_urls_are_rejected_before_decoding() { + let payload = "A".repeat(4_000); + let reference = MediaReference::DataUrl(format!("data:image/png;base64,{payload}")); + assert!(matches!( + reference.resolve(100).await, + Err(Error::TooLarge { limit: 100 }) + )); +} diff --git a/crates/tinyinference-image/src/transport.rs b/crates/tinyinference-image/src/transport.rs new file mode 100644 index 0000000..d86ef6e --- /dev/null +++ b/crates/tinyinference-image/src/transport.rs @@ -0,0 +1,520 @@ +//! Authenticated HTTP transport for OpenRouter's media APIs. +//! +//! One transport serves two deployments that speak the same wire format: +//! +//! - **Direct** — `https://openrouter.ai/api/v1` with the caller's OpenRouter +//! API key ([`MediaAuth::ApiKey`]). +//! - **Proxied** — a host backend that forwards OpenRouter's request bodies +//! verbatim and returns OpenRouter's response body, optionally wrapped in a +//! `{"success": true, "data": …}` envelope that is unwrapped transparently +//! ([`unwrap_envelope`]) (for example TinyHumans' +//! `/agent-integrations/openrouter`), authenticated with a host-owned bearer +//! ([`MediaAuth::Bearer`]) so the credential lifecycle stays in the host. +//! +//! Retry policy is billing-aware. `GET` calls (model listings, job polls, +//! content downloads) retry on 429/5xx/transport failures. A `POST` that may +//! have started a paid generation is retried **only** on HTTP 429, where the +//! provider rejected the request before doing any work; a 5xx or a dropped +//! connection after a submit is surfaced as-is, because resubmitting could +//! bill a second generation. + +use std::sync::Arc; +use std::time::Duration; + +use bytes::Bytes; +use reqwest::header::{AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue}; +use serde::Serialize; +use serde::de::DeserializeOwned; +use tinyinference_core::retry_after::{MAX_RETRIES, backoff_ms_for_attempt}; +use tinyinference_core::sanitize::sanitize_api_error; + +use crate::{Error, Result}; + +/// OpenRouter's public API base URL. +pub const OPENROUTER_BASE_URL: &str = "https://openrouter.ai/api/v1"; +/// Environment variable read by [`MediaAuth::from_env`]. +pub const OPENROUTER_API_KEY_ENV: &str = "OPENROUTER_API_KEY"; +/// Default cap on a downloaded or decoded media body (512 MiB). +pub const DEFAULT_MAX_MEDIA_BYTES: usize = 512 * 1024 * 1024; +/// Cap on an error body read into a message. +const MAX_ERROR_BODY_BYTES: usize = 16 * 1024; + +/// Resolves the current bearer token for each request. +pub type BearerResolver = Arc Result + Send + Sync>; + +/// How requests authenticate. +#[derive(Clone)] +pub enum MediaAuth { + /// A static API key sent as `Authorization: Bearer `. + ApiKey(String), + /// A host-owned resolver called before every request. + Bearer(BearerResolver), +} + +impl MediaAuth { + /// Reads an OpenRouter API key from [`OPENROUTER_API_KEY_ENV`]. + /// + /// # Errors + /// + /// [`Error::Auth`] when the variable is unset or blank. + pub fn from_env() -> Result { + match std::env::var(OPENROUTER_API_KEY_ENV) { + Ok(key) if !key.trim().is_empty() => Ok(Self::ApiKey(key.trim().to_owned())), + _ => Err(Error::Auth(format!("{OPENROUTER_API_KEY_ENV} is not set"))), + } + } + + fn token(&self) -> Result { + let token = match self { + Self::ApiKey(key) => key.clone(), + Self::Bearer(resolve) => resolve()?, + }; + let token = token.trim().to_owned(); + if token.is_empty() { + return Err(Error::Auth( + "no credential available for media generation".into(), + )); + } + Ok(token) + } +} + +impl std::fmt::Debug for MediaAuth { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::ApiKey(_) => formatter.write_str("MediaAuth::ApiKey()"), + Self::Bearer(_) => formatter.write_str("MediaAuth::Bearer()"), + } + } +} + +/// Whether a request may start billable work, which decides its retry policy. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Billing { + /// Safe to repeat: listings, polls, downloads. + Idempotent, + /// May start a paid generation: only a 429 is retried. + Billable, +} + +/// HTTP client bound to one OpenRouter-compatible base URL and credential. +#[derive(Clone)] +pub struct MediaTransport { + client: reqwest::Client, + base_url: String, + auth: MediaAuth, + headers: HeaderMap, + max_retries: u32, + max_media_bytes: usize, +} + +impl MediaTransport { + /// Creates a transport for OpenRouter's public API. + #[must_use] + pub fn new(auth: MediaAuth) -> Self { + Self { + client: reqwest::Client::new(), + base_url: OPENROUTER_BASE_URL.to_owned(), + auth, + headers: HeaderMap::new(), + max_retries: MAX_RETRIES, + max_media_bytes: DEFAULT_MAX_MEDIA_BYTES, + } + } + + /// Points the transport at another OpenRouter-compatible base URL, such as + /// a host backend that proxies OpenRouter's media routes. + #[must_use] + pub fn with_base_url(mut self, base_url: impl AsRef) -> Self { + self.base_url = base_url.as_ref().trim().trim_end_matches('/').to_owned(); + self + } + + /// Uses a caller-supplied HTTP client (timeouts, proxies, TLS roots). + #[must_use] + pub fn with_client(mut self, client: reqwest::Client) -> Self { + self.client = client; + self + } + + /// Adds a header sent on every request (for example `HTTP-Referer`, + /// `X-Title`, or a host's client-identification header). + /// + /// Invalid header names or values are ignored with a warning rather than + /// failing construction. + #[must_use] + pub fn with_header(mut self, name: &str, value: &str) -> Self { + match ( + HeaderName::from_bytes(name.as_bytes()), + HeaderValue::from_str(value), + ) { + (Ok(name), Ok(value)) => { + self.headers.insert(name, value); + } + _ => tracing::warn!( + header = name, + "[tinyinference-image] ignoring invalid header" + ), + } + self + } + + /// Sets how many times a retryable failure is retried (default 3). + #[must_use] + pub fn with_max_retries(mut self, max_retries: u32) -> Self { + self.max_retries = max_retries; + self + } + + /// Caps the size of any single media body (default 512 MiB). + #[must_use] + pub fn with_max_media_bytes(mut self, max_media_bytes: usize) -> Self { + self.max_media_bytes = max_media_bytes; + self + } + + /// The configured base URL, without a trailing slash. + #[must_use] + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// The configured media size cap in bytes. + #[must_use] + pub fn max_media_bytes(&self) -> usize { + self.max_media_bytes + } + + /// Sends a JSON `POST` that may start a paid generation. + /// + /// # Errors + /// + /// [`Error::Auth`], [`Error::Http`], [`Error::Transport`] or + /// [`Error::Decode`]. Only HTTP 429 is retried. + pub async fn post_json(&self, path: &str, body: &B) -> Result + where + B: Serialize + ?Sized + Sync, + T: DeserializeOwned, + { + let body = serde_json::to_vec(body)?; + let response = self + .send(reqwest::Method::POST, path, Some(body), Billing::Billable) + .await?; + decode_json_with_limit(response, self.json_limit()).await + } + + /// Sends an idempotent JSON `GET`. + /// + /// # Errors + /// + /// [`Error::Auth`], [`Error::Http`], [`Error::Transport`] or + /// [`Error::Decode`] after retries are exhausted. + pub async fn get_json(&self, path: &str) -> Result { + let response = self + .send(reqwest::Method::GET, path, None, Billing::Idempotent) + .await?; + decode_json_with_limit(response, self.json_limit()).await + } + + /// Cap on a JSON body: the configured media cap plus base64's 4/3 + /// inflation and a little envelope headroom, so lowering + /// [`MediaTransport::with_max_media_bytes`] also bounds image JSON. + fn json_limit(&self) -> usize { + self.max_media_bytes + .saturating_div(3) + .saturating_mul(4) + .saturating_add(64 * 1024) + } + + /// The base URL with any userinfo and query string redacted, for logs. + #[must_use] + pub fn redacted_base_url(&self) -> String { + tinyinference_core::sanitize::redact_url(&self.base_url) + } + + /// Downloads a binary body with an idempotent `GET`, enforcing the media + /// size cap. Returns the bytes and the response `Content-Type`, if any. + /// + /// # Errors + /// + /// [`Error::TooLarge`] when the body exceeds the cap, plus the errors of + /// [`MediaTransport::get_json`]. + pub async fn get_bytes(&self, path: &str) -> Result<(Bytes, Option)> { + // A dropped connection mid-body is as transient as one before the + // headers, and the GET is idempotent, so the whole download is retried + // under the same budget as a failed request. + let mut attempt = 0u32; + loop { + match self.get_bytes_once(path).await { + Err(Error::Transport(message)) if attempt < self.max_retries => { + let delay = backoff_ms_for_attempt(attempt, None); + tracing::debug!( + path, + attempt, + delay_ms = delay, + error = %message, + "[tinyinference-image] retrying interrupted download" + ); + tokio::time::sleep(Duration::from_millis(delay)).await; + attempt += 1; + } + other => return other, + } + } + } + + async fn get_bytes_once(&self, path: &str) -> Result<(Bytes, Option)> { + let token = self.auth.token()?; + let mut response = self + .send(reqwest::Method::GET, path, None, Billing::Idempotent) + .await?; + let content_type = response + .headers() + .get(CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + if response + .content_length() + .is_some_and(|length| length > self.max_media_bytes as u64) + { + return Err(Error::TooLarge { + limit: self.max_media_bytes, + }); + } + let mut body = Vec::new(); + while let Some(chunk) = response + .chunk() + .await + .map_err(|error| Error::Transport(self.scrub(&error.to_string(), Some(&token))))? + { + if body.len() + chunk.len() > self.max_media_bytes { + return Err(Error::TooLarge { + limit: self.max_media_bytes, + }); + } + body.extend_from_slice(&chunk); + } + Ok((Bytes::from(body), content_type)) + } + + fn url(&self, path: &str) -> String { + format!("{}/{}", self.base_url, path.trim_start_matches('/')) + } + + fn scrub(&self, text: &str, token: Option<&str>) -> String { + let mut text = text.to_owned(); + if let Some(token) = token + && !token.is_empty() + { + text = text.replace(token, "[REDACTED]"); + } + sanitize_api_error(&text) + } + + async fn send( + &self, + method: reqwest::Method, + path: &str, + body: Option>, + billing: Billing, + ) -> Result { + let url = self.url(path); + let mut attempt = 0u32; + loop { + let token = self.auth.token()?; + let mut request = self + .client + .request(method.clone(), &url) + .headers(self.headers.clone()) + .header(AUTHORIZATION, format!("Bearer {token}")); + if let Some(body) = &body { + request = request + .header(CONTENT_TYPE, "application/json") + .body(body.clone()); + } + tracing::debug!( + method = %method, + path, + attempt, + "[tinyinference-image] media request" + ); + let outcome = request.send().await; + let (retry, retry_after, error) = match outcome { + Ok(response) if response.status().is_success() => return Ok(response), + Ok(response) => { + let status = response.status().as_u16(); + let retry_after = response + .headers() + .get(reqwest::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + let message = self.error_message(response, &token).await; + let error = match status { + 401 | 403 => Error::Auth(format!("HTTP {status}: {message}")), + _ => Error::Http { status, message }, + }; + let retry = match billing { + Billing::Billable => status == 429, + Billing::Idempotent => status == 429 || status >= 500, + }; + (retry, retry_after, error) + } + Err(error) => { + let error = Error::Transport(self.scrub(&error.to_string(), Some(&token))); + (billing == Billing::Idempotent, None, error) + } + }; + if !retry || attempt >= self.max_retries { + tracing::warn!( + method = %method, + path, + attempt, + error = %error, + "[tinyinference-image] media request failed" + ); + return Err(error); + } + let delay = backoff_ms_for_attempt(attempt, retry_after.as_deref()); + tracing::debug!( + method = %method, + path, + attempt, + delay_ms = delay, + "[tinyinference-image] retrying media request" + ); + tokio::time::sleep(Duration::from_millis(delay)).await; + attempt += 1; + } + } + + async fn error_message(&self, mut response: reqwest::Response, token: &str) -> String { + let mut body = Vec::new(); + loop { + match response.chunk().await { + Ok(Some(chunk)) => { + let room = MAX_ERROR_BODY_BYTES.saturating_sub(body.len()); + body.extend_from_slice(&chunk[..chunk.len().min(room)]); + if body.len() >= MAX_ERROR_BODY_BYTES { + break; + } + } + Ok(None) => break, + Err(error) => { + return self.scrub(&format!("response stream error: {error}"), Some(token)); + } + } + } + let text = String::from_utf8_lossy(&body).into_owned(); + let message = serde_json::from_str::(&text) + .ok() + .and_then(|value| { + value + .pointer("/error/message") + .or_else(|| value.get("message")) + .or_else(|| value.get("error")) + .and_then(|message| message.as_str().map(str::to_owned)) + }) + .unwrap_or(text); + self.scrub(&message, Some(token)) + } +} + +impl std::fmt::Debug for MediaTransport { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("MediaTransport") + .field("base_url", &self.redacted_base_url()) + .field("auth", &self.auth) + .field( + "headers", + &self + .headers + .keys() + .map(HeaderName::as_str) + .collect::>(), + ) + .field("max_retries", &self.max_retries) + .field("max_media_bytes", &self.max_media_bytes) + .finish() + } +} + +async fn decode_json_with_limit( + mut response: reqwest::Response, + limit: usize, +) -> Result { + let mut body = Vec::new(); + loop { + match response.chunk().await { + Ok(Some(chunk)) => { + let room = limit.saturating_sub(body.len()); + if room == 0 { + return Err(Error::TooLarge { limit }); + } + body.extend_from_slice(&chunk[..chunk.len().min(room)]); + } + Ok(None) => break, + Err(error) => return Err(Error::Transport(format!("response stream error: {error}"))), + } + } + let value: serde_json::Value = serde_json::from_slice(&body).map_err(|error| { + Error::Decode(format!( + "unexpected response body ({} bytes): {error}", + body.len() + )) + })?; + serde_json::from_value(unwrap_envelope(value)?) + .map_err(|error| Error::Decode(format!("unexpected response shape: {error}"))) +} + +/// Unwraps a proxying backend's `{"success": bool, "data": …}` envelope. +/// +/// OpenRouter's own media responses never carry a top-level boolean +/// `success`, so its presence identifies the envelope unambiguously; any other +/// body passes through untouched. +/// +/// # Errors +/// +/// [`Error::Http`] for a `success: false` envelope delivered with a 2xx +/// status, carrying the envelope's sanitized error message. +pub fn unwrap_envelope(value: serde_json::Value) -> Result { + let serde_json::Value::Object(mut map) = value else { + return Ok(value); + }; + match map.get("success").and_then(serde_json::Value::as_bool) { + Some(true) => Ok(map.remove("data").unwrap_or(serde_json::Value::Null)), + Some(false) => { + let message = map + .get("error") + .and_then(|error| { + error.as_str().map(str::to_owned).or_else(|| { + error + .pointer("/message") + .and_then(|m| m.as_str()) + .map(str::to_owned) + }) + }) + .or_else(|| { + map.get("message") + .and_then(|m| m.as_str()) + .map(str::to_owned) + }) + .unwrap_or_else(|| "request failed".to_owned()); + Err(Error::Http { + status: 200, + message: sanitize_api_error(&message), + }) + } + None => Ok(serde_json::Value::Object(map)), + } +} + +/// Strips an `openrouter/` routing prefix from a model id. +/// +/// Hosts often qualify OpenRouter slugs (`openrouter/bytedance/seedance-2.0-mini`) +/// to say which provider serves them; the wire format wants the bare slug. +#[must_use] +pub fn wire_model_id(model: &str) -> &str { + let model = model.trim(); + model.strip_prefix("openrouter/").unwrap_or(model) +} diff --git a/crates/tinyinference-image/src/types.rs b/crates/tinyinference-image/src/types.rs new file mode 100644 index 0000000..abf32c9 --- /dev/null +++ b/crates/tinyinference-image/src/types.rs @@ -0,0 +1,190 @@ +//! Public request, response and model-listing types for image generation. + +use serde_json::Value; + +use crate::capabilities::ModelCapabilities; +use crate::media::GeneratedMedia; +use crate::reference::MediaReference; +use crate::{Error, Result}; + +/// Maximum images per request accepted by OpenRouter's image API. +pub const MAX_IMAGES_PER_REQUEST: u32 = 10; + +/// A provider-neutral image generation request. +/// +/// Output-shape fields accept loose spellings (`"16x9"`, `"landscape"`, +/// `"2k"`, `"1024×1024"`); providers normalize them with the helpers in +/// [`reference`](mod@crate::reference) and forward anything unrecognized unchanged. +#[derive(Debug, Clone, Default, PartialEq)] +pub struct ImageRequest { + /// Model id; `None` uses the generator's default. An `openrouter/` prefix + /// is accepted and stripped. + pub model: Option, + /// Text description of the desired image, or the edit instruction when + /// references are supplied. + pub prompt: String, + /// Number of images (1–10); providers may return fewer. + pub n: Option, + /// Pixel size (`1536x1024`) or tier shorthand (`2K`). + pub size: Option, + /// Resolution tier (`512`, `1K`, `2K`, `4K`). + pub resolution: Option, + /// Aspect ratio (`16:9`, `auto`, …). + pub aspect_ratio: Option, + /// Rendering quality (`auto`, `low`, `medium`, `high`). + pub quality: Option, + /// Output encoding (`png`, `jpeg`, `webp`, `svg`). + pub output_format: Option, + /// Background treatment (`auto`, `transparent`, `opaque`). + pub background: Option, + /// Deterministic seed, where supported. + pub seed: Option, + /// Reference images for image-to-image generation and editing. + pub references: Vec, + /// Stable end-user identifier for provider abuse detection. + pub user: Option, + /// Observability grouping id (never sent to the upstream model provider). + pub session_id: Option, + /// Provider-specific extra fields merged into the wire body. + pub extra: serde_json::Map, +} + +impl ImageRequest { + /// Creates a request for `prompt` with every other field defaulted. + #[must_use] + pub fn new(prompt: impl Into) -> Self { + Self { + prompt: prompt.into(), + ..Self::default() + } + } + + /// Sets the model id. + #[must_use] + pub fn with_model(mut self, model: impl Into) -> Self { + self.model = Some(model.into()); + self + } + + /// Sets the number of images. + #[must_use] + pub fn with_n(mut self, n: u32) -> Self { + self.n = Some(n); + self + } + + /// Sets the aspect ratio. + #[must_use] + pub fn with_aspect_ratio(mut self, aspect_ratio: impl Into) -> Self { + self.aspect_ratio = Some(aspect_ratio.into()); + self + } + + /// Sets the resolution tier. + #[must_use] + pub fn with_resolution(mut self, resolution: impl Into) -> Self { + self.resolution = Some(resolution.into()); + self + } + + /// Sets the pixel size or tier shorthand. + #[must_use] + pub fn with_size(mut self, size: impl Into) -> Self { + self.size = Some(size.into()); + self + } + + /// Sets the seed. + #[must_use] + pub fn with_seed(mut self, seed: i64) -> Self { + self.seed = Some(seed); + self + } + + /// Adds a reference image. + #[must_use] + pub fn with_reference(mut self, reference: MediaReference) -> Self { + self.references.push(reference); + self + } + + /// Checks the fields that are invalid for every model. + /// + /// # Errors + /// + /// [`Error::Validation`] for a blank prompt or an out-of-range `n`. + pub fn validate(&self) -> Result<()> { + if self.prompt.trim().is_empty() { + return Err(Error::Validation("prompt is required".into())); + } + if let Some(n) = self.n + && !(1..=MAX_IMAGES_PER_REQUEST).contains(&n) + { + return Err(Error::Validation(format!( + "n must be between 1 and {MAX_IMAGES_PER_REQUEST}, got {n}" + ))); + } + Ok(()) + } +} + +/// The result of a successful image generation. Always carries at least one +/// image: an accepted request that returns none fails with +/// [`Error::NoMedia`] instead. +#[derive(Debug, Clone, PartialEq)] +pub struct ImageResponse { + /// Wire model id that served the request. + pub model: String, + /// Generated images, in provider order. + pub images: Vec, + /// Provider-reported cost in USD, when available. + pub cost_usd: Option, + /// Provider creation timestamp (Unix seconds), when available. + pub created: Option, +} + +/// One entry from a media model listing. +#[derive(Debug, Clone, PartialEq)] +pub struct MediaModel { + /// Model slug to pass as `model`. + pub id: String, + /// Human-readable name, when provided. + pub name: Option, + /// Short description, when provided. + pub description: Option, + /// Advertised capabilities (fields are `None` when not advertised). + pub capabilities: ModelCapabilities, + /// The raw listing record, for fields this crate does not model. + pub raw: Value, +} + +impl MediaModel { + /// Reads every record in a listing body. Accepts `{ "data": [...] }` (both + /// OpenRouter and proxying backends) or a bare array. + #[must_use] + pub fn parse_listing(body: &Value, capabilities: fn(&Value) -> ModelCapabilities) -> Vec { + let records = body + .get("data") + .and_then(Value::as_array) + .or_else(|| body.as_array()); + records + .map(|records| { + records + .iter() + .filter_map(|record| { + let id = record.get("id").and_then(Value::as_str)?.to_owned(); + let text = + |key: &str| record.get(key).and_then(Value::as_str).map(str::to_owned); + Some(Self { + id, + name: text("name").or_else(|| text("display_name")), + description: text("description"), + capabilities: capabilities(record), + raw: record.clone(), + }) + }) + .collect() + }) + .unwrap_or_default() + } +} diff --git a/crates/tinyinference-video/Cargo.toml b/crates/tinyinference-video/Cargo.toml new file mode 100644 index 0000000..7ae540b --- /dev/null +++ b/crates/tinyinference-video/Cargo.toml @@ -0,0 +1,29 @@ +[package] +name = "tinyinference-video" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +repository.workspace = true +description = "Provider-neutral asynchronous video generation (submit, poll, download) for Rust." +documentation = "https://docs.rs/tinyinference-video" +readme = "../../README.md" +keywords = ["video-generation", "inference", "openrouter", "media"] +categories = ["api-bindings", "asynchronous", "multimedia::video"] + +[dependencies] +async-trait = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } +thiserror = { workspace = true } +tinyinference-image = { version = "0.3.0", path = "../tinyinference-image" } +tokio = { workspace = true } +tracing = { workspace = true } + +[dev-dependencies] +axum = { workspace = true } +tokio = { workspace = true, features = ["net", "time"] } +tokio-test = "0.4" + +[lints] +workspace = true diff --git a/crates/tinyinference-video/examples/live_openrouter_video.rs b/crates/tinyinference-video/examples/live_openrouter_video.rs new file mode 100644 index 0000000..cbbd7ad --- /dev/null +++ b/crates/tinyinference-video/examples/live_openrouter_video.rs @@ -0,0 +1,111 @@ +//! Live smoke test: generate a short video through OpenRouter and save it. +//! +//! Network- and credential-gated, and billed. Reads `OPENROUTER_API_KEY` from +//! the environment, falling back to the workspace `.env` file; exits cleanly +//! (status 0, "skipped") when neither provides one. +//! +//! ```sh +//! cargo run -p tinyinference-video --example live_openrouter_video +//! # image-to-video: pass a first frame (path or URL) +//! LIVE_FIRST_FRAME=target/live-media/image-….png cargo run -p tinyinference-video --example live_openrouter_video +//! # resume a job that timed out, without paying again +//! LIVE_RESUME_JOB=gen-vid-… cargo run -p tinyinference-video --example live_openrouter_video +//! ``` +//! +//! Output lands in `target/live-media/`. + +use std::path::PathBuf; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use tinyinference_video::{ + MediaAuth, MediaReference, OpenRouterVideoGenerator, VideoGenerator, VideoJobStatus, + VideoRequest, WaitPolicy, wait_for_job, +}; + +fn api_key() -> Option { + if let Ok(key) = std::env::var("OPENROUTER_API_KEY") + && !key.trim().is_empty() + { + return Some(key.trim().to_owned()); + } + let env_file = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../../.env"); + std::fs::read_to_string(env_file) + .ok()? + .lines() + .find_map(|line| { + let value = line.trim().strip_prefix("OPENROUTER_API_KEY=")?; + let value = value.trim().trim_matches('"').trim_matches('\''); + (!value.is_empty()).then(|| value.to_owned()) + }) +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let Some(key) = api_key() else { + println!("skipped: OPENROUTER_API_KEY is not set (env or workspace .env)"); + return Ok(()); + }; + let generator = OpenRouterVideoGenerator::new(MediaAuth::ApiKey(key)); + let model = + std::env::var("LIVE_VIDEO_MODEL").unwrap_or_else(|_| generator.default_model().to_owned()); + let started = Instant::now(); + let wait = WaitPolicy::new(Duration::from_secs(5), Duration::from_secs(900)).with_progress( + Arc::new(move |status: &VideoJobStatus| { + println!( + " [{:>5.1}s] {} state={} outputs={}", + started.elapsed().as_secs_f64(), + status.id, + status.state, + status.outputs + ); + }), + ); + + let response = if let Ok(job_id) = std::env::var("LIVE_RESUME_JOB") { + println!("resuming job {job_id}"); + wait_for_job(&generator, &job_id, &model, &wait).await? + } else { + let prompt = std::env::var("LIVE_PROMPT").unwrap_or_else(|_| { + "Anime style: two cheerful engineers shake hands in front of a glowing server as \ + confetti falls, gentle camera push-in" + .to_owned() + }); + let mut request = VideoRequest::new(prompt) + .with_model(&model) + .with_duration(4) + .with_resolution("480p") + .with_aspect_ratio("16:9"); + if let Ok(frame) = std::env::var("LIVE_FIRST_FRAME") { + println!("image-to-video with first frame {frame}"); + request = request.with_first_frame(MediaReference::parse(&frame)); + } + generator.generate(request, &wait).await? + }; + + let out_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../../target/live-media"); + for (index, video) in response.videos.iter().enumerate() { + let path = video + .persist( + &out_dir, + &format!("video-{}-{index}", response.job_id), + "mp4", + ) + .await?; + println!( + "saved {} ({} bytes, {})", + path.display(), + video.data.len(), + video.media_type + ); + } + println!( + "job={} model={} videos={} cost_usd={:?} elapsed={:.1}s", + response.job_id, + response.model, + response.videos.len(), + response.cost_usd, + started.elapsed().as_secs_f64() + ); + Ok(()) +} diff --git a/crates/tinyinference-video/src/error.rs b/crates/tinyinference-video/src/error.rs new file mode 100644 index 0000000..920528d --- /dev/null +++ b/crates/tinyinference-video/src/error.rs @@ -0,0 +1,70 @@ +//! Error type for video generation. + +use thiserror::Error; + +/// Result returned by TinyInference video APIs. +pub type Result = std::result::Result; + +/// A normalized video-generation failure. +/// +/// Video jobs are billed on submit. Every failure that happens *after* a +/// successful submit names the job id and says not to resubmit, because a new +/// submit is a new, separately billed generation; a job that is still running +/// can be resumed with [`crate::wait_for_job`] instead. +#[derive(Debug, Error)] +pub enum Error { + /// A failure before or during submit (validation, capability, auth, + /// transport). Nothing was billed unless the variant says so. + #[error(transparent)] + Media(#[from] tinyinference_image::Error), + /// A failure after the job was accepted (polling or download). + #[error( + "video job {job_id} was accepted and billed, but {stage} failed: {source}; do not resubmit — resume by job id or report this to the user" + )] + Job { + /// Provider job id. + job_id: String, + /// What was being done (`polling`, `downloading output 0`). + stage: String, + /// The underlying failure (boxed to keep `Result` small). + #[source] + source: Box, + }, + /// The provider reported a terminal failure for the job. + #[error( + "video job {job_id} ended as {state}: {message}; it was accepted and billed and cannot be resubmitted — resume by job id or report this to the user" + )] + JobFailed { + /// Provider job id. + job_id: String, + /// Terminal state (`failed`, `cancelled`, `expired`). + state: String, + /// Provider-reported reason, or a placeholder. + message: String, + }, + /// The wait budget elapsed before the job delivered a video. + #[error( + "video job {job_id} did not deliver within {waited_secs}s (last state: {last_state}); it was accepted and billed and may still be running — do not resubmit; resume by job id or report this to the user" + )] + Timeout { + /// Provider job id. + job_id: String, + /// Seconds waited. + waited_secs: u64, + /// Last observed state. + last_state: String, + }, +} + +impl Error { + /// The job id, when the failure happened after a billed submit. + #[must_use] + pub fn job_id(&self) -> Option<&str> { + match self { + Self::Job { job_id, .. } + | Self::JobFailed { job_id, .. } + | Self::Timeout { job_id, .. } => Some(job_id), + Self::Media(_) => None, + } + } +} diff --git a/crates/tinyinference-video/src/job_test.rs b/crates/tinyinference-video/src/job_test.rs new file mode 100644 index 0000000..e8bf78e --- /dev/null +++ b/crates/tinyinference-video/src/job_test.rs @@ -0,0 +1,257 @@ +//! Tests for the submit → poll → download job loop. + +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::time::Duration; + +use async_trait::async_trait; +use tinyinference_image::{GeneratedMedia, MediaModel}; + +use crate::{ + Error, JobState, MockVideoGenerator, MockVideoScript, Result, VideoGenerator, VideoJob, + VideoJobStatus, VideoRequest, WaitPolicy, wait_for_job, +}; + +fn fast(timeout_ms: u64) -> WaitPolicy { + WaitPolicy::new(Duration::from_millis(1), Duration::from_millis(timeout_ms)) +} + +#[tokio::test] +async fn delivers_after_pending_and_in_progress() { + let generator = MockVideoGenerator::new(MockVideoScript::delivers()); + let seen = Arc::new(AtomicUsize::new(0)); + let counter = seen.clone(); + let wait = fast(5_000).with_progress(Arc::new(move |_status: &VideoJobStatus| { + counter.fetch_add(1, Ordering::SeqCst); + })); + let response = generator + .generate(VideoRequest::new("a cat surfing"), &wait) + .await + .unwrap(); + assert_eq!(response.job_id, "mock-job"); + assert_eq!(response.videos.len(), 1); + assert_eq!(response.videos[0].media_type, "video/mp4"); + assert_eq!(seen.load(Ordering::SeqCst), 3); +} + +/// Regression (R1): a `completed` status that reports no outputs yet is not a +/// terminal result — the loop keeps polling until outputs appear. +#[tokio::test] +async fn completed_without_outputs_keeps_polling_until_outputs_appear() { + let generator = MockVideoGenerator::new(MockVideoScript { + polls: vec![ + (JobState::Completed, 0), + (JobState::Completed, 0), + (JobState::Completed, 2), + ], + error: None, + }); + let response = generator + .generate(VideoRequest::new("x"), &fast(5_000)) + .await + .unwrap(); + assert_eq!(response.videos.len(), 2); +} + +/// A generator whose poll script is fixed and whose downloads can be failed. +struct Scripted { + status: (JobState, usize), + content_ok: bool, + poll_error: Option tinyinference_image::Error>, + polls: AtomicUsize, +} + +impl Scripted { + fn new(state: JobState, outputs: usize, content_ok: bool) -> Self { + Self { + status: (state, outputs), + content_ok, + poll_error: None, + polls: AtomicUsize::new(0), + } + } +} + +#[async_trait] +impl VideoGenerator for Scripted { + fn name(&self) -> &str { + "scripted" + } + fn default_model(&self) -> &str { + "scripted/video" + } + async fn submit(&self, _request: VideoRequest) -> Result { + Ok(VideoJob { + id: "job-1".into(), + model: "scripted/video".into(), + state: JobState::Pending, + }) + } + async fn poll(&self, job_id: &str) -> Result { + let call = self.polls.fetch_add(1, Ordering::SeqCst); + if let Some(make) = self.poll_error + && call == 0 + { + return Err(Error::Media(make())); + } + Ok(VideoJobStatus { + id: job_id.into(), + state: self.status.0.clone(), + outputs: self.status.1, + output_indices: Vec::new(), + cost_usd: None, + error: Some("content policy".into()), + }) + } + async fn content(&self, job_id: &str, _index: usize) -> Result { + if self.content_ok { + Ok(GeneratedMedia::new("video/mp4", &b"mp4"[..])) + } else { + Err(Error::Media(tinyinference_image::Error::Http { + status: 404, + message: format!("no content for {job_id}"), + })) + } + } + async fn list_models(&self) -> Result> { + Ok(Vec::new()) + } +} + +/// Regression (R1): a provider that says `completed` but never lists outputs +/// still delivers through one direct download at the deadline. +#[tokio::test(start_paused = true)] +async fn completed_without_listed_outputs_falls_back_to_direct_download() { + let generator = Scripted::new(JobState::Completed, 0, true); + let response = wait_for_job(&generator, "job-1", "m", &fast(100)) + .await + .unwrap(); + assert_eq!(response.videos.len(), 1); +} + +/// Regression (R1/R2): when nothing is ever delivered, the result is a timeout +/// that names the billed job and says not to resubmit — never a success. +#[tokio::test(start_paused = true)] +async fn completed_without_any_output_times_out_naming_the_job() { + let generator = Scripted::new(JobState::Completed, 0, false); + let error = wait_for_job(&generator, "job-1", "m", &fast(100)) + .await + .unwrap_err(); + assert!(matches!(error, Error::Timeout { .. }), "{error:?}"); + assert_eq!(error.job_id(), Some("job-1")); + let message = error.to_string(); + assert!( + message.contains("job-1") && message.contains("do not resubmit"), + "{message}" + ); +} + +#[tokio::test(start_paused = true)] +async fn in_progress_past_the_deadline_times_out() { + let generator = Scripted::new(JobState::InProgress, 0, true); + let error = wait_for_job(&generator, "job-1", "m", &fast(100)) + .await + .unwrap_err(); + match error { + Error::Timeout { last_state, .. } => assert_eq!(last_state, "in_progress"), + other => panic!("expected timeout, got {other:?}"), + } +} + +#[tokio::test] +async fn terminal_failure_reports_state_and_reason() { + let generator = Scripted::new(JobState::Failed, 0, true); + let error = wait_for_job(&generator, "job-1", "m", &fast(5_000)) + .await + .unwrap_err(); + match &error { + Error::JobFailed { + job_id, + state, + message, + } => { + assert_eq!((job_id.as_str(), state.as_str()), ("job-1", "failed")); + assert_eq!(message, "content policy"); + } + other => panic!("expected JobFailed, got {other:?}"), + } +} + +#[tokio::test] +async fn transient_poll_errors_are_retried() { + let mut generator = Scripted::new(JobState::Completed, 1, true); + generator.poll_error = Some(|| tinyinference_image::Error::Http { + status: 503, + message: "busy".into(), + }); + let response = wait_for_job(&generator, "job-1", "m", &fast(5_000)) + .await + .unwrap(); + assert_eq!(response.videos.len(), 1); + assert_eq!(generator.polls.load(Ordering::SeqCst), 2); +} + +/// Regression (R2): a hard failure after submit carries the job id. +#[tokio::test] +async fn hard_poll_errors_name_the_billed_job() { + let mut generator = Scripted::new(JobState::Completed, 1, true); + generator.poll_error = Some(|| tinyinference_image::Error::Http { + status: 404, + message: "unknown job".into(), + }); + let error = wait_for_job(&generator, "job-1", "m", &fast(5_000)) + .await + .unwrap_err(); + assert!( + matches!(error, Error::Job { ref stage, .. } if stage == "polling"), + "{error:?}" + ); + assert_eq!(error.job_id(), Some("job-1")); + assert!(error.to_string().contains("do not resubmit")); +} + +#[tokio::test] +async fn download_failures_name_the_output() { + let generator = Scripted::new(JobState::Completed, 1, false); + let error = wait_for_job(&generator, "job-1", "m", &fast(5_000)) + .await + .unwrap_err(); + assert!( + matches!(error, Error::Job { ref stage, .. } if stage == "downloading output 0"), + "{error:?}" + ); +} + +#[tokio::test] +async fn requests_need_a_prompt_or_an_image() { + let generator = MockVideoGenerator::new(MockVideoScript::delivers()); + let empty = VideoRequest::default(); + assert!(matches!( + generator.generate(empty, &fast(10)).await, + Err(Error::Media(tinyinference_image::Error::Validation(_))) + )); + let image_only = VideoRequest::default().with_first_frame( + tinyinference_image::MediaReference::Url("https://x.test/f.png".into()), + ); + generator.generate(image_only, &fast(5_000)).await.unwrap(); + assert!(matches!( + generator + .generate(VideoRequest::new("x").with_duration(0), &fast(10)) + .await, + Err(Error::Media(tinyinference_image::Error::Validation(_))) + )); +} + +#[test] +fn job_states_parse_provider_spellings() { + assert_eq!(JobState::parse("in_progress"), JobState::InProgress); + assert_eq!(JobState::parse("QUEUED"), JobState::Pending); + assert_eq!(JobState::parse("succeeded"), JobState::Completed); + assert_eq!(JobState::parse("canceled"), JobState::Cancelled); + assert!(JobState::parse("expired").is_terminal_failure()); + assert_eq!( + JobState::parse("warming"), + JobState::Other("warming".into()) + ); + assert!(!JobState::parse("warming").is_terminal_failure()); +} diff --git a/crates/tinyinference-video/src/lib.rs b/crates/tinyinference-video/src/lib.rs new file mode 100644 index 0000000..ba7143b --- /dev/null +++ b/crates/tinyinference-video/src/lib.rs @@ -0,0 +1,316 @@ +//! Provider-neutral asynchronous video generation for TinyInference. +//! +//! Video generation is a job: **submit** (billed), **poll** until the job +//! delivers, then **download** each output. [`VideoGenerator`] exposes the three +//! steps, and [`VideoGenerator::generate`] runs them end to end through +//! [`wait_for_job`], which is also callable on its own to resume a job by id. +//! +//! The wait loop returns only on a *delivered* outcome — `completed` **with** +//! outputs — or a terminal failure. A `completed` status that reports no outputs +//! yet keeps polling instead of ending the job, because providers can flip the +//! status before the artifact exists and a caller told "failed" will resubmit +//! and pay again. Every failure after a successful submit names the job id and +//! says not to resubmit. +//! +//! Reference assets, output-shape normalization and the OpenRouter transport +//! come from [`tinyinference_image`], so image and video generation share one +//! vocabulary and one credential model. + +mod error; +mod mock; +pub mod openrouter; +mod types; + +pub use error::{Error, Result}; +pub use mock::{MockVideoGenerator, MockVideoScript}; +pub use openrouter::{DEFAULT_VIDEO_MODEL, OpenRouterVideoGenerator}; +pub use tinyinference_image::{ + GeneratedMedia, MediaAuth, MediaModel, MediaReference, MediaTransport, ModelCapabilities, +}; +pub use types::{ + JobState, ProgressFn, VideoJob, VideoJobStatus, VideoRequest, VideoResponse, WaitPolicy, +}; + +use tokio::time::Instant; + +use async_trait::async_trait; + +/// A video generation provider. +#[async_trait] +pub trait VideoGenerator: Send + Sync { + /// Short provider name for logs and diagnostics (`"openrouter"`). + fn name(&self) -> &str; + + /// Model used when a request names none. + fn default_model(&self) -> &str; + + /// Submits a job. This is the billed step. + /// + /// # Errors + /// + /// [`Error::Media`] for validation, capability, auth or provider failures. + async fn submit(&self, request: VideoRequest) -> Result; + + /// Reads a job's current state. + /// + /// # Errors + /// + /// [`Error::Media`] for provider or decode failures. + async fn poll(&self, job_id: &str) -> Result; + + /// Downloads output `index` of a completed job. + /// + /// # Errors + /// + /// [`Error::Media`] for provider, size-cap or decode failures. + async fn content(&self, job_id: &str, index: usize) -> Result; + + /// Lists the models this provider can generate with. + /// + /// # Errors + /// + /// Provider or decode errors from the listing endpoint. + async fn list_models(&self) -> Result>; + + /// Submits `request` and waits for delivery under `wait`. + /// + /// # Errors + /// + /// As for [`VideoGenerator::submit`] and [`wait_for_job`]. + async fn generate(&self, request: VideoRequest, wait: &WaitPolicy) -> Result { + let job = self.submit(request).await?; + tracing::info!( + provider = self.name(), + job_id = %job.id, + model = %job.model, + state = %job.state, + "[tinyinference-video] job submitted" + ); + wait_for_job(self, &job.id, &job.model, wait).await + } +} + +/// Waits for job `job_id` to deliver and downloads every output. +/// +/// Use this directly to resume a job that an earlier call submitted and then +/// timed out on, without paying for a new generation. +/// +/// Transient poll failures (rate limits, 5xx, transport) are retried until the +/// deadline. A `completed` job with no outputs keeps polling; at the deadline +/// one direct download of output 0 is attempted before giving up, so a +/// provider that never lists outputs still delivers. +/// +/// # Errors +/// +/// [`Error::JobFailed`] for a terminal failure, [`Error::Timeout`] when the +/// budget elapses, and [`Error::Job`] for a non-transient poll or download +/// failure. All carry the job id. +pub async fn wait_for_job( + generator: &G, + job_id: &str, + model: &str, + wait: &WaitPolicy, +) -> Result { + let started = Instant::now(); + let mut last_state = JobState::Pending; + let mut last_cost_usd: Option = None; + let mut last_poll_error: Option = None; + loop { + let elapsed = started.elapsed(); + if elapsed >= wait.timeout { + return deadline_outcome( + generator, + job_id, + model, + &last_state, + last_cost_usd, + last_poll_error.as_deref(), + elapsed, + ) + .await; + } + + let remaining = wait.timeout - elapsed; + match tokio::time::timeout(remaining, generator.poll(job_id)).await { + Ok(Ok(status)) => { + if let Some(progress) = &wait.progress { + progress(&status); + } + tracing::debug!( + job_id, + state = %status.state, + outputs = status.outputs, + "[tinyinference-video] poll" + ); + last_poll_error = None; + last_cost_usd = status.cost_usd; + if status.is_delivered() { + return download_all( + generator, + job_id, + model, + &status.download_indices(), + status.cost_usd, + started, + wait, + ) + .await; + } + if status.state.is_terminal_failure() { + tracing::warn!(job_id, state = %status.state, "[tinyinference-video] job failed"); + return Err(Error::JobFailed { + job_id: job_id.to_owned(), + state: status.state.to_string(), + message: status + .error + .unwrap_or_else(|| "no reason given by the provider".into()), + }); + } + if status.state == JobState::Completed { + tracing::warn!( + job_id, + "[tinyinference-video] job reports completed with no outputs yet; still polling" + ); + } + last_state = status.state; + } + Ok(Err(Error::Media(error))) if error.is_retryable() => { + tracing::warn!(job_id, %error, "[tinyinference-video] transient poll failure"); + last_poll_error = Some(error.to_string()); + } + Ok(Err(Error::Media(error))) => { + return Err(Error::Job { + job_id: job_id.to_owned(), + stage: "polling".into(), + source: Box::new(error), + }); + } + Ok(Err(other)) => return Err(other), + Err(_) => { + tracing::warn!(job_id, "[tinyinference-video] poll request timeout"); + last_poll_error = Some("poll request timed out".into()); + } + } + + let elapsed = started.elapsed(); + if elapsed >= wait.timeout { + return deadline_outcome( + generator, + job_id, + model, + &last_state, + last_cost_usd, + last_poll_error.as_deref(), + elapsed, + ) + .await; + } + tokio::time::sleep(wait.interval.min(wait.timeout - elapsed)).await; + } +} + +/// Grace allowed for the one direct download attempted when the wait budget +/// runs out on a job that reported `completed` without listing outputs. It is +/// separate from the (already spent) wait budget so the attempt can finish. +const FALLBACK_DOWNLOAD_GRACE: std::time::Duration = std::time::Duration::from_secs(30); + +/// What to return once the wait budget is spent: the output of a `completed` +/// job via one bounded direct download, otherwise [`Error::Timeout`]. +async fn deadline_outcome( + generator: &G, + job_id: &str, + model: &str, + last_state: &JobState, + last_cost_usd: Option, + last_poll_error: Option<&str>, + elapsed: std::time::Duration, +) -> Result { + if *last_state == JobState::Completed + && let Ok(Ok(video)) = + tokio::time::timeout(FALLBACK_DOWNLOAD_GRACE, generator.content(job_id, 0)).await + { + tracing::info!( + job_id, + "[tinyinference-video] completed job without listed outputs delivered on direct download" + ); + return Ok(VideoResponse { + job_id: job_id.to_owned(), + model: model.to_owned(), + videos: vec![video], + cost_usd: last_cost_usd, + }); + } + tracing::warn!( + job_id, + last_state = %last_state, + last_poll_error = ?last_poll_error, + waited_secs = elapsed.as_secs(), + "[tinyinference-video] wait budget elapsed" + ); + Err(Error::Timeout { + job_id: job_id.to_owned(), + waited_secs: elapsed.as_secs(), + last_state: last_state.to_string(), + }) +} + +async fn download_all( + generator: &G, + job_id: &str, + model: &str, + indices: &[usize], + cost_usd: Option, + started: Instant, + wait: &WaitPolicy, +) -> Result { + let mut videos = Vec::with_capacity(indices.len()); + for &index in indices { + let elapsed = started.elapsed(); + if elapsed >= wait.timeout { + return Err(Error::Timeout { + job_id: job_id.to_owned(), + waited_secs: elapsed.as_secs(), + last_state: "downloading".to_owned(), + }); + } + let remaining = wait.timeout - elapsed; + match tokio::time::timeout(remaining, generator.content(job_id, index)).await { + Ok(Ok(video)) => videos.push(video), + Ok(Err(Error::Media(source))) => { + return Err(Error::Job { + job_id: job_id.to_owned(), + stage: format!("downloading output {index}"), + source: Box::new(source), + }); + } + Ok(Err(other)) => return Err(other), + Err(_) => { + return Err(Error::Timeout { + job_id: job_id.to_owned(), + waited_secs: started.elapsed().as_secs(), + last_state: format!("downloading output {index}"), + }); + } + } + } + tracing::info!( + job_id, + videos = videos.len(), + cost_usd, + "[tinyinference-video] job delivered" + ); + Ok(VideoResponse { + job_id: job_id.to_owned(), + model: model.to_owned(), + videos, + cost_usd, + }) +} + +#[cfg(test)] +#[path = "job_test.rs"] +mod job_test; + +#[cfg(test)] +#[path = "openrouter_test.rs"] +mod openrouter_test; diff --git a/crates/tinyinference-video/src/mock.rs b/crates/tinyinference-video/src/mock.rs new file mode 100644 index 0000000..19ba851 --- /dev/null +++ b/crates/tinyinference-video/src/mock.rs @@ -0,0 +1,139 @@ +//! Scripted, offline video generator for tests. + +use std::collections::VecDeque; +use std::sync::Mutex; + +use async_trait::async_trait; +use tinyinference_image::{GeneratedMedia, MediaModel, ModelCapabilities}; + +use crate::types::{JobState, VideoJob, VideoJobStatus, VideoRequest}; +use crate::{Result, VideoGenerator}; + +/// Four bytes of an MP4 `ftyp` box — enough for type sniffing in tests. +const FAKE_MP4: &[u8] = b"\x00\x00\x00\x18ftypmp42"; + +/// A poll script: each poll pops the next state; the last one repeats. +#[derive(Debug, Clone)] +pub struct MockVideoScript { + /// `(state, outputs)` returned by successive polls. + pub polls: Vec<(JobState, usize)>, + /// Error message reported with a terminal failure. + pub error: Option, +} + +impl MockVideoScript { + /// A job that goes pending → in progress → completed with one output. + #[must_use] + pub fn delivers() -> Self { + Self { + polls: vec![ + (JobState::Pending, 0), + (JobState::InProgress, 0), + (JobState::Completed, 1), + ], + error: None, + } + } +} + +/// Replays a [`MockVideoScript`] and records submitted requests. +#[derive(Debug)] +pub struct MockVideoGenerator { + polls: Mutex>, + last: Mutex<(JobState, usize)>, + error: Option, + requests: Mutex>, +} + +impl MockVideoGenerator { + /// Creates a mock that replays `script`. + #[must_use] + pub fn new(script: MockVideoScript) -> Self { + let last = script + .polls + .last() + .cloned() + .unwrap_or((JobState::Completed, 1)); + Self { + polls: Mutex::new(script.polls.into()), + last: Mutex::new(last), + error: script.error, + requests: Mutex::new(Vec::new()), + } + } + + /// Requests submitted so far. + /// + /// # Panics + /// + /// If a previous holder of the internal lock panicked. + #[must_use] + pub fn requests(&self) -> Vec { + self.requests.lock().expect("mock lock poisoned").clone() + } +} + +#[async_trait] +impl VideoGenerator for MockVideoGenerator { + fn name(&self) -> &str { + "mock" + } + + fn default_model(&self) -> &str { + "mock/video" + } + + async fn submit(&self, request: VideoRequest) -> Result { + request.validate()?; + let model = request + .model + .clone() + .unwrap_or_else(|| self.default_model().to_owned()); + self.requests + .lock() + .expect("mock lock poisoned") + .push(request); + Ok(VideoJob { + id: "mock-job".into(), + model, + state: JobState::Pending, + }) + } + + async fn poll(&self, job_id: &str) -> Result { + let next = self.polls.lock().expect("mock lock poisoned").pop_front(); + let (state, outputs) = match next { + Some(step) => { + *self.last.lock().expect("mock lock poisoned") = step.clone(); + step + } + None => self.last.lock().expect("mock lock poisoned").clone(), + }; + let error = state + .is_terminal_failure() + .then(|| self.error.clone()) + .flatten(); + Ok(VideoJobStatus { + id: job_id.to_owned(), + state, + outputs, + output_indices: Vec::new(), + cost_usd: Some(0.0), + error, + }) + } + + async fn content(&self, _job_id: &str, _index: usize) -> Result { + Ok(GeneratedMedia::new("video/mp4", FAKE_MP4)) + } + + async fn list_models(&self) -> Result> { + Ok(vec![MediaModel { + id: self.default_model().to_owned(), + name: Some("Mock video".into()), + description: None, + capabilities: ModelCapabilities::default(), + raw: serde_json::Value::Null, + }]) + } +} diff --git a/crates/tinyinference-video/src/openrouter.rs b/crates/tinyinference-video/src/openrouter.rs new file mode 100644 index 0000000..1321860 --- /dev/null +++ b/crates/tinyinference-video/src/openrouter.rs @@ -0,0 +1,374 @@ +//! OpenRouter video generation (`POST /videos`, `GET /videos/{id}`, +//! `GET /videos/{id}/content`). + +use std::collections::HashMap; + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::{Map, Value, json}; +use tinyinference_image::reference::{ + DEFAULT_MAX_REFERENCE_BYTES, normalize_aspect_ratio, normalize_size, normalize_video_resolution, +}; +use tinyinference_image::transport::wire_model_id; +use tinyinference_image::{ + GeneratedMedia, MediaAuth, MediaModel, MediaTransport, ModelCapabilities, +}; +use tokio::sync::Mutex; + +use crate::types::{JobState, VideoJob, VideoJobStatus, VideoRequest}; +use crate::{Error, Result, VideoGenerator}; + +/// Default video model: Seedance 2.0 Mini (text/image-to-video, first and last +/// frame control, 4–15 s, 480p/720p, optional audio). +pub const DEFAULT_VIDEO_MODEL: &str = "bytedance/seedance-2.0-mini"; + +/// Video generator for OpenRouter's video API, or any backend that proxies it +/// verbatim. +#[derive(Debug)] +pub struct OpenRouterVideoGenerator { + transport: MediaTransport, + default_model: String, + check_capabilities: bool, + max_reference_bytes: usize, + capabilities: Mutex>>, +} + +#[derive(Deserialize)] +struct WireJob { + id: String, + #[serde(default)] + status: Option, + #[serde(default)] + unsigned_urls: Vec, + #[serde(default)] + usage: Option, + #[serde(default)] + error: Option, +} + +#[derive(Deserialize)] +struct WireUsage { + #[serde(default)] + cost: Option, +} + +impl OpenRouterVideoGenerator { + /// Creates a generator against OpenRouter's public API. + #[must_use] + pub fn new(auth: MediaAuth) -> Self { + Self::with_transport(MediaTransport::new(auth)) + } + + /// Creates a generator from a pre-configured transport. + #[must_use] + pub fn with_transport(transport: MediaTransport) -> Self { + Self { + transport, + default_model: DEFAULT_VIDEO_MODEL.to_owned(), + check_capabilities: true, + max_reference_bytes: DEFAULT_MAX_REFERENCE_BYTES, + capabilities: Mutex::new(None), + } + } + + /// Creates a generator from `OPENROUTER_API_KEY`. + /// + /// # Errors + /// + /// [`Error::Media`] wrapping an auth error when the variable is unset. + pub fn from_env() -> Result { + Ok(Self::new(MediaAuth::from_env()?)) + } + + /// Sets the model used when a request names none. + #[must_use] + pub fn with_default_model(mut self, model: impl Into) -> Self { + self.default_model = model.into(); + self + } + + /// Enables or disables pre-flight validation against the model listing. + #[must_use] + pub fn with_capability_check(mut self, enabled: bool) -> Self { + self.check_capabilities = enabled; + self + } + + /// Caps the size of each inlined reference or frame (default 20 MiB). + #[must_use] + pub fn with_max_reference_bytes(mut self, max_reference_bytes: usize) -> Self { + self.max_reference_bytes = max_reference_bytes; + self + } + + /// The underlying transport. + #[must_use] + pub fn transport(&self) -> &MediaTransport { + &self.transport + } + + async fn capabilities_for(&self, model: &str) -> Option { + let mut cache = self.capabilities.lock().await; + if cache.is_none() { + match self.list_models().await { + Ok(models) => { + *cache = Some( + models + .into_iter() + .map(|model| (wire_model_id(&model.id).to_owned(), model.capabilities)) + .collect(), + ); + } + Err(error) => { + tracing::debug!( + %error, + "[tinyinference-video] model listing unavailable; skipping capability check" + ); + return None; + } + } + } + cache.as_ref()?.get(model).cloned() + } +} + +fn validate_against( + model: &str, + body: &Value, + request: &VideoRequest, + caps: &ModelCapabilities, +) -> tinyinference_image::Result<()> { + let field = |key: &str| body.get(key).and_then(Value::as_str); + if let Some(value) = field("resolution") { + ModelCapabilities::check_one_of(model, "resolution", value, caps.resolutions.as_deref())?; + } + if let Some(value) = field("aspect_ratio") { + ModelCapabilities::check_one_of( + model, + "aspect_ratio", + value, + caps.aspect_ratios.as_deref(), + )?; + } + if let (Some(duration), Some(allowed)) = (request.duration_s, caps.durations.as_ref()) { + let allowed: Vec = allowed.iter().map(u32::to_string).collect(); + ModelCapabilities::check_one_of(model, "duration", &duration.to_string(), Some(&allowed))?; + } + for (role, present) in [ + ("first_frame", request.first_frame.is_some()), + ("last_frame", request.last_frame.is_some()), + ] { + if present { + ModelCapabilities::check_one_of( + model, + "frame_images", + role, + caps.frame_images.as_deref(), + )?; + } + } + if request.generate_audio == Some(true) { + ModelCapabilities::check_flag(model, "generate_audio", caps.generate_audio)?; + } + if request.seed.is_some() { + ModelCapabilities::check_flag(model, "seed", caps.seed)?; + } + Ok(()) +} + +/// Builds the wire body for `request`. +/// +/// # Errors +/// +/// Reference resolution errors from [`tinyinference_image::MediaReference::resolve`]. +pub async fn build_video_body( + model: &str, + request: &VideoRequest, + max_reference_bytes: usize, +) -> Result { + let mut body = Map::new(); + body.insert("model".into(), json!(model)); + if let Some(prompt) = request.prompt.as_deref().filter(|p| !p.trim().is_empty()) { + body.insert("prompt".into(), json!(prompt)); + } + if let Some(duration) = request.duration_s { + body.insert("duration".into(), json!(duration)); + } + let normalized = |value: &Option, normalize: fn(&str) -> Option| { + value + .as_deref() + .map(|raw| normalize(raw).unwrap_or_else(|| raw.trim().to_owned())) + }; + if let Some(resolution) = normalized(&request.resolution, normalize_video_resolution) { + body.insert("resolution".into(), json!(resolution)); + } + if let Some(aspect_ratio) = normalized(&request.aspect_ratio, normalize_aspect_ratio) { + body.insert("aspect_ratio".into(), json!(aspect_ratio)); + } + if let Some(size) = normalized(&request.size, normalize_size) { + body.insert("size".into(), json!(size)); + } + if let Some(generate_audio) = request.generate_audio { + body.insert("generate_audio".into(), json!(generate_audio)); + } + if let Some(seed) = request.seed { + body.insert("seed".into(), json!(seed)); + } + let mut frames = Vec::new(); + for (role, frame) in [ + ("first_frame", &request.first_frame), + ("last_frame", &request.last_frame), + ] { + if let Some(frame) = frame { + let url = frame.resolve(max_reference_bytes).await?; + frames.push(json!({ + "type": "image_url", + "image_url": { "url": url }, + "frame_type": role, + })); + } + } + if !frames.is_empty() { + body.insert("frame_images".into(), Value::Array(frames)); + } + if !request.references.is_empty() { + let mut parts = Vec::with_capacity(request.references.len()); + for reference in &request.references { + parts.push(reference.to_content_part(max_reference_bytes).await?); + } + body.insert("input_references".into(), Value::Array(parts)); + } + for (key, value) in [ + ("previous_job_id", &request.previous_job_id), + ("user", &request.user), + ("session_id", &request.session_id), + ] { + if let Some(value) = value { + body.insert(key.into(), json!(value)); + } + } + for (key, value) in &request.extra { + body.entry(key.clone()).or_insert_with(|| value.clone()); + } + Ok(Value::Object(body)) +} + +/// Refuses job ids that could escape the `videos/{id}` path segment. +fn checked_job_id(job_id: &str) -> tinyinference_image::Result<&str> { + let valid = !job_id.is_empty() + && job_id.len() <= 256 + && job_id + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'); + if valid { + Ok(job_id) + } else { + Err(tinyinference_image::Error::Validation(format!( + "invalid video job id {job_id:?}" + ))) + } +} + +fn error_text(error: Option) -> Option { + match error? { + Value::String(text) => Some(text), + Value::Null => None, + other => other + .pointer("/message") + .and_then(Value::as_str) + .map(str::to_owned) + .or_else(|| Some(other.to_string())), + } +} + +#[async_trait] +impl VideoGenerator for OpenRouterVideoGenerator { + fn name(&self) -> &str { + "openrouter" + } + + fn default_model(&self) -> &str { + &self.default_model + } + + async fn submit(&self, request: VideoRequest) -> Result { + request.validate()?; + let model = + wire_model_id(request.model.as_deref().unwrap_or(&self.default_model)).to_owned(); + let body = build_video_body(&model, &request, self.max_reference_bytes).await?; + if self.check_capabilities + && let Some(caps) = self.capabilities_for(&model).await + { + validate_against(&model, &body, &request, &caps)?; + } + tracing::info!( + model = %model, + duration_s = request.duration_s, + first_frame = request.first_frame.is_some(), + last_frame = request.last_frame.is_some(), + references = request.references.len(), + base_url = %self.transport.redacted_base_url(), + "[tinyinference-video] submitting job" + ); + let job: WireJob = self.transport.post_json("videos", &body).await?; + checked_job_id(&job.id)?; + Ok(VideoJob { + id: job.id, + model, + state: JobState::parse(job.status.as_deref().unwrap_or("pending")), + }) + } + + async fn poll(&self, job_id: &str) -> Result { + let job_id = checked_job_id(job_id)?; + let job: WireJob = self.transport.get_json(&format!("videos/{job_id}")).await?; + let output_indices: Vec = job + .unsigned_urls + .iter() + .enumerate() + .filter(|(_, url)| !url.trim().is_empty()) + .map(|(index, _)| index) + .collect(); + Ok(VideoJobStatus { + id: job.id, + state: JobState::parse(job.status.as_deref().unwrap_or("pending")), + outputs: output_indices.len(), + output_indices, + cost_usd: job.usage.and_then(|usage| usage.cost), + error: error_text(job.error), + }) + } + + async fn content(&self, job_id: &str, index: usize) -> Result { + let job_id = checked_job_id(job_id)?; + let (data, content_type) = self + .transport + .get_bytes(&format!("videos/{job_id}/content?index={index}")) + .await?; + if content_type.as_deref().is_some_and(|value| { + let value = value.trim().to_ascii_lowercase(); + value.starts_with("application/json") || value.contains("+json") + }) { + return Err(Error::Media(tinyinference_image::Error::NoMedia { + request_id: Some(job_id.to_owned()), + })); + } + if data.is_empty() { + return Err(Error::Media(tinyinference_image::Error::NoMedia { + request_id: Some(job_id.to_owned()), + })); + } + let media_type = content_type + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| "video/mp4".to_owned()); + Ok(GeneratedMedia::new(media_type, data)) + } + + async fn list_models(&self) -> Result> { + let body: Value = self.transport.get_json("videos/models").await?; + Ok(MediaModel::parse_listing( + &body, + ModelCapabilities::from_video_model, + )) + } +} diff --git a/crates/tinyinference-video/src/openrouter_test.rs b/crates/tinyinference-video/src/openrouter_test.rs new file mode 100644 index 0000000..8390715 --- /dev/null +++ b/crates/tinyinference-video/src/openrouter_test.rs @@ -0,0 +1,300 @@ +//! Offline tests for the OpenRouter video generator. + +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use axum::Router; +use axum::extract::{Path, Query, Request, State}; +use axum::http::StatusCode; +use axum::response::IntoResponse; +use axum::routing::{get, post}; +use serde_json::{Value, json}; +use tinyinference_image::{MediaAuth, MediaReference, MediaTransport}; + +use crate::{Error, OpenRouterVideoGenerator, VideoGenerator, VideoRequest, WaitPolicy}; + +#[derive(Clone)] +struct Server { + submits: Arc>>, + polls: Arc, + /// Poll replies, in order; the last one repeats. + script: Arc>, + listing: Option, + content_queries: Arc>>, + /// Wrap JSON replies in the proxying backend's `{success, data}` envelope. + envelope: bool, +} + +fn wrap(envelope: bool, body: Value) -> Value { + if envelope { + json!({ "success": true, "data": body }) + } else { + body + } +} + +async fn start( + prefix: &str, + script: Vec<(StatusCode, Value)>, + listing: Option, +) -> (String, Server) { + let server = Server { + submits: Arc::default(), + polls: Arc::default(), + script: Arc::new(script), + listing, + content_queries: Arc::default(), + envelope: prefix.starts_with("/agent-integrations"), + }; + let router = Router::new() + .route( + &format!("{prefix}/videos"), + post(|State(s): State, request: Request| async move { + let body = axum::body::to_bytes(request.into_body(), usize::MAX).await.unwrap(); + s.submits.lock().unwrap().push(serde_json::from_slice(&body).unwrap()); + axum::Json(wrap( + s.envelope, + json!({ + "id": "gen-vid-1-abc", "polling_url": "/api/v1/videos/gen-vid-1-abc", "status": "pending" + }), + )) + }), + ) + .route( + &format!("{prefix}/videos/models"), + get(|State(s): State| async move { + match s.listing { + Some(listing) => axum::Json(listing).into_response(), + None => StatusCode::NOT_FOUND.into_response(), + } + }), + ) + .route( + &format!("{prefix}/videos/{{id}}"), + get(|State(s): State, Path(id): Path| async move { + let call = s.polls.fetch_add(1, Ordering::SeqCst); + let (status, mut body) = s.script[call.min(s.script.len() - 1)].clone(); + body["id"] = json!(id); + let body = if status.is_success() { wrap(s.envelope, body) } else { body }; + (status, [("retry-after", "0")], axum::Json(body)).into_response() + }), + ) + .route( + &format!("{prefix}/videos/{{id}}/content"), + get( + |State(s): State, Query(q): Query>| async move { + s.content_queries + .lock() + .unwrap() + .push(q.get("index").cloned().unwrap_or_default()); + ([("content-type", "video/mp4")], &b"\x00\x00\x00\x18ftypmp42"[..]).into_response() + }, + ), + ) + .with_state(server.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, router).await.unwrap() }); + (format!("http://{address}{prefix}"), server) +} + +fn generator(base_url: &str) -> OpenRouterVideoGenerator { + OpenRouterVideoGenerator::with_transport( + MediaTransport::new(MediaAuth::ApiKey("sk-or-test-key-123456".into())) + .with_base_url(base_url) + .with_max_retries(2), + ) +} + +fn fast() -> WaitPolicy { + WaitPolicy::new(Duration::from_millis(1), Duration::from_secs(5)) +} + +fn poll(status: &str, urls: &[&str]) -> (StatusCode, Value) { + ( + StatusCode::OK, + json!({ "status": status, "unsigned_urls": urls, "usage": { "cost": 0.42 } }), + ) +} + +/// Regression (R1) end to end over HTTP: `completed` with an empty +/// `unsigned_urls` is polled through, then both outputs are downloaded. +#[tokio::test] +async fn full_lifecycle_polls_through_completed_without_urls() { + let (base, server) = start( + "/api/v1", + vec![ + poll("pending", &[]), + poll("in_progress", &[]), + poll("completed", &[]), + poll( + "completed", + &["https://cdn.test/0.mp4", "https://cdn.test/1.mp4"], + ), + ], + None, + ) + .await; + let response = generator(&base) + .generate( + VideoRequest::new("a lighthouse in a storm") + .with_model("openrouter/bytedance/seedance-2.0-mini") + .with_duration(5) + .with_resolution("720") + .with_aspect_ratio("landscape") + .with_audio(true) + .with_first_frame(MediaReference::Url("https://x.test/first.png".into())) + .with_last_frame(MediaReference::DataUrl( + "data:image/png;base64,iVBORw0KGgo=".into(), + )) + .with_reference(MediaReference::Url("https://x.test/style.mp4".into())), + &fast(), + ) + .await + .unwrap(); + + assert_eq!(response.job_id, "gen-vid-1-abc"); + assert_eq!(response.model, "bytedance/seedance-2.0-mini"); + assert_eq!(response.videos.len(), 2); + assert_eq!(response.cost_usd, Some(0.42)); + assert_eq!(server.polls.load(Ordering::SeqCst), 4); + assert_eq!(*server.content_queries.lock().unwrap(), vec!["0", "1"]); + + let body = server.submits.lock().unwrap()[0].clone(); + assert_eq!(body["model"], "bytedance/seedance-2.0-mini"); + assert_eq!(body["duration"], 5); + assert_eq!(body["resolution"], "720p"); + assert_eq!(body["aspect_ratio"], "16:9"); + assert_eq!(body["generate_audio"], true); + assert_eq!(body["frame_images"][0]["frame_type"], "first_frame"); + assert_eq!( + body["frame_images"][0]["image_url"]["url"], + "https://x.test/first.png" + ); + assert_eq!(body["frame_images"][1]["frame_type"], "last_frame"); + assert_eq!(body["input_references"][0]["type"], "video_url"); + assert_eq!( + body["input_references"][0]["video_url"]["url"], + "https://x.test/style.mp4" + ); +} + +#[tokio::test] +async fn proxied_base_url_unwraps_the_backend_envelope() { + let (base, server) = start( + "/agent-integrations/openrouter", + vec![poll("pending", &[]), poll("completed", &["u"])], + None, + ) + .await; + let response = generator(&base) + .generate(VideoRequest::new("x"), &fast()) + .await + .unwrap(); + assert_eq!(response.job_id, "gen-vid-1-abc"); + assert_eq!(response.videos.len(), 1); + assert_eq!(response.cost_usd, Some(0.42)); + assert_eq!(server.submits.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn failed_job_surfaces_the_provider_error() { + let (base, _server) = start( + "/api/v1", + vec![( + StatusCode::OK, + json!({ "status": "failed", "error": "prompt rejected by safety filter" }), + )], + None, + ) + .await; + let error = generator(&base) + .generate(VideoRequest::new("x"), &fast()) + .await + .unwrap_err(); + match error { + Error::JobFailed { message, .. } => assert_eq!(message, "prompt rejected by safety filter"), + other => panic!("expected JobFailed, got {other:?}"), + } +} + +#[tokio::test] +async fn transient_poll_server_errors_are_retried() { + let (base, server) = start( + "/api/v1", + vec![ + ( + StatusCode::BAD_GATEWAY, + json!({"error": {"message": "busy"}}), + ), + poll("completed", &["u"]), + ], + None, + ) + .await; + generator(&base) + .generate(VideoRequest::new("x"), &fast()) + .await + .unwrap(); + assert_eq!(server.polls.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn unsupported_duration_fails_before_submit() { + let listing = json!({ "data": [{ + "id": "bytedance/seedance-2.0-mini", + "supported_durations": [4, 5, 6], + "supported_resolutions": ["480p", "720p"], + "supported_frame_images": ["first_frame", "last_frame"], + "generate_audio": true, + "seed": true + }]}); + let (base, server) = start("/api/v1", vec![poll("completed", &["u"])], Some(listing)).await; + let error = generator(&base) + .submit(VideoRequest::new("x").with_duration(20)) + .await + .unwrap_err(); + assert!( + matches!(&error, Error::Media(tinyinference_image::Error::Unsupported { field, .. }) if field == "duration"), + "{error:?}" + ); + let error = generator(&base) + .submit(VideoRequest::new("x").with_resolution("1080p")) + .await + .unwrap_err(); + assert!( + matches!(&error, Error::Media(tinyinference_image::Error::Unsupported { field, .. }) if field == "resolution"), + "{error:?}" + ); + assert!(server.submits.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn hostile_job_ids_are_rejected_locally() { + let generator = generator("http://127.0.0.1:9"); + for id in ["../admin", "a/b", "", "x?y=1"] { + assert!(matches!( + generator.poll(id).await, + Err(Error::Media(tinyinference_image::Error::Validation(_))) + )); + } +} + +/// Sparse `unsigned_urls` (a blank slot before a populated one) download the +/// populated slot's index, not `0`. +#[tokio::test] +async fn sparse_output_slots_download_their_own_index() { + let (base, server) = start( + "/api/v1", + vec![poll("completed", &["", "https://cdn.test/1.mp4"])], + None, + ) + .await; + let response = generator(&base) + .generate(VideoRequest::new("x"), &fast()) + .await + .unwrap(); + assert_eq!(response.videos.len(), 1); + assert_eq!(*server.content_queries.lock().unwrap(), vec!["1"]); +} diff --git a/crates/tinyinference-video/src/types.rs b/crates/tinyinference-video/src/types.rs new file mode 100644 index 0000000..f05ff44 --- /dev/null +++ b/crates/tinyinference-video/src/types.rs @@ -0,0 +1,319 @@ +//! Public request, job, and response types for video generation. + +use std::sync::Arc; +use std::time::Duration; + +use serde_json::Value; +use tinyinference_image::{GeneratedMedia, MediaReference}; + +use crate::Result; + +/// A provider-neutral video generation request. +/// +/// Output-shape fields accept loose spellings (`"720"`, `"full hd"`, +/// `"landscape"`); providers normalize them with +/// [`tinyinference_image::reference`] and forward anything unrecognized. +#[derive(Debug, Clone, Default, PartialEq)] +pub struct VideoRequest { + /// Model id; `None` uses the generator's default. An `openrouter/` prefix + /// is accepted and stripped. + pub model: Option, + /// Text prompt. Optional only when a frame image or reference drives the + /// generation on its own. + pub prompt: Option, + /// Clip duration in seconds. + pub duration_s: Option, + /// Output resolution (`480p`, `720p`, `1080p`, `4K`, …). + pub resolution: Option, + /// Aspect ratio (`16:9`, `9:16`, …). + pub aspect_ratio: Option, + /// Exact pixel size (`1280x720`); interchangeable with resolution + + /// aspect ratio. + pub size: Option, + /// Whether to generate an audio track, where supported. + pub generate_audio: Option, + /// Deterministic seed, where supported. + pub seed: Option, + /// Image to use as the first frame (image-to-video). + pub first_frame: Option, + /// Image to use as the last frame. + pub last_frame: Option, + /// Reference assets (image, video or audio) guiding subject or style. + pub references: Vec, + /// A completed job to edit or extend, for models that support it. + pub previous_job_id: Option, + /// Stable end-user identifier for provider abuse detection. + pub user: Option, + /// Observability grouping id (never sent to the upstream model provider). + pub session_id: Option, + /// Provider-specific extra fields merged into the wire body. + pub extra: serde_json::Map, +} + +impl VideoRequest { + /// Creates a text-to-video request. + #[must_use] + pub fn new(prompt: impl Into) -> Self { + Self { + prompt: Some(prompt.into()), + ..Self::default() + } + } + + /// Sets the model id. + #[must_use] + pub fn with_model(mut self, model: impl Into) -> Self { + self.model = Some(model.into()); + self + } + + /// Sets the duration in seconds. + #[must_use] + pub fn with_duration(mut self, seconds: u32) -> Self { + self.duration_s = Some(seconds); + self + } + + /// Sets the resolution. + #[must_use] + pub fn with_resolution(mut self, resolution: impl Into) -> Self { + self.resolution = Some(resolution.into()); + self + } + + /// Sets the aspect ratio. + #[must_use] + pub fn with_aspect_ratio(mut self, aspect_ratio: impl Into) -> Self { + self.aspect_ratio = Some(aspect_ratio.into()); + self + } + + /// Sets audio generation. + #[must_use] + pub fn with_audio(mut self, generate_audio: bool) -> Self { + self.generate_audio = Some(generate_audio); + self + } + + /// Sets the first frame. + #[must_use] + pub fn with_first_frame(mut self, frame: MediaReference) -> Self { + self.first_frame = Some(frame); + self + } + + /// Sets the last frame. + #[must_use] + pub fn with_last_frame(mut self, frame: MediaReference) -> Self { + self.last_frame = Some(frame); + self + } + + /// Adds a reference asset. + #[must_use] + pub fn with_reference(mut self, reference: MediaReference) -> Self { + self.references.push(reference); + self + } + + /// Checks the fields that are invalid for every model. + /// + /// # Errors + /// + /// [`tinyinference_image::Error::Validation`] when there is neither a + /// prompt nor an image input, or the duration is zero. + pub fn validate(&self) -> Result<()> { + let has_prompt = self.prompt.as_deref().is_some_and(|p| !p.trim().is_empty()); + let has_image_input = + self.first_frame.is_some() || self.last_frame.is_some() || !self.references.is_empty(); + if !has_prompt && !has_image_input { + return Err(validation("a prompt or an input image is required")); + } + if self.duration_s == Some(0) { + return Err(validation("duration must be at least 1 second")); + } + Ok(()) + } +} + +fn validation(message: &str) -> crate::Error { + crate::Error::Media(tinyinference_image::Error::Validation(message.to_owned())) +} + +/// Lifecycle state of a video job. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum JobState { + /// Queued, not started. + Pending, + /// Generating. + InProgress, + /// Finished; output should be downloadable. + Completed, + /// Terminal failure. + Failed, + /// Cancelled before completion. + Cancelled, + /// Output expired before it was collected. + Expired, + /// A state this crate does not know; treated as in-flight. + Other(String), +} + +impl JobState { + /// Parses a provider status string (case-insensitive). + #[must_use] + pub fn parse(value: &str) -> Self { + match value.trim().to_ascii_lowercase().as_str() { + "pending" | "queued" => Self::Pending, + "in_progress" | "processing" | "running" => Self::InProgress, + "completed" | "succeeded" | "success" => Self::Completed, + "failed" | "error" => Self::Failed, + "cancelled" | "canceled" => Self::Cancelled, + "expired" => Self::Expired, + other => Self::Other(other.to_owned()), + } + } + + /// Whether the job can no longer produce output. + #[must_use] + pub fn is_terminal_failure(&self) -> bool { + matches!(self, Self::Failed | Self::Cancelled | Self::Expired) + } + + /// The canonical lowercase name. + #[must_use] + pub fn as_str(&self) -> &str { + match self { + Self::Pending => "pending", + Self::InProgress => "in_progress", + Self::Completed => "completed", + Self::Failed => "failed", + Self::Cancelled => "cancelled", + Self::Expired => "expired", + Self::Other(other) => other, + } + } +} + +impl std::fmt::Display for JobState { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(self.as_str()) + } +} + +/// A submitted job. +#[derive(Debug, Clone, PartialEq)] +pub struct VideoJob { + /// Provider job id; use it to poll, download, or resume. + pub id: String, + /// Wire model id the job runs on. + pub model: String, + /// State reported at submit. + pub state: JobState, +} + +/// One poll result. +#[derive(Debug, Clone, PartialEq)] +pub struct VideoJobStatus { + /// Provider job id. + pub id: String, + /// Current state. + pub state: JobState, + /// How many outputs the provider reports as ready (`unsigned_urls`). + pub outputs: usize, + /// Provider slot index of each ready output. Slots can be sparse (a blank + /// entry between populated ones), so downloads use these indices rather + /// than `0..outputs`. + pub output_indices: Vec, + /// Provider-reported cost in USD, once known. + pub cost_usd: Option, + /// Provider-reported error, for failed jobs. + pub error: Option, +} + +impl VideoJobStatus { + /// Whether the job finished *and* has output to download. + /// + /// A `completed` status with no outputs is not delivered: providers can + /// flip the status before the artifact is materialized, and treating that + /// as terminal turns a paid, about-to-deliver job into a false failure. + #[must_use] + pub fn is_delivered(&self) -> bool { + self.state == JobState::Completed && self.outputs > 0 + } + + /// The provider slots to download: `output_indices` when known, else the + /// dense range `0..outputs`. + #[must_use] + pub fn download_indices(&self) -> Vec { + if self.output_indices.is_empty() { + (0..self.outputs).collect() + } else { + self.output_indices.clone() + } + } +} + +/// Called with every poll result, for host progress reporting. +pub type ProgressFn = Arc; + +/// How long and how often to wait for a job. +#[derive(Clone)] +pub struct WaitPolicy { + /// Delay between polls. + pub interval: Duration, + /// Total wait budget measured from the first poll. + pub timeout: Duration, + /// Optional progress observer. + pub progress: Option, +} + +impl WaitPolicy { + /// Polls every `interval` for at most `timeout`. + #[must_use] + pub fn new(interval: Duration, timeout: Duration) -> Self { + Self { + interval, + timeout, + progress: None, + } + } + + /// Installs a progress observer. + #[must_use] + pub fn with_progress(mut self, progress: ProgressFn) -> Self { + self.progress = Some(progress); + self + } +} + +impl Default for WaitPolicy { + /// Polls every 5 seconds for up to 10 minutes. + fn default() -> Self { + Self::new(Duration::from_secs(5), Duration::from_secs(600)) + } +} + +impl std::fmt::Debug for WaitPolicy { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("WaitPolicy") + .field("interval", &self.interval) + .field("timeout", &self.timeout) + .field("progress", &self.progress.is_some()) + .finish() + } +} + +/// A delivered video generation. Always carries at least one video. +#[derive(Debug, Clone, PartialEq)] +pub struct VideoResponse { + /// Provider job id. + pub job_id: String, + /// Wire model id. + pub model: String, + /// Downloaded videos, in provider order. + pub videos: Vec, + /// Provider-reported cost in USD, when available. + pub cost_usd: Option, +}