use std::future::Future;
use std::sync::Arc;
use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use url::Url;
use crate::dynamic::BoxFuture;
use crate::error::ProviderError;
use crate::image_model::AspectRatio;
use crate::image_model::ImageSize;
use crate::json::JsonValue;
use crate::language_model::ResponseMetadata;
use crate::shared::FileData;
use crate::shared::Headers;
use crate::shared::MediaType;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
pub trait VideoModel: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn model_id(&self) -> &ModelId;
fn max_videos_per_call(&self) -> Option<usize>;
fn supports_generate(&self) -> bool {
false
}
fn do_generate(
&self,
options: VideoOptions,
) -> impl Future<Output = Result<VideoResult, ProviderError>> + Send {
let _ = options;
std::future::ready(Err(ProviderError::unsupported(
"synchronous video generation",
)))
}
fn supports_operations(&self) -> bool {
false
}
fn do_start(
&self,
options: VideoStartOptions,
) -> impl Future<Output = Result<VideoStartResult, ProviderError>> + Send {
let _ = options;
std::future::ready(Err(ProviderError::unsupported(
"asynchronous video generation",
)))
}
fn do_status(
&self,
options: VideoStatusOptions,
) -> impl Future<Output = Result<VideoStatusResult, ProviderError>> + Send {
let _ = options;
std::future::ready(Err(ProviderError::unsupported(
"asynchronous video generation",
)))
}
fn supports_webhook(&self) -> bool {
false
}
fn handle_webhook(
&self,
factory: WebhookFactory,
) -> impl Future<Output = Result<WebhookHandle, ProviderError>> + Send {
factory()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(untagged)]
pub enum VideoAspectRatio {
Ratio(AspectRatio),
#[serde(with = "adaptive")]
Adaptive,
}
mod adaptive {
pub(super) fn serialize<S: serde::Serializer>(serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str("adaptive")
}
pub(super) fn deserialize<'de, D: serde::Deserializer<'de>>(
deserializer: D,
) -> Result<(), D::Error> {
let text = <std::borrow::Cow<'de, str> as serde::Deserialize>::deserialize(deserializer)?;
if text == "adaptive" {
Ok(())
} else {
Err(serde::de::Error::custom("expected `adaptive`"))
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FrameType {
FirstFrame,
LastFrame,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct VideoFile {
pub data: FileData,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub media_type: Option<MediaType>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_options: Option<ProviderOptions>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct FrameImage {
pub image: VideoFile,
pub frame_type: FrameType,
}
#[derive(Debug, Clone)]
pub struct VideoOptions {
pub prompt: Option<String>,
pub n: u32,
pub aspect_ratio: Option<VideoAspectRatio>,
pub resolution: Option<ImageSize>,
pub duration: Option<f64>,
pub fps: Option<u32>,
pub seed: Option<u64>,
pub image: Option<VideoFile>,
pub frame_images: Vec<FrameImage>,
pub input_references: Vec<VideoFile>,
pub generate_audio: Option<bool>,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
impl VideoOptions {
#[must_use]
pub fn new(prompt: impl Into<String>) -> Self {
Self {
prompt: Some(prompt.into()),
..Self::default()
}
}
}
impl Default for VideoOptions {
fn default() -> Self {
Self {
prompt: None,
n: 1,
aspect_ratio: None,
resolution: None,
duration: None,
fps: None,
seed: None,
image: None,
frame_images: Vec::new(),
input_references: Vec::new(),
generate_audio: None,
provider_options: ProviderOptions::new(),
headers: Headers::new(),
cancellation: CancellationToken::new(),
}
}
}
#[derive(Debug, Clone)]
pub struct VideoStartOptions {
pub options: VideoOptions,
pub webhook_url: Option<Url>,
}
#[derive(Debug, Clone)]
pub struct VideoStatusOptions {
pub operation: JsonValue,
pub headers: Headers,
pub cancellation: CancellationToken,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct VideoData {
pub data: FileData,
pub media_type: MediaType,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VideoResult {
pub videos: Vec<VideoData>,
#[serde(default)]
pub warnings: Vec<Warning>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
pub response: ResponseMetadata,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VideoStartResult {
pub operation: JsonValue,
#[serde(default)]
pub warnings: Vec<Warning>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
pub response: ResponseMetadata,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "status", rename_all = "lowercase")]
#[non_exhaustive]
pub enum VideoStatusResult {
Pending {
#[serde(default)]
warnings: Vec<Warning>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
response: ResponseMetadata,
},
Completed {
videos: Vec<VideoData>,
#[serde(default)]
warnings: Vec<Warning>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
response: ResponseMetadata,
},
Error {
error: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
response: ResponseMetadata,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct WebhookPayload {
pub headers: Headers,
pub body: JsonValue,
}
pub struct WebhookHandle {
pub url: Url,
pub received: BoxFuture<'static, Result<WebhookPayload, ProviderError>>,
}
impl std::fmt::Debug for WebhookHandle {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebhookHandle")
.field("url", &self.url)
.field("received", &"<future>")
.finish()
}
}
pub type WebhookFactory =
Arc<dyn Fn() -> BoxFuture<'static, Result<WebhookHandle, ProviderError>> + Send + Sync>;