1use 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
57const DEFAULT_VIDEO_MEDIA_TYPE: &str = "video/mp4";
59
60const GENERIC_MEDIA_TYPE: &str = "application/octet-stream";
62
63const IDEMPOTENCY_KEY: &str = "idempotency-key";
65
66#[derive(Debug, Clone, PartialEq, Eq)]
68pub struct PollConfig {
69 pub interval: Duration,
71 pub timeout: Duration,
74 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#[derive(Debug, Clone, PartialEq, Eq)]
90pub struct GeneratedVideo {
91 pub data: Bytes,
93 pub media_type: MediaType,
95}
96
97#[derive(Debug, Clone, PartialEq)]
99pub struct GenerateVideoResult {
100 pub videos: Vec<GeneratedVideo>,
102 pub warnings: Vec<Warning>,
104 pub responses: Vec<ResponseMetadata>,
106 pub provider_metadata: ProviderMetadata,
108}
109
110impl GenerateVideoResult {
111 #[must_use]
113 pub fn video(&self) -> Option<&GeneratedVideo> {
114 self.videos.first()
115 }
116}
117
118#[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
142pub 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 #[must_use]
190 pub fn prompt(mut self, prompt: impl Into<String>) -> Self {
191 self.prompt = Some(prompt.into());
192 self
193 }
194
195 #[must_use]
197 pub fn n(mut self, n: u32) -> Self {
198 self.n = n;
199 self
200 }
201
202 #[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 #[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 #[must_use]
218 pub fn resolution(mut self, resolution: ImageSize) -> Self {
219 self.resolution = Some(resolution);
220 self
221 }
222
223 #[must_use]
225 pub fn duration(mut self, seconds: f64) -> Self {
226 self.duration = Some(seconds);
227 self
228 }
229
230 #[must_use]
232 pub fn fps(mut self, fps: u32) -> Self {
233 self.fps = Some(fps);
234 self
235 }
236
237 #[must_use]
239 pub fn seed(mut self, seed: u64) -> Self {
240 self.seed = Some(seed);
241 self
242 }
243
244 #[must_use]
246 pub fn image(mut self, image: VideoFile) -> Self {
247 self.image = Some(image);
248 self
249 }
250
251 #[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 #[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 #[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 #[must_use]
274 pub fn poll(mut self, poll: PollConfig) -> Self {
275 self.poll = Some(poll);
276 self
277 }
278
279 #[must_use]
282 pub fn webhook(mut self, factory: WebhookFactory) -> Self {
283 self.webhook = Some(factory);
284 self
285 }
286
287 #[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
306struct 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 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}