1use std::fmt;
2
3use base64::{engine::general_purpose::STANDARD, Engine as _};
4use serde::{Deserialize, Serialize};
5
6use crate::{ArtifactKind, ArtifactReference};
7
8#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
9#[serde(rename_all = "camelCase", deny_unknown_fields)]
10pub struct PutArtifactRequest {
11 pub kind: ArtifactKind,
12 pub mime_type: String,
13 pub data_base64: String,
14 #[serde(default, skip_serializing_if = "Option::is_none")]
15 pub transform: Option<crate::ImageTransform>,
16 #[serde(default, skip_serializing_if = "Option::is_none")]
17 pub duration_millis: Option<u64>,
18}
19
20impl PutArtifactRequest {
21 pub fn decode(&self) -> Result<Vec<u8>, PutArtifactValidationError> {
22 if let Some(transform) = &self.transform {
23 if self.kind != ArtifactKind::Image {
24 return Err(invalid("image transform requires IMAGE kind"));
25 }
26 transform.validate().map_err(invalid)?;
27 }
28 if self.mime_type.trim() != self.mime_type || self.mime_type.is_empty() {
29 return Err(invalid("artifact put MIME type is invalid"));
30 }
31 match self.kind {
32 ArtifactKind::Image if !self.mime_type.starts_with("image/") => {
33 return Err(invalid("artifact put image MIME type is invalid"));
34 }
35 ArtifactKind::Audio if !self.mime_type.starts_with("audio/") => {
36 return Err(invalid("artifact put audio MIME type is invalid"));
37 }
38 ArtifactKind::Video if !self.mime_type.starts_with("video/") => {
39 return Err(invalid("artifact put video MIME type is invalid"));
40 }
41 ArtifactKind::File if self.mime_type.starts_with("video/") => {
42 return Err(invalid("video artifacts are not supported"));
43 }
44 _ => {}
45 }
46 if !matches!(self.kind, ArtifactKind::Audio | ArtifactKind::Video)
47 && self.duration_millis.is_some()
48 {
49 return Err(invalid(
50 "artifact put duration is only valid for audio or video",
51 ));
52 }
53 if self.duration_millis == Some(0) {
54 return Err(invalid("artifact put duration is invalid"));
55 }
56 let bytes = STANDARD
57 .decode(&self.data_base64)
58 .map_err(|_| invalid("artifact put content is invalid base64"))?;
59 if bytes.is_empty() {
60 return Err(invalid("artifact put content is empty"));
61 }
62 Ok(bytes)
63 }
64}
65
66#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
67#[serde(rename_all = "camelCase", deny_unknown_fields)]
68pub struct PutArtifactResponse {
69 pub uri: String,
70 pub kind: ArtifactKind,
71 pub mime_type: String,
72 pub size_bytes: u64,
73 #[serde(default, skip_serializing_if = "Option::is_none")]
74 pub width: Option<u32>,
75 #[serde(default, skip_serializing_if = "Option::is_none")]
76 pub height: Option<u32>,
77 #[serde(default, skip_serializing_if = "Option::is_none")]
78 pub duration_millis: Option<u64>,
79}
80
81impl PutArtifactResponse {
82 pub fn validate(&self) -> Result<ArtifactReference, PutArtifactValidationError> {
83 let reference = ArtifactReference::parse(&self.uri)
84 .map_err(|_| invalid("artifact put URI is invalid"))?;
85 let metadata = reference.metadata();
86 if self.kind != metadata.kind()
87 || self.mime_type != metadata.mime_type()
88 || self.size_bytes != metadata.size_bytes()
89 || self.width != metadata.width()
90 || self.height != metadata.height()
91 || self.duration_millis != metadata.duration_millis()
92 {
93 return Err(invalid(
94 "artifact put metadata does not match its canonical reference",
95 ));
96 }
97 Ok(reference)
98 }
99}
100
101#[derive(Debug, Clone, PartialEq, Eq)]
102pub struct PutArtifactValidationError(&'static str);
103
104impl fmt::Display for PutArtifactValidationError {
105 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
106 formatter.write_str(self.0)
107 }
108}
109
110impl std::error::Error for PutArtifactValidationError {}
111
112fn invalid(message: &'static str) -> PutArtifactValidationError {
113 PutArtifactValidationError(message)
114}
115
116#[cfg(test)]
117mod tests {
118 use super::*;
119
120 #[test]
121 fn request_rejects_kind_mime_mismatch() {
122 let request = PutArtifactRequest {
123 kind: ArtifactKind::Audio,
124 mime_type: "image/png".to_string(),
125 data_base64: "eA==".to_string(),
126 duration_millis: None,
127 transform: None,
128 };
129 assert!(request.decode().is_err());
130 }
131}