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_metadata::merge_video_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            let result = tokio::select! {
388                biased;
389                () = self.cancellation.cancelled() => return Err(Error::Cancelled),
390                result = wait => result,
391            };
392            match result {
393                Ok(Ok(_payload)) => {}
394                Ok(Err(error)) => return Err(Error::from(error)),
395                Err(_) => {
396                    return Err(Error::Timeout {
397                        scope: TimeoutScope::Total,
398                        elapsed: started.elapsed(),
399                    });
400                }
401            }
402        }
403        let waits_for_webhook = webhook_url.is_some();
404
405        let mut attempts: u32 = 0;
406        loop {
407            if !waits_for_webhook {
408                if Instant::now() >= deadline {
409                    return Err(Error::Timeout {
410                        scope: TimeoutScope::Total,
411                        elapsed: started.elapsed(),
412                    });
413                }
414                if let Some(max_attempts) = self.poll.max_attempts
415                    && attempts >= max_attempts
416                {
417                    return Err(Error::message(format!(
418                        "video generation did not complete after {max_attempts} status requests"
419                    )));
420                }
421                let sleep =
422                    tokio::time::sleep_until((Instant::now() + self.poll.interval).min(deadline));
423                tokio::select! {
424                    () = sleep => {}
425                    () = self.cancellation.cancelled() => return Err(Error::Cancelled),
426                }
427                if Instant::now() >= deadline {
428                    return Err(Error::Timeout {
429                        scope: TimeoutScope::Total,
430                        elapsed: started.elapsed(),
431                    });
432                }
433            }
434            attempts = attempts.saturating_add(1);
435            let operation = start.operation.clone();
436            let headers = self.options.headers.clone();
437            let status_token = self.cancellation.child_token();
438            let _cancel_status_on_drop = status_token.clone().drop_guard();
439            let status_call = retry(&self.retry_policy, &status_token, |_| {
440                let options = VideoStatusOptions {
441                    operation: operation.clone(),
442                    headers: headers.clone(),
443                    cancellation: status_token.child_token(),
444                };
445                async move { model.do_status(options).await.map_err(Error::from) }
446            });
447            let status = tokio::select! {
448                biased;
449                () = self.cancellation.cancelled() => return Err(Error::Cancelled),
450                result = tokio::time::timeout_at(deadline, status_call) => {
451                    result.map_err(|_| Error::Timeout {
452                        scope: TimeoutScope::Total,
453                        elapsed: started.elapsed(),
454                    })??
455                }
456            };
457            match status {
458                VideoStatusResult::Error { error, .. } => {
459                    return Err(Error::message(format!("video generation failed: {error}")));
460                }
461                VideoStatusResult::Pending {
462                    warnings: status_warnings,
463                    provider_metadata: status_metadata,
464                    ..
465                } => {
466                    warnings.extend(status_warnings);
467                    if let Some(metadata) = status_metadata {
468                        merge_video_metadata(
469                            provider_metadata.get_or_insert_with(ProviderMetadata::new),
470                            &metadata,
471                        );
472                    }
473                    if waits_for_webhook {
474                        return Err(Error::message(
475                            "video generation did not complete after the webhook notification",
476                        ));
477                    }
478                }
479                VideoStatusResult::Completed {
480                    videos,
481                    warnings: status_warnings,
482                    provider_metadata: status_metadata,
483                    response,
484                } => {
485                    warnings.extend(status_warnings);
486                    if let Some(metadata) = status_metadata {
487                        merge_video_metadata(
488                            provider_metadata.get_or_insert_with(ProviderMetadata::new),
489                            &metadata,
490                        );
491                    }
492                    return Ok(VideoResult {
493                        videos,
494                        warnings,
495                        provider_metadata,
496                        response,
497                    });
498                }
499                #[allow(unreachable_patterns, reason = "VideoStatusResult is non-exhaustive")]
500                _ => {
501                    return Err(Error::message("unknown video status"));
502                }
503            }
504        }
505    }
506}
507
508fn usable_media_type(media_type: &MediaType) -> bool {
509    !media_type.as_str().is_empty() && media_type.as_str() != GENERIC_MEDIA_TYPE
510}
511
512async fn run(builder: GenerateVideo) -> Result<GenerateVideoResult, Error> {
513    if builder.n == 0 {
514        return Err(Error::invalid_argument("n", "must be at least 1"));
515    }
516    if builder.max_videos_per_call == Some(0) {
517        return Err(Error::invalid_argument(
518            "max_videos_per_call",
519            "must be at least 1",
520        ));
521    }
522    let model = resolve_model(&builder.model, ProviderRegistry::video_model)?;
523    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
524    let span = spans::modality_span("video", &identity);
525    let base = builder.base.clone();
526    base.run(|_, token| run_calls(model, identity, builder, token).instrument(span))
527        .await
528}
529
530async fn run_calls(
531    model: Arc<dyn DynVideoModel>,
532    identity: ModelIdentity,
533    builder: GenerateVideo,
534    cancellation: CancellationToken,
535) -> Result<GenerateVideoResult, Error> {
536    let supports_generate = model.supports_generate();
537    let supports_operations = model.supports_operations();
538    let wants_operations = builder.poll.is_some() || builder.webhook.is_some();
539    if !supports_generate && !supports_operations {
540        return Err(Error::from(ProviderError::unsupported(format!(
541            "video generation (model `{}` implements neither synchronous nor asynchronous generation)",
542            identity.model_id
543        ))));
544    }
545    let use_operations = supports_operations && (wants_operations || !supports_generate);
546    if wants_operations && !supports_operations {
547        spans::log_warnings(
548            &[Warning::other(
549                "poll/webhook options were provided but the model does not support \
550                 asynchronous operations; falling back to synchronous generation",
551            )],
552            &identity,
553        );
554    }
555
556    let per_call = builder
557        .max_videos_per_call
558        .map(|limit| usize::try_from(limit).unwrap_or(usize::MAX))
559        .or_else(|| model.max_videos_per_call().filter(|limit| *limit > 0))
560        .unwrap_or(1);
561    let n = usize::try_from(builder.n).unwrap_or(usize::MAX);
562    let counts: Vec<u32> = (0..n.div_ceil(per_call))
563        .map(|index| {
564            let remaining = n.saturating_sub(index.saturating_mul(per_call));
565            u32::try_from(remaining.min(per_call)).unwrap_or(u32::MAX)
566        })
567        .collect();
568
569    let mut warnings = Vec::new();
570    let input_references = if builder.frame_images.is_empty() {
571        builder.input_references.clone()
572    } else {
573        if !builder.input_references.is_empty() {
574            warnings.push(Warning::other(
575                "inputReferences were ignored because frameImages were provided; frameImages and inputReferences cannot be combined.",
576            ));
577        }
578        Vec::new()
579    };
580    let first_frame = builder
581        .frame_images
582        .iter()
583        .find(|frame| frame.frame_type == FrameType::FirstFrame)
584        .map(|frame| frame.image.clone());
585    if builder.image.is_some() && first_frame.is_some() {
586        warnings.push(Warning::other(
587            "prompt.image was ignored because a first_frame frameImage was provided; the first_frame frameImage takes precedence as the start image.",
588        ));
589    }
590    let image = first_frame.or_else(|| builder.image.clone());
591
592    let template = VideoOptions {
593        prompt: builder.prompt.clone(),
594        n: 1,
595        aspect_ratio: builder.aspect_ratio,
596        resolution: builder.resolution,
597        duration: builder.duration,
598        fps: builder.fps,
599        seed: builder.seed,
600        image,
601        frame_images: builder.frame_images.clone(),
602        input_references,
603        generate_audio: builder.generate_audio,
604        provider_options: builder.base.provider_options.clone(),
605        headers: builder.base.request_headers(),
606        cancellation: cancellation.clone(),
607    };
608    let poll = builder.poll.clone().unwrap_or_default();
609    let make_task = |count: u32| VideoCallTask {
610        model: Arc::clone(&model),
611        options: VideoOptions {
612            n: count,
613            ..template.clone()
614        },
615        retry_policy: builder.base.retry_policy.clone(),
616        cancellation: cancellation.clone(),
617        use_operations,
618        poll: poll.clone(),
619        webhook: builder.webhook.clone(),
620    };
621
622    let mut results: Vec<Option<VideoResult>> = (0..counts.len()).map(|_| None).collect();
623    if let [count] = counts.as_slice() {
624        results = vec![Some(make_task(*count).run().await?)];
625    } else {
626        let mut tasks: JoinSet<(usize, Result<VideoResult, Error>)> = JoinSet::new();
627        for (index, count) in counts.iter().enumerate() {
628            let task = make_task(*count);
629            tasks.spawn(async move { (index, task.run().await) });
630        }
631        while let Some(joined) = tasks.join_next().await {
632            let (index, result) =
633                joined.map_err(|error| Error::message(format!("video task failed: {error}")))?;
634            if let Some(slot) = results.get_mut(index) {
635                *slot = Some(result?);
636            }
637        }
638    }
639
640    let mut downloader: Option<Arc<dyn DownloadFn>> = builder.download.clone();
641    let mut videos: Vec<GeneratedVideo> = Vec::new();
642    let mut responses: Vec<ResponseMetadata> = Vec::new();
643    let mut provider_metadata = ProviderMetadata::new();
644    for result in results.into_iter().flatten() {
645        for video in result.videos {
646            let reported = usable_media_type(&video.media_type).then_some(video.media_type);
647            let (data, downloaded_media_type) = match video.data {
648                FileData::Bytes { data } => (data, None),
649                FileData::Url { url } => {
650                    let download = match &downloader {
651                        Some(download) => Arc::clone(download),
652                        None => {
653                            let default: Arc<dyn DownloadFn> =
654                                Arc::new(DefaultDownloader::try_default()?);
655                            downloader = Some(Arc::clone(&default));
656                            default
657                        }
658                    };
659                    let mut downloaded = download
660                        .download(
661                            vec![DownloadRequest {
662                                url: url.clone(),
663                                is_url_supported_by_model: false,
664                            }],
665                            cancellation.clone(),
666                        )
667                        .await?;
668                    match downloaded.pop().flatten() {
669                        Some(file) => (file.data, file.media_type.filter(usable_media_type)),
670                        None => {
671                            return Err(Error::download(
672                                url,
673                                None,
674                                Some("the download function returned no data".into()),
675                            ));
676                        }
677                    }
678                }
679                #[allow(unreachable_patterns, reason = "FileData is non-exhaustive")]
680                _ => {
681                    return Err(Error::invalid_data_content(
682                        "video data must be bytes or a URL",
683                        None,
684                    ));
685                }
686            };
687            let media_type = reported
688                .or(downloaded_media_type)
689                .or_else(|| detect_media_type_for(&data, "video"))
690                .unwrap_or_else(|| MediaType::new(DEFAULT_VIDEO_MEDIA_TYPE));
691            videos.push(GeneratedVideo { data, media_type });
692        }
693        warnings.extend(result.warnings);
694        responses.push(result.response);
695        if let Some(metadata) = &result.provider_metadata {
696            merge_video_metadata(&mut provider_metadata, metadata);
697        }
698    }
699    if videos.is_empty() {
700        return Err(Error::NoVideoGenerated { responses });
701    }
702    spans::log_warnings(&warnings, &identity);
703    Ok(GenerateVideoResult {
704        videos,
705        warnings,
706        responses,
707        provider_metadata,
708    })
709}