1use crate::{
7 GenerateRequest, GenerationReference, GenerationReferenceAuthority, GenerationReferenceKind,
8 GenerationReferenceProvenance, MoldClient, ReferenceUploadCapabilities,
9 ReferenceUploadCompleteResponse, ReferenceUploadSessionRequest, ReferenceUploadSessionResponse,
10};
11use anyhow::{ensure, Context, Result};
12use sha2::{Digest, Sha256};
13use std::{
14 collections::HashSet,
15 fs::File,
16 time::{SystemTime, UNIX_EPOCH},
17};
18
19pub const PROTOCOL_VERSION: u32 = 2;
20pub const SESSION_PATH: &str = "/api/generate/reference-upload-sessions";
21pub const UPLOAD_PATH: &str = "/api/generate/reference-upload";
22pub const SESSION_HANDLE_HEADER: &str = "x-mold-reference-upload-session";
23pub const UPLOAD_HANDLE_HEADER: &str = "x-mold-reference-upload";
24
25enum ReferenceUploadBody {
26 OpenFile(File),
27 Bytes(Vec<u8>),
28}
29
30pub struct ReferenceUploadSource {
32 body: ReferenceUploadBody,
33 length: u64,
34}
35
36impl std::fmt::Debug for ReferenceUploadSource {
37 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
38 formatter
39 .debug_struct("ReferenceUploadSource")
40 .field("body", &"<redacted>")
41 .field("length", &self.length)
42 .finish()
43 }
44}
45
46impl ReferenceUploadSource {
47 pub fn open_file(file: File) -> Result<Self> {
48 let metadata = file
49 .metadata()
50 .context("failed to inspect reference file")?;
51 ensure!(
52 metadata.is_file() && metadata.len() > 0,
53 "reference upload source is not a non-empty regular file"
54 );
55 Ok(Self {
56 body: ReferenceUploadBody::OpenFile(file),
57 length: metadata.len(),
58 })
59 }
60
61 pub fn bytes(bytes: Vec<u8>) -> Result<Self> {
62 ensure!(!bytes.is_empty(), "reference upload source is empty");
63 let length = u64::try_from(bytes.len()).context("reference upload is too large")?;
64 Ok(Self {
65 body: ReferenceUploadBody::Bytes(bytes),
66 length,
67 })
68 }
69
70 fn sha256(&self) -> Result<String> {
71 match &self.body {
72 ReferenceUploadBody::OpenFile(file) => crate::secure_file::sha256_open_file(file),
73 ReferenceUploadBody::Bytes(bytes) => Ok(format!("{:x}", Sha256::digest(bytes))),
74 }
75 }
76}
77
78pub struct ReferenceUploadLease {
80 request: GenerateRequest,
81 pub expires_at_ms: u64,
82 pub request_scope_sha256: String,
83 cancellation: SessionCancellationGuard,
84}
85
86impl std::fmt::Debug for ReferenceUploadLease {
87 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88 formatter
89 .debug_struct("ReferenceUploadLease")
90 .field("request", &"<redacted frozen request>")
91 .field("session_handle", &"<redacted>")
92 .field("expires_at_ms", &self.expires_at_ms)
93 .field("request_scope_sha256", &self.request_scope_sha256)
94 .finish()
95 }
96}
97
98impl ReferenceUploadLease {
99 pub fn request(&self) -> &GenerateRequest {
102 &self.request
103 }
104
105 pub fn mark_consumed(&mut self) {
108 self.cancellation.disarm();
109 }
110
111 pub async fn cancel(&mut self) -> Result<()> {
114 self.cancellation.cancel().await
115 }
116}
117
118struct SessionCancellationGuard {
119 client: MoldClient,
120 session_handle: Option<String>,
121}
122
123impl SessionCancellationGuard {
124 fn new(client: MoldClient, session_handle: &str) -> Self {
125 Self {
126 client,
127 session_handle: valid_handle(session_handle).then(|| session_handle.to_string()),
128 }
129 }
130
131 fn disarm(&mut self) {
132 self.session_handle = None;
133 }
134
135 async fn cancel(&mut self) -> Result<()> {
136 let Some(session_handle) = self.session_handle.clone() else {
137 return Ok(());
138 };
139 self.client
140 .cancel_reference_upload_session(&session_handle)
141 .await?;
142 self.disarm();
143 Ok(())
144 }
145}
146
147impl Drop for SessionCancellationGuard {
148 fn drop(&mut self) {
149 let Some(session_handle) = self.session_handle.take() else {
150 return;
151 };
152 let Ok(runtime) = tokio::runtime::Handle::try_current() else {
153 return;
156 };
157 let client = self.client.clone();
158 runtime.spawn(async move {
159 let _ = client
160 .cancel_reference_upload_session(&session_handle)
161 .await;
162 });
163 }
164}
165
166fn valid_sha256(value: &str) -> bool {
167 value.len() == 64 && value.bytes().all(|byte| byte.is_ascii_hexdigit())
168}
169
170fn valid_handle(value: &str) -> bool {
171 !value.is_empty()
172 && value.len() <= crate::minimax_h3::MAX_REFERENCE_UPLOAD_HANDLE_BYTES
173 && value
174 .bytes()
175 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b':' | b'.'))
176}
177
178fn now_ms() -> Result<u64> {
179 let millis = SystemTime::now()
180 .duration_since(UNIX_EPOCH)
181 .context("system clock is before the Unix epoch")?
182 .as_millis();
183 u64::try_from(millis).context("system clock is outside the supported range")
184}
185
186fn ensure_session_live(expires_at_ms: u64, observed_at_ms: u64) -> Result<()> {
187 ensure!(
188 expires_at_ms > observed_at_ms,
189 "host returned an expired reference-upload session"
190 );
191 Ok(())
192}
193
194fn validate_capabilities(capabilities: &ReferenceUploadCapabilities) -> Result<()> {
195 ensure!(
196 capabilities.available
197 && capabilities.protocol_version == PROTOCOL_VERSION
198 && capabilities.requires_api_key,
199 "host does not advertise authenticated reference-upload protocol V2"
200 );
201 ensure!(
202 capabilities.session_path == SESSION_PATH
203 && capabilities.upload_path == UPLOAD_PATH
204 && capabilities
205 .session_handle_header
206 .eq_ignore_ascii_case(SESSION_HANDLE_HEADER)
207 && capabilities
208 .upload_handle_header
209 .eq_ignore_ascii_case(UPLOAD_HANDLE_HEADER),
210 "host advertised an incompatible reference-upload V2 wire contract"
211 );
212 ensure!(
213 capabilities.max_file_bytes > 0
214 && capabilities.max_session_bytes >= capabilities.max_file_bytes
215 && capabilities.session_ttl_ms > 0,
216 "host advertised invalid reference-upload V2 limits"
217 );
218 Ok(())
219}
220
221fn validate_session(
222 session: &ReferenceUploadSessionResponse,
223 expected_instance_id: &str,
224 expected_count: usize,
225 now_ms: u64,
226) -> Result<Vec<String>> {
227 ensure!(
228 session.instance_id == expected_instance_id,
229 "reference-upload session came from a different Mold instance"
230 );
231 ensure_session_live(session.expires_at_ms, now_ms)?;
232 ensure!(
233 valid_sha256(&session.request_scope_sha256)
234 && valid_handle(&session.session_handle)
235 && session.uploads.len() == expected_count,
236 "host returned an incomplete reference-upload session"
237 );
238 let mut ordered = vec![None; expected_count];
239 let mut handles = HashSet::new();
240 for slot in &session.uploads {
241 let index = usize::try_from(slot.reference)
242 .ok()
243 .and_then(|reference| reference.checked_sub(1))
244 .filter(|index| *index < expected_count)
245 .context("host returned an out-of-range reference-upload slot")?;
246 ensure!(
247 valid_handle(&slot.handle)
248 && handles.insert(slot.handle.as_str())
249 && ordered[index].replace(slot.handle.clone()).is_none(),
250 "host returned an invalid or duplicate reference-upload slot"
251 );
252 }
253 ordered
254 .into_iter()
255 .enumerate()
256 .map(|(index, handle)| {
257 handle.with_context(|| format!("host omitted reference-upload slot {}", index + 1))
258 })
259 .collect()
260}
261
262fn canonical_reference(
263 provisional: &GenerationReference,
264 complete: &ReferenceUploadCompleteResponse,
265 expected_instance_id: &str,
266 expected_reference: u32,
267 expected_session_complete: bool,
268 upload_handle: String,
269) -> Result<GenerationReference> {
270 ensure!(
271 complete.instance_id == expected_instance_id
272 && complete.reference == expected_reference
273 && complete.metadata.index == expected_reference
274 && complete.metadata.kind == provisional.kind()
275 && valid_sha256(&complete.request_scope_sha256)
276 && complete.session_complete == expected_session_complete,
277 "reference {expected_reference} upload response did not match its bound V2 session"
278 );
279 let expected = provisional
280 .redacted_metadata(usize::try_from(expected_reference.saturating_sub(1))?)
281 .with_context(|| format!("reference {expected_reference} lost digest provenance"))?;
282 ensure!(
283 complete
284 .metadata
285 .sha256
286 .eq_ignore_ascii_case(&expected.sha256)
287 && complete.metadata.name == expected.name,
288 "reference {expected_reference} upload response changed immutable provenance"
289 );
290 let provenance = GenerationReferenceProvenance {
291 name: complete.metadata.name.clone(),
292 sha256: Some(complete.metadata.sha256.to_ascii_lowercase()),
293 };
294 let descriptor = match provisional {
295 GenerationReference::Image { mime_type, .. } => {
296 ensure!(
297 complete.metadata.kind == GenerationReferenceKind::Image
298 && complete.metadata.mime_type == *mime_type,
299 "reference {expected_reference} upload response changed its image type"
300 );
301 GenerationReference::Image {
302 media: GenerationReferenceAuthority::Descriptor,
303 provenance,
304 mime_type: complete.metadata.mime_type.clone(),
305 width: complete
306 .metadata
307 .width
308 .filter(|value| *value > 0)
309 .context("canonical image width is missing")?,
310 height: complete
311 .metadata
312 .height
313 .filter(|value| *value > 0)
314 .context("canonical image height is missing")?,
315 }
316 }
317 GenerationReference::Video { mime_type, .. } => {
318 ensure!(
319 mime_type == "video/mp4"
320 && complete.metadata.kind == GenerationReferenceKind::Video
321 && complete.metadata.mime_type == "video/mp4",
322 "reference {expected_reference} upload response changed its video type"
323 );
324 let fps = complete
325 .metadata
326 .fps
327 .filter(|value| value.is_finite() && *value > 0.0)
328 .context("canonical video fps is missing")?;
329 let (audio_duration_ms, audio_sample_count, audio_sample_rate, audio_channels) =
330 if complete.metadata.has_audio {
331 let channels = complete
332 .metadata
333 .audio_channels
334 .filter(|value| matches!(value, 1 | 2))
335 .context("canonical video audio channel count is missing")?;
336 (
337 Some(
338 complete
339 .metadata
340 .audio_duration_ms
341 .filter(|value| *value > 0)
342 .context("canonical video audio duration is missing")?,
343 ),
344 Some(
345 complete
346 .metadata
347 .audio_sample_count
348 .filter(|value| *value > 0)
349 .context("canonical video audio sample count is missing")?,
350 ),
351 Some(
352 complete
353 .metadata
354 .audio_sample_rate
355 .filter(|value| *value > 0)
356 .context("canonical video audio sample rate is missing")?,
357 ),
358 Some(channels),
359 )
360 } else {
361 ensure!(
362 complete.metadata.audio_duration_ms.is_none()
363 && complete.metadata.audio_sample_count.is_none()
364 && complete.metadata.audio_sample_rate.is_none()
365 && complete.metadata.audio_channels.is_none(),
366 "canonical silent video unexpectedly contains audio facts"
367 );
368 (None, None, None, None)
369 };
370 GenerationReference::Video {
371 media: GenerationReferenceAuthority::Descriptor,
372 provenance,
373 mime_type: "video/mp4".to_string(),
374 width: complete
375 .metadata
376 .width
377 .filter(|value| *value > 0)
378 .context("canonical video width is missing")?,
379 height: complete
380 .metadata
381 .height
382 .filter(|value| *value > 0)
383 .context("canonical video height is missing")?,
384 frame_count: Some(
385 complete
386 .metadata
387 .frame_count
388 .filter(|value| *value > 0)
389 .context("canonical video frame count is missing")?,
390 ),
391 duration_ms: complete
392 .metadata
393 .duration_ms
394 .filter(|value| *value > 0)
395 .context("canonical video duration is missing")?,
396 fps,
397 has_audio: complete.metadata.has_audio,
398 audio_duration_ms,
399 audio_sample_count,
400 audio_sample_rate,
401 audio_channels,
402 }
403 }
404 GenerationReference::Audio { mime_type, .. } => {
405 let declared = mime_type.trim().to_ascii_lowercase();
406 ensure!(
407 matches!(
408 declared.as_str(),
409 "audio/wav" | "audio/x-wav" | "audio/wave"
410 ) && complete.metadata.kind == GenerationReferenceKind::Audio
411 && complete.metadata.mime_type == "audio/wav",
412 "reference {expected_reference} upload response changed its audio type"
413 );
414 GenerationReference::Audio {
415 media: GenerationReferenceAuthority::Descriptor,
416 provenance,
417 mime_type: "audio/wav".to_string(),
418 duration_ms: complete
419 .metadata
420 .duration_ms
421 .filter(|value| *value > 0)
422 .context("canonical audio duration is missing")?,
423 sample_rate: complete
424 .metadata
425 .sample_rate
426 .filter(|value| *value > 0)
427 .context("canonical audio sample rate is missing")?,
428 channels: complete
429 .metadata
430 .channels
431 .filter(|value| matches!(value, 1 | 2))
432 .context("canonical audio channel count is missing")?,
433 sample_count: Some(
434 complete
435 .metadata
436 .sample_count
437 .filter(|value| *value > 0)
438 .context("canonical audio sample count is missing")?,
439 ),
440 }
441 }
442 };
443 crate::minimax_h3::reference_prepared_shape(&descriptor).map_err(anyhow::Error::new)?;
444 Ok(with_authority(
445 descriptor,
446 GenerationReferenceAuthority::Upload {
447 handle: upload_handle,
448 },
449 ))
450}
451
452fn with_authority(
453 reference: GenerationReference,
454 media: GenerationReferenceAuthority,
455) -> GenerationReference {
456 match reference {
457 GenerationReference::Image {
458 provenance,
459 mime_type,
460 width,
461 height,
462 ..
463 } => GenerationReference::Image {
464 media,
465 provenance,
466 mime_type,
467 width,
468 height,
469 },
470 GenerationReference::Video {
471 provenance,
472 mime_type,
473 width,
474 height,
475 frame_count,
476 duration_ms,
477 fps,
478 has_audio,
479 audio_duration_ms,
480 audio_sample_count,
481 audio_sample_rate,
482 audio_channels,
483 ..
484 } => GenerationReference::Video {
485 media,
486 provenance,
487 mime_type,
488 width,
489 height,
490 frame_count,
491 duration_ms,
492 fps,
493 has_audio,
494 audio_duration_ms,
495 audio_sample_count,
496 audio_sample_rate,
497 audio_channels,
498 },
499 GenerationReference::Audio {
500 provenance,
501 mime_type,
502 duration_ms,
503 sample_rate,
504 channels,
505 sample_count,
506 ..
507 } => GenerationReference::Audio {
508 media,
509 provenance,
510 mime_type,
511 duration_ms,
512 sample_rate,
513 channels,
514 sample_count,
515 },
516 }
517}
518
519fn mime_type(reference: &GenerationReference) -> &str {
520 match reference {
521 GenerationReference::Image { mime_type, .. }
522 | GenerationReference::Video { mime_type, .. }
523 | GenerationReference::Audio { mime_type, .. } => mime_type,
524 }
525}
526
527fn rebind_canonical_request(
528 mut scoped_request: GenerateRequest,
529 references: Vec<GenerationReference>,
530) -> GenerateRequest {
531 scoped_request.references = Some(references);
532 scoped_request
533}
534
535impl MoldClient {
536 pub async fn bind_reference_uploads_v2(
540 &self,
541 request: &GenerateRequest,
542 sources: Vec<ReferenceUploadSource>,
543 ) -> Result<ReferenceUploadLease> {
544 ensure!(
545 self.has_api_key(),
546 "reference uploads require a client configured with a valid API key"
547 );
548 let (capabilities, status) =
549 tokio::try_join!(self.server_capabilities(), self.server_status())?;
550 let capabilities = capabilities.reference_uploads;
551 validate_capabilities(&capabilities)?;
552 let expected_instance_id = status
553 .instance_id
554 .filter(|value| !value.trim().is_empty())
555 .context("host status does not expose an exact Mold instance identity")?;
556 let scoped_request = request.clone();
557 let descriptors = scoped_request
558 .references
559 .clone()
560 .context("reference uploads require ordered H3 descriptors")?;
561 ensure!(
562 !descriptors.is_empty() && descriptors.len() == sources.len(),
563 "reference descriptor/body count mismatch"
564 );
565 crate::minimax_h3::validate_reference_descriptors(&descriptors)
566 .map_err(anyhow::Error::new)?;
567
568 let mut total_bytes = 0_u64;
569 for (index, (descriptor, source)) in descriptors.iter().zip(&sources).enumerate() {
570 ensure!(
571 source.length <= capabilities.max_file_bytes,
572 "reference {} exceeds this host's per-file upload limit",
573 index + 1
574 );
575 total_bytes = total_bytes
576 .checked_add(source.length)
577 .context("combined reference size overflowed")?;
578 ensure!(
579 total_bytes <= capabilities.max_session_bytes,
580 "ordered references exceed this host's upload-session limit"
581 );
582 let expected_digest = descriptor
583 .content_sha256()
584 .with_context(|| format!("reference {} has no content digest", index + 1))?;
585 ensure!(
586 source.sha256()?.eq_ignore_ascii_case(&expected_digest),
587 "reference {} changed after it was probed",
588 index + 1
589 );
590 }
591
592 let upload_references = (1..=descriptors.len())
593 .map(|index| u32::try_from(index).context("too many reference uploads"))
594 .collect::<Result<Vec<_>>>()?;
595 let session = self
596 .create_reference_upload_session(&ReferenceUploadSessionRequest {
597 request: scoped_request.clone(),
598 upload_references,
599 })
600 .await?;
601 let mut cancellation = SessionCancellationGuard::new(self.clone(), &session.session_handle);
604 let handles = match validate_session(
605 &session,
606 &expected_instance_id,
607 descriptors.len(),
608 now_ms()?,
609 ) {
610 Ok(handles) => handles,
611 Err(error) => {
612 let _ = cancellation.cancel().await;
613 return Err(error);
614 }
615 };
616
617 let result: Result<(Vec<GenerationReference>, String)> = async {
618 let upload_count = sources.len();
619 let mut canonical = Vec::with_capacity(upload_count);
620 let mut canonical_scope = session.request_scope_sha256.to_ascii_lowercase();
621 for (index, ((descriptor, source), handle)) in
622 descriptors.iter().zip(sources).zip(handles).enumerate()
623 {
624 let reference = u32::try_from(index + 1).context("too many references")?;
625 let content_type = mime_type(descriptor);
626 let completed = match source.body {
627 ReferenceUploadBody::OpenFile(file) => {
628 self.upload_reference_open_file(&handle, file, content_type)
629 .await
630 }
631 ReferenceUploadBody::Bytes(bytes) => {
632 self.upload_reference_bytes(&handle, bytes, content_type)
633 .await
634 }
635 }
636 .with_context(|| format!("reference {reference} upload failed"))?;
637 let bound = canonical_reference(
638 descriptor,
639 &completed,
640 &expected_instance_id,
641 reference,
642 index + 1 == upload_count,
643 handle,
644 )?;
645 canonical_scope = completed.request_scope_sha256.to_ascii_lowercase();
646 canonical.push(bound);
647 }
648 crate::minimax_h3::validate_references(&canonical).map_err(anyhow::Error::new)?;
649 ensure_session_live(session.expires_at_ms, now_ms()?).context(
650 "reference-upload session expired before its canonical request could be returned",
651 )?;
652 Ok((canonical, canonical_scope))
653 }
654 .await;
655
656 match result {
657 Ok((references, request_scope_sha256)) => {
658 let request = rebind_canonical_request(scoped_request, references);
659 Ok(ReferenceUploadLease {
660 request,
661 expires_at_ms: session.expires_at_ms,
662 request_scope_sha256,
663 cancellation,
664 })
665 }
666 Err(error) => {
667 let _ = cancellation.cancel().await;
668 Err(error)
669 }
670 }
671 }
672}
673
674#[cfg(test)]
675mod tests {
676 use super::*;
677
678 fn request_with_references(references: Vec<GenerationReference>) -> GenerateRequest {
679 let mut request: GenerateRequest = serde_json::from_value(serde_json::json!({
680 "prompt": "keep this request frozen",
681 "model": crate::minimax_h3::REF2VA_COMFY,
682 "width": crate::minimax_h3::DEFAULT_WIDTH,
683 "height": crate::minimax_h3::DEFAULT_HEIGHT,
684 "steps": crate::minimax_h3::DEFAULT_STEPS,
685 "guidance": 0.0,
686 "seed": 77,
687 "batch_size": 1,
688 "output_format": "mp4",
689 "strength": 1.0,
690 "frames": crate::minimax_h3::MIN_FRAMES,
691 "fps": crate::minimax_h3::FIXED_FPS,
692 "enable_audio": true
693 }))
694 .unwrap();
695 request.references = Some(references);
696 request
697 }
698
699 fn audio_descriptor() -> GenerationReference {
700 GenerationReference::Audio {
701 media: GenerationReferenceAuthority::Descriptor,
702 provenance: GenerationReferenceProvenance {
703 name: Some("voice.wav".into()),
704 sha256: Some("a".repeat(64)),
705 },
706 mime_type: "audio/x-wav".into(),
707 duration_ms: 1_000,
708 sample_rate: 44_100,
709 channels: 2,
710 sample_count: Some(44_100),
711 }
712 }
713
714 #[test]
715 fn canonical_audio_response_rebinds_provisional_facts() {
716 let completed = ReferenceUploadCompleteResponse {
717 instance_id: "instance".into(),
718 reference: 1,
719 metadata: crate::GenerationReferenceMetadata {
720 kind: GenerationReferenceKind::Audio,
721 index: 1,
722 name: Some("voice.wav".into()),
723 sha256: "a".repeat(64),
724 mime_type: "audio/wav".into(),
725 width: None,
726 height: None,
727 frame_count: None,
728 duration_ms: Some(998),
729 fps: None,
730 has_audio: false,
731 audio_duration_ms: None,
732 audio_sample_count: None,
733 audio_sample_rate: None,
734 audio_channels: None,
735 sample_rate: Some(48_000),
736 channels: Some(2),
737 sample_count: Some(47_904),
738 prepared_shape: None,
739 },
740 request_scope_sha256: "b".repeat(64),
741 session_complete: true,
742 };
743 let bound = canonical_reference(
744 &audio_descriptor(),
745 &completed,
746 "instance",
747 1,
748 true,
749 "mru_handle".into(),
750 )
751 .unwrap();
752 assert!(matches!(
753 bound,
754 GenerationReference::Audio {
755 media: GenerationReferenceAuthority::Upload { .. },
756 duration_ms: 998,
757 sample_rate: 48_000,
758 sample_count: Some(47_904),
759 ..
760 }
761 ));
762 }
763
764 #[test]
765 fn lease_debug_redacts_frozen_request_and_session_authority() {
766 let mut request = request_with_references(vec![with_authority(
767 audio_descriptor(),
768 GenerationReferenceAuthority::Upload {
769 handle: "mru_private_upload".into(),
770 },
771 )]);
772 request.prompt = "private prompt".into();
773 request.source_image = Some(vec![17, 42, 99]);
774 let lease = ReferenceUploadLease {
775 request,
776 expires_at_ms: u64::MAX,
777 request_scope_sha256: "c".repeat(64),
778 cancellation: SessionCancellationGuard::new(
779 MoldClient::new("http://127.0.0.1:9"),
780 "mrs_private_session",
781 ),
782 };
783
784 let debug = format!("{lease:?}");
785 assert!(debug.contains("<redacted frozen request>"));
786 assert!(!debug.contains("private prompt"));
787 assert!(!debug.contains("mru_private_upload"));
788 assert!(!debug.contains("mrs_private_session"));
789 assert!(!debug.contains("17, 42, 99"));
790 }
791
792 #[test]
793 fn completion_order_and_scope_fail_closed() {
794 let mut completed = ReferenceUploadCompleteResponse {
795 instance_id: "instance".into(),
796 reference: 1,
797 metadata: audio_descriptor().redacted_metadata(0).unwrap(),
798 request_scope_sha256: "b".repeat(64),
799 session_complete: true,
800 };
801 assert!(canonical_reference(
802 &audio_descriptor(),
803 &completed,
804 "instance",
805 1,
806 false,
807 "mru_handle".into(),
808 )
809 .is_err());
810 completed.session_complete = false;
811 completed.request_scope_sha256 = "bad".into();
812 assert!(canonical_reference(
813 &audio_descriptor(),
814 &completed,
815 "instance",
816 1,
817 false,
818 "mru_handle".into(),
819 )
820 .is_err());
821 }
822
823 #[test]
824 fn canonical_rebind_changes_only_the_reference_authority() {
825 let original = request_with_references(vec![audio_descriptor()]);
826 let bound_reference = with_authority(
827 audio_descriptor(),
828 GenerationReferenceAuthority::Upload {
829 handle: "mru_handle".into(),
830 },
831 );
832 let rebound = rebind_canonical_request(original.clone(), vec![bound_reference.clone()]);
833 let mut expected = original;
834 expected.references = Some(vec![bound_reference]);
835 assert_eq!(
836 serde_json::to_value(rebound).unwrap(),
837 serde_json::to_value(expected).unwrap()
838 );
839 }
840
841 #[test]
842 fn final_expiry_check_rejects_a_session_that_elapsed_during_upload() {
843 assert!(ensure_session_live(101, 100).is_ok());
844 assert!(ensure_session_live(100, 100).is_err());
845 assert!(ensure_session_live(99, 100).is_err());
846 }
847
848 #[tokio::test]
849 async fn unauthenticated_client_fails_before_any_host_or_media_work() {
850 let server = wiremock::MockServer::start().await;
851 let request = request_with_references(vec![audio_descriptor()]);
852 let error = MoldClient::new(&server.uri())
853 .bind_reference_uploads_v2(
854 &request,
855 vec![ReferenceUploadSource::bytes(b"unread media".to_vec()).unwrap()],
856 )
857 .await
858 .unwrap_err();
859 assert!(error
860 .to_string()
861 .contains("configured with a valid API key"));
862 assert!(server.received_requests().await.unwrap().is_empty());
863 }
864
865 #[tokio::test]
866 async fn dropping_armed_session_guard_schedules_cancellation() {
867 use wiremock::matchers::{header, method, path};
868 use wiremock::{Mock, MockServer, ResponseTemplate};
869
870 let server = MockServer::start().await;
871 Mock::given(method("DELETE"))
872 .and(path(SESSION_PATH))
873 .and(header("x-api-key", "sekrit"))
874 .and(header(SESSION_HANDLE_HEADER, "mrs_abort"))
875 .respond_with(ResponseTemplate::new(204))
876 .expect(1)
877 .mount(&server)
878 .await;
879 let client = MoldClient::with_api_key(&server.uri(), "sekrit".into());
880 drop(SessionCancellationGuard::new(client, "mrs_abort"));
881 for _ in 0..20 {
882 if !server.received_requests().await.unwrap().is_empty() {
883 break;
884 }
885 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
886 }
887 assert_eq!(server.received_requests().await.unwrap().len(), 1);
888 }
889}