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::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
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 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}