Skip to main content

ferrin_core/
video.rs

1//! Video generation (feature `video`): [`generate_video`] supports the
2//! synchronous flow (`do_generate`) and the asynchronous flow
3//! (`do_start`/`do_status`, with polling or a webhook).
4//!
5//! Design: `docs/01-architecture/11-other-modalities.md` ยง6.
6
7use std::fmt;
8use std::future::IntoFuture;
9use std::sync::Arc;
10use std::time::Duration;
11
12use bytes::Bytes;
13use ferrin_provider_util::ids::generate_id;
14use ferrin_provider_util::media_type::detect_media_type_for;
15use ferrin_spec::BoxFuture;
16use ferrin_spec::DynVideoModel;
17use ferrin_spec::FileData;
18use ferrin_spec::ImageSize;
19use ferrin_spec::MediaType;
20use ferrin_spec::ProviderMetadata;
21use ferrin_spec::ResponseMetadata;
22use ferrin_spec::VideoModelRef;
23use ferrin_spec::Warning;
24use ferrin_spec::error::ProviderError;
25pub use ferrin_spec::video_model::FrameImage;
26pub use ferrin_spec::video_model::FrameType;
27pub use ferrin_spec::video_model::VideoAspectRatio;
28pub use ferrin_spec::video_model::VideoFile;
29use ferrin_spec::video_model::VideoOptions;
30use ferrin_spec::video_model::VideoResult;
31use ferrin_spec::video_model::VideoStartOptions;
32use ferrin_spec::video_model::VideoStatusOptions;
33use ferrin_spec::video_model::VideoStatusResult;
34pub use ferrin_spec::video_model::WebhookFactory;
35pub use ferrin_spec::video_model::WebhookHandle;
36pub use ferrin_spec::video_model::WebhookPayload;
37use tokio::task::JoinSet;
38use tokio::time::Instant;
39use tokio_util::sync::CancellationToken;
40use tracing::Instrument;
41
42use crate::error::Error;
43use crate::modality::ModalityOptions;
44use crate::modality::impl_modality_builder;
45use crate::modality::merge_provider_metadata;
46use crate::prompt::DefaultDownloader;
47use crate::prompt::DownloadFn;
48use crate::prompt::DownloadRequest;
49use crate::registry::ProviderRegistry;
50use crate::registry::default::resolve_model;
51use crate::retry::RetryPolicy;
52use crate::retry::retry;
53use crate::telemetry::ModelIdentity;
54use crate::telemetry::spans;
55use crate::timeout::TimeoutScope;
56
57/// Media type used when nothing better is known.
58const DEFAULT_VIDEO_MEDIA_TYPE: &str = "video/mp4";
59
60/// Media type treated as unknown.
61const GENERIC_MEDIA_TYPE: &str = "application/octet-stream";
62
63/// Header carrying the idempotency key of `do_start`.
64const IDEMPOTENCY_KEY: &str = "idempotency-key";
65
66/// Polling configuration of the asynchronous flow.
67#[derive(Debug, Clone, PartialEq, Eq)]
68pub struct PollConfig {
69    /// Delay between status requests (default 5 s).
70    pub interval: Duration,
71    /// Maximum time to wait for completion, polling or via webhook
72    /// (default 10 min).
73    pub timeout: Duration,
74    /// Maximum number of status requests (default unlimited).
75    pub max_attempts: Option<u32>,
76}
77
78impl Default for PollConfig {
79    fn default() -> Self {
80        Self {
81            interval: Duration::from_secs(5),
82            timeout: Duration::from_secs(600),
83            max_attempts: None,
84        }
85    }
86}
87
88/// A generated video.
89#[derive(Debug, Clone, PartialEq, Eq)]
90pub struct GeneratedVideo {
91    /// Video bytes.
92    pub data: Bytes,
93    /// Media type (reported, downloaded, detected, or `video/mp4`).
94    pub media_type: MediaType,
95}
96
97/// Result of [`generate_video`].
98#[derive(Debug, Clone, PartialEq)]
99pub struct GenerateVideoResult {
100    /// All videos, in call order.
101    pub videos: Vec<GeneratedVideo>,
102    /// Warnings of all calls.
103    pub warnings: Vec<Warning>,
104    /// Response metadata of all calls.
105    pub responses: Vec<ResponseMetadata>,
106    /// Provider metadata merged over all calls.
107    pub provider_metadata: ProviderMetadata,
108}
109
110impl GenerateVideoResult {
111    /// The first video.
112    #[must_use]
113    pub fn video(&self) -> Option<&GeneratedVideo> {
114        self.videos.first()
115    }
116}
117
118/// Generates videos from a text prompt.
119#[must_use]
120pub fn generate_video(model: impl Into<VideoModelRef>, prompt: impl Into<String>) -> GenerateVideo {
121    GenerateVideo {
122        model: model.into(),
123        prompt: Some(prompt.into()),
124        n: 1,
125        max_videos_per_call: None,
126        aspect_ratio: None,
127        resolution: None,
128        duration: None,
129        fps: None,
130        seed: None,
131        image: None,
132        frame_images: Vec::new(),
133        input_references: Vec::new(),
134        generate_audio: None,
135        poll: None,
136        webhook: None,
137        download: None,
138        base: ModalityOptions::default(),
139    }
140}
141
142/// Builder returned by [`generate_video`]; `.await` runs the calls.
143pub struct GenerateVideo {
144    model: VideoModelRef,
145    prompt: Option<String>,
146    n: u32,
147    max_videos_per_call: Option<u32>,
148    aspect_ratio: Option<VideoAspectRatio>,
149    resolution: Option<ImageSize>,
150    duration: Option<f64>,
151    fps: Option<u32>,
152    seed: Option<u64>,
153    image: Option<VideoFile>,
154    frame_images: Vec<FrameImage>,
155    input_references: Vec<VideoFile>,
156    generate_audio: Option<bool>,
157    poll: Option<PollConfig>,
158    webhook: Option<WebhookFactory>,
159    download: Option<Arc<dyn DownloadFn>>,
160    base: ModalityOptions,
161}
162
163impl fmt::Debug for GenerateVideo {
164    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
165        f.debug_struct("GenerateVideo")
166            .field("model", &self.model)
167            .field("prompt", &self.prompt)
168            .field("n", &self.n)
169            .field("max_videos_per_call", &self.max_videos_per_call)
170            .field("aspect_ratio", &self.aspect_ratio)
171            .field("resolution", &self.resolution)
172            .field("duration", &self.duration)
173            .field("fps", &self.fps)
174            .field("seed", &self.seed)
175            .field("image", &self.image)
176            .field("frame_images", &self.frame_images)
177            .field("input_references", &self.input_references)
178            .field("generate_audio", &self.generate_audio)
179            .field("poll", &self.poll)
180            .field("has_webhook", &self.webhook.is_some())
181            .field("has_download", &self.download.is_some())
182            .field("base", &self.base)
183            .finish()
184    }
185}
186
187impl GenerateVideo {
188    /// Sets the text prompt.
189    #[must_use]
190    pub fn prompt(mut self, prompt: impl Into<String>) -> Self {
191        self.prompt = Some(prompt.into());
192        self
193    }
194
195    /// Number of videos to generate (default 1).
196    #[must_use]
197    pub fn n(mut self, n: u32) -> Self {
198        self.n = n;
199        self
200    }
201
202    /// Overrides the model's videos-per-call limit.
203    #[must_use]
204    pub fn max_videos_per_call(mut self, max_videos_per_call: u32) -> Self {
205        self.max_videos_per_call = Some(max_videos_per_call);
206        self
207    }
208
209    /// Aspect ratio.
210    #[must_use]
211    pub fn aspect_ratio(mut self, aspect_ratio: VideoAspectRatio) -> Self {
212        self.aspect_ratio = Some(aspect_ratio);
213        self
214    }
215
216    /// Resolution (`1280x720`).
217    #[must_use]
218    pub fn resolution(mut self, resolution: ImageSize) -> Self {
219        self.resolution = Some(resolution);
220        self
221    }
222
223    /// Duration in seconds.
224    #[must_use]
225    pub fn duration(mut self, seconds: f64) -> Self {
226        self.duration = Some(seconds);
227        self
228    }
229
230    /// Frames per second.
231    #[must_use]
232    pub fn fps(mut self, fps: u32) -> Self {
233        self.fps = Some(fps);
234        self
235    }
236
237    /// Seed for reproducible generation.
238    #[must_use]
239    pub fn seed(mut self, seed: u64) -> Self {
240        self.seed = Some(seed);
241        self
242    }
243
244    /// Source image (image-to-video).
245    #[must_use]
246    pub fn image(mut self, image: VideoFile) -> Self {
247        self.image = Some(image);
248        self
249    }
250
251    /// First/last frame images.
252    #[must_use]
253    pub fn frame_images(mut self, frame_images: Vec<FrameImage>) -> Self {
254        self.frame_images = frame_images;
255        self
256    }
257
258    /// Reference inputs.
259    #[must_use]
260    pub fn input_references(mut self, input_references: Vec<VideoFile>) -> Self {
261        self.input_references = input_references;
262        self
263    }
264
265    /// Whether to generate audio.
266    #[must_use]
267    pub fn generate_audio(mut self, generate_audio: bool) -> Self {
268        self.generate_audio = Some(generate_audio);
269        self
270    }
271
272    /// Uses the asynchronous flow with polling.
273    #[must_use]
274    pub fn poll(mut self, poll: PollConfig) -> Self {
275        self.poll = Some(poll);
276        self
277    }
278
279    /// Uses the asynchronous flow with a webhook created by `factory`
280    /// (falls back to polling when the model does not support webhooks).
281    #[must_use]
282    pub fn webhook(mut self, factory: WebhookFactory) -> Self {
283        self.webhook = Some(factory);
284        self
285    }
286
287    /// Sets the function used to fetch videos returned as URLs.
288    #[must_use]
289    pub fn download(mut self, download: Arc<dyn DownloadFn>) -> Self {
290        self.download = Some(download);
291        self
292    }
293}
294
295impl_modality_builder!(GenerateVideo);
296
297impl IntoFuture for GenerateVideo {
298    type Output = Result<GenerateVideoResult, Error>;
299    type IntoFuture = BoxFuture<'static, Self::Output>;
300
301    fn into_future(self) -> Self::IntoFuture {
302        Box::pin(run(self))
303    }
304}
305
306/// One provider call (owned so that it can run on a task).
307struct VideoCallTask {
308    model: Arc<dyn DynVideoModel>,
309    options: VideoOptions,
310    retry_policy: RetryPolicy,
311    cancellation: CancellationToken,
312    use_operations: bool,
313    poll: PollConfig,
314    webhook: Option<WebhookFactory>,
315}
316
317impl VideoCallTask {
318    async fn run(self) -> Result<VideoResult, Error> {
319        if self.use_operations {
320            self.run_operations().await
321        } else {
322            let model = &self.model;
323            let options = &self.options;
324            let token = &self.cancellation;
325            retry(&self.retry_policy, &self.cancellation, |_| {
326                let mut options = options.clone();
327                options.cancellation = token.child_token();
328                async move { model.do_generate(options).await.map_err(Error::from) }
329            })
330            .await
331        }
332    }
333
334    async fn run_operations(self) -> Result<VideoResult, Error> {
335        let mut warnings: Vec<Warning> = Vec::new();
336        let mut webhook_url = None;
337        let mut received = None;
338        if let Some(factory) = self.webhook.clone() {
339            if self.model.supports_webhook() {
340                let handle = self
341                    .model
342                    .handle_webhook(factory)
343                    .await
344                    .map_err(Error::from)?;
345                webhook_url = Some(handle.url);
346                received = Some(handle.received);
347            } else {
348                warnings.push(Warning::unsupported_with_details(
349                    "webhook",
350                    "This model does not support webhooks. Falling back to polling.",
351                ));
352            }
353        }
354
355        // `do_start` is billable: one idempotency key per logical start,
356        // minted outside the retry loop; a caller-supplied key wins.
357        let mut start_options = self.options.clone();
358        if !start_options.headers.contains(IDEMPOTENCY_KEY) {
359            start_options.headers = start_options
360                .headers
361                .with(IDEMPOTENCY_KEY, &format!("ferrin_vid_{}", generate_id()));
362        }
363        let model = &self.model;
364        let token = &self.cancellation;
365        let start = retry(&self.retry_policy, &self.cancellation, |_| {
366            let mut options = start_options.clone();
367            options.cancellation = token.child_token();
368            let webhook_url = webhook_url.clone();
369            async move {
370                model
371                    .do_start(VideoStartOptions {
372                        options,
373                        webhook_url,
374                    })
375                    .await
376                    .map_err(Error::from)
377            }
378        })
379        .await?;
380        warnings.extend(start.warnings);
381        let mut provider_metadata = start.provider_metadata;
382        let started = Instant::now();
383        let deadline = started + self.poll.timeout;
384
385        if let Some(received) = received {
386            let wait = tokio::time::timeout_at(deadline, received);
387            match wait.await {
388                Ok(Ok(_payload)) => {}
389                Ok(Err(error)) => return Err(Error::from(error)),
390                Err(_) => {
391                    return Err(Error::Timeout {
392                        scope: TimeoutScope::Total,
393                        elapsed: started.elapsed(),
394                    });
395                }
396            }
397        }
398        let waits_for_webhook = webhook_url.is_some();
399
400        let mut attempts: u32 = 0;
401        loop {
402            if !waits_for_webhook {
403                if Instant::now() >= deadline {
404                    return Err(Error::Timeout {
405                        scope: TimeoutScope::Total,
406                        elapsed: started.elapsed(),
407                    });
408                }
409                if let Some(max_attempts) = self.poll.max_attempts
410                    && attempts >= max_attempts
411                {
412                    return Err(Error::message(format!(
413                        "video generation did not complete after {max_attempts} status requests"
414                    )));
415                }
416                let sleep =
417                    tokio::time::sleep_until((Instant::now() + self.poll.interval).min(deadline));
418                tokio::select! {
419                    () = sleep => {}
420                    () = self.cancellation.cancelled() => return Err(Error::Cancelled),
421                }
422                if Instant::now() >= deadline {
423                    return Err(Error::Timeout {
424                        scope: TimeoutScope::Total,
425                        elapsed: started.elapsed(),
426                    });
427                }
428            }
429            attempts = attempts.saturating_add(1);
430            let operation = start.operation.clone();
431            let headers = self.options.headers.clone();
432            let status = retry(&self.retry_policy, &self.cancellation, |_| {
433                let options = VideoStatusOptions {
434                    operation: operation.clone(),
435                    headers: headers.clone(),
436                    cancellation: token.child_token(),
437                };
438                async move { model.do_status(options).await.map_err(Error::from) }
439            })
440            .await?;
441            match status {
442                VideoStatusResult::Error { error, .. } => {
443                    return Err(Error::message(format!("video generation failed: {error}")));
444                }
445                VideoStatusResult::Pending {
446                    warnings: status_warnings,
447                    provider_metadata: status_metadata,
448                    ..
449                } => {
450                    warnings.extend(status_warnings);
451                    if let Some(metadata) = status_metadata {
452                        merge_provider_metadata(
453                            provider_metadata.get_or_insert_with(ProviderMetadata::new),
454                            &metadata,
455                        );
456                    }
457                    if waits_for_webhook {
458                        return Err(Error::message(
459                            "video generation did not complete after the webhook notification",
460                        ));
461                    }
462                }
463                VideoStatusResult::Completed {
464                    videos,
465                    warnings: status_warnings,
466                    provider_metadata: status_metadata,
467                    response,
468                } => {
469                    warnings.extend(status_warnings);
470                    if let Some(metadata) = status_metadata {
471                        merge_provider_metadata(
472                            provider_metadata.get_or_insert_with(ProviderMetadata::new),
473                            &metadata,
474                        );
475                    }
476                    return Ok(VideoResult {
477                        videos,
478                        warnings,
479                        provider_metadata,
480                        response,
481                    });
482                }
483                #[allow(unreachable_patterns, reason = "VideoStatusResult is non-exhaustive")]
484                _ => {
485                    return Err(Error::message("unknown video status"));
486                }
487            }
488        }
489    }
490}
491
492fn usable_media_type(media_type: &MediaType) -> bool {
493    !media_type.as_str().is_empty() && media_type.as_str() != GENERIC_MEDIA_TYPE
494}
495
496async fn run(builder: GenerateVideo) -> Result<GenerateVideoResult, Error> {
497    if builder.n == 0 {
498        return Err(Error::invalid_argument("n", "must be at least 1"));
499    }
500    if builder.max_videos_per_call == Some(0) {
501        return Err(Error::invalid_argument(
502            "max_videos_per_call",
503            "must be at least 1",
504        ));
505    }
506    let model = resolve_model(&builder.model, ProviderRegistry::video_model)?;
507    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
508    let span = spans::modality_span("video", &identity);
509    let base = builder.base.clone();
510    base.run(|_, token| run_calls(model, identity, builder, token).instrument(span))
511        .await
512}
513
514async fn run_calls(
515    model: Arc<dyn DynVideoModel>,
516    identity: ModelIdentity,
517    builder: GenerateVideo,
518    cancellation: CancellationToken,
519) -> Result<GenerateVideoResult, Error> {
520    let supports_generate = model.supports_generate();
521    let supports_operations = model.supports_operations();
522    let wants_operations = builder.poll.is_some() || builder.webhook.is_some();
523    if !supports_generate && !supports_operations {
524        return Err(Error::from(ProviderError::unsupported(format!(
525            "video generation (model `{}` implements neither synchronous nor asynchronous generation)",
526            identity.model_id
527        ))));
528    }
529    let use_operations = supports_operations && (wants_operations || !supports_generate);
530    if wants_operations && !supports_operations {
531        spans::log_warnings(
532            &[Warning::other(
533                "poll/webhook options were provided but the model does not support \
534                 asynchronous operations; falling back to synchronous generation",
535            )],
536            &identity,
537        );
538    }
539
540    let per_call = builder
541        .max_videos_per_call
542        .map(|limit| usize::try_from(limit).unwrap_or(usize::MAX))
543        .or_else(|| model.max_videos_per_call().filter(|limit| *limit > 0))
544        .unwrap_or(1);
545    let n = usize::try_from(builder.n).unwrap_or(usize::MAX);
546    let counts: Vec<u32> = (0..n.div_ceil(per_call))
547        .map(|index| {
548            let remaining = n.saturating_sub(index.saturating_mul(per_call));
549            u32::try_from(remaining.min(per_call)).unwrap_or(u32::MAX)
550        })
551        .collect();
552
553    let template = VideoOptions {
554        prompt: builder.prompt.clone(),
555        n: 1,
556        aspect_ratio: builder.aspect_ratio,
557        resolution: builder.resolution,
558        duration: builder.duration,
559        fps: builder.fps,
560        seed: builder.seed,
561        image: builder.image.clone(),
562        frame_images: builder.frame_images.clone(),
563        input_references: builder.input_references.clone(),
564        generate_audio: builder.generate_audio,
565        provider_options: builder.base.provider_options.clone(),
566        headers: builder.base.request_headers(),
567        cancellation: cancellation.clone(),
568    };
569    let poll = builder.poll.clone().unwrap_or_default();
570    let make_task = |count: u32| VideoCallTask {
571        model: Arc::clone(&model),
572        options: VideoOptions {
573            n: count,
574            ..template.clone()
575        },
576        retry_policy: builder.base.retry_policy.clone(),
577        cancellation: cancellation.clone(),
578        use_operations,
579        poll: poll.clone(),
580        webhook: builder.webhook.clone(),
581    };
582
583    let mut results: Vec<Option<VideoResult>> = (0..counts.len()).map(|_| None).collect();
584    if let [count] = counts.as_slice() {
585        results = vec![Some(make_task(*count).run().await?)];
586    } else {
587        let mut tasks: JoinSet<(usize, Result<VideoResult, Error>)> = JoinSet::new();
588        for (index, count) in counts.iter().enumerate() {
589            let task = make_task(*count);
590            tasks.spawn(async move { (index, task.run().await) });
591        }
592        while let Some(joined) = tasks.join_next().await {
593            let (index, result) =
594                joined.map_err(|error| Error::message(format!("video task failed: {error}")))?;
595            if let Some(slot) = results.get_mut(index) {
596                *slot = Some(result?);
597            }
598        }
599    }
600
601    let mut downloader: Option<Arc<dyn DownloadFn>> = builder.download.clone();
602    let mut videos: Vec<GeneratedVideo> = Vec::new();
603    let mut warnings: Vec<Warning> = Vec::new();
604    let mut responses: Vec<ResponseMetadata> = Vec::new();
605    let mut provider_metadata = ProviderMetadata::new();
606    for result in results.into_iter().flatten() {
607        for video in result.videos {
608            let reported = usable_media_type(&video.media_type).then_some(video.media_type);
609            let (data, downloaded_media_type) = match video.data {
610                FileData::Bytes { data } => (data, None),
611                FileData::Url { url } => {
612                    let download = match &downloader {
613                        Some(download) => Arc::clone(download),
614                        None => {
615                            let default: Arc<dyn DownloadFn> =
616                                Arc::new(DefaultDownloader::try_default()?);
617                            downloader = Some(Arc::clone(&default));
618                            default
619                        }
620                    };
621                    let mut downloaded = download
622                        .download(
623                            vec![DownloadRequest {
624                                url: url.clone(),
625                                is_url_supported_by_model: false,
626                            }],
627                            cancellation.clone(),
628                        )
629                        .await?;
630                    match downloaded.pop().flatten() {
631                        Some(file) => (file.data, file.media_type.filter(usable_media_type)),
632                        None => {
633                            return Err(Error::download(
634                                url,
635                                None,
636                                Some("the download function returned no data".into()),
637                            ));
638                        }
639                    }
640                }
641                #[allow(unreachable_patterns, reason = "FileData is non-exhaustive")]
642                _ => {
643                    return Err(Error::invalid_data_content(
644                        "video data must be bytes or a URL",
645                        None,
646                    ));
647                }
648            };
649            let media_type = reported
650                .or(downloaded_media_type)
651                .or_else(|| detect_media_type_for(&data, "video"))
652                .unwrap_or_else(|| MediaType::new(DEFAULT_VIDEO_MEDIA_TYPE));
653            videos.push(GeneratedVideo { data, media_type });
654        }
655        warnings.extend(result.warnings);
656        responses.push(result.response);
657        if let Some(metadata) = &result.provider_metadata {
658            merge_provider_metadata(&mut provider_metadata, metadata);
659        }
660    }
661    if videos.is_empty() {
662        return Err(Error::NoVideoGenerated { responses });
663    }
664    spans::log_warnings(&warnings, &identity);
665    Ok(GenerateVideoResult {
666        videos,
667        warnings,
668        responses,
669        provider_metadata,
670    })
671}