Skip to main content

zai_rs/model/moderation/
models.rs

1//! Content-moderation wire types.
2
3use serde::{Deserialize, Deserializer, Serialize, de::Error as _, ser::SerializeStruct};
4use validator::Validate;
5
6/// Content moderation model type.
7#[derive(Debug, Clone, Serialize, Deserialize, Default)]
8pub enum ModerationModel {
9    /// Current moderation model.
10    #[serde(rename = "moderation")]
11    #[default]
12    Moderation,
13}
14
15/// Moderation input content.
16#[derive(Clone, Serialize, Deserialize)]
17#[serde(untagged)]
18pub enum ModerationInput {
19    /// Text content.
20    Text(String),
21    /// Multimedia content with its type and URL.
22    Multimedia(MultimediaInput),
23    /// Multiple structured text and multimedia items.
24    Items(Vec<ModerationItem>),
25}
26
27impl std::fmt::Debug for ModerationInput {
28    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
29        match self {
30            Self::Text(_) => formatter.debug_tuple("Text").field(&"[REDACTED]").finish(),
31            Self::Multimedia(value) => formatter.debug_tuple("Multimedia").field(value).finish(),
32            Self::Items(values) => formatter
33                .debug_struct("Items")
34                .field("len", &values.len())
35                .finish(),
36        }
37    }
38}
39
40/// Multimedia input for content moderation.
41#[derive(Clone, Validate)]
42pub struct MultimediaInput {
43    /// Content type.
44    pub content_type: MediaType,
45    /// URL of the multimedia content.
46    #[validate(url)]
47    pub url: String,
48}
49
50impl std::fmt::Debug for MultimediaInput {
51    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
52        formatter
53            .debug_struct("MultimediaInput")
54            .field("content_type", &self.content_type)
55            .field("url", &"[REDACTED]")
56            .finish()
57    }
58}
59
60#[derive(Serialize, Deserialize)]
61struct UrlValue<T> {
62    url: T,
63}
64
65impl Serialize for MultimediaInput {
66    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
67    where
68        S: serde::Serializer,
69    {
70        let mut state = serializer.serialize_struct("MultimediaInput", 2)?;
71        state.serialize_field("type", &self.content_type)?;
72        let value = UrlValue { url: &self.url };
73        state.serialize_field(self.content_type.content_key(), &value)?;
74        state.end()
75    }
76}
77
78impl<'de> Deserialize<'de> for MultimediaInput {
79    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
80    where
81        D: serde::Deserializer<'de>,
82    {
83        #[derive(Deserialize)]
84        struct WireInput {
85            #[serde(rename = "type")]
86            content_type: MediaType,
87            image_url: Option<UrlValue<String>>,
88            audio_url: Option<UrlValue<String>>,
89            video_url: Option<UrlValue<String>>,
90        }
91
92        let wire = WireInput::deserialize(deserializer)?;
93        let value = match wire.content_type {
94            MediaType::Image => wire.image_url,
95            MediaType::Audio => wire.audio_url,
96            MediaType::Video => wire.video_url,
97        }
98        .ok_or_else(|| {
99            serde::de::Error::custom(format_args!(
100                "missing {} object for moderation input",
101                wire.content_type.content_key()
102            ))
103        })?;
104
105        Ok(Self {
106            content_type: wire.content_type,
107            url: value.url,
108        })
109    }
110}
111
112/// Media types supported for moderation.
113#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
114pub enum MediaType {
115    /// Image content.
116    #[serde(rename = "image_url")]
117    Image,
118    /// Audio content.
119    #[serde(rename = "audio_url")]
120    Audio,
121    /// Video content.
122    #[serde(rename = "video_url")]
123    Video,
124}
125
126impl MediaType {
127    const fn content_key(self) -> &'static str {
128        match self {
129            Self::Image => "image_url",
130            Self::Audio => "audio_url",
131            Self::Video => "video_url",
132        }
133    }
134}
135
136/// One structured item in a batch moderation request.
137#[derive(Debug, Clone, Serialize, Deserialize)]
138#[serde(untagged)]
139pub enum ModerationItem {
140    /// Text item encoded as `{ "type": "text", "text": "..." }`.
141    Text(ModerationTextItem),
142    /// Image, audio, or video item.
143    Multimedia(MultimediaInput),
144}
145
146/// Structured text input used inside a batch moderation request.
147#[derive(Clone, Serialize, Deserialize)]
148pub struct ModerationTextItem {
149    #[serde(rename = "type")]
150    content_type: TextMediaType,
151    /// Text to moderate.
152    pub text: String,
153}
154
155impl std::fmt::Debug for ModerationTextItem {
156    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
157        formatter
158            .debug_struct("ModerationTextItem")
159            .field("content_type", &self.content_type)
160            .field("text", &"[REDACTED]")
161            .finish()
162    }
163}
164
165#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
166enum TextMediaType {
167    #[serde(rename = "text")]
168    Text,
169}
170
171impl ModerationItem {
172    /// Create a structured text item.
173    pub fn text(text: impl Into<String>) -> Self {
174        Self::Text(ModerationTextItem {
175            content_type: TextMediaType::Text,
176            text: text.into(),
177        })
178    }
179
180    /// Create a structured multimedia item.
181    pub fn multimedia(content_type: MediaType, url: impl Into<String>) -> Self {
182        Self::Multimedia(MultimediaInput {
183            content_type,
184            url: url.into(),
185        })
186    }
187}
188
189/// Content moderation request.
190#[derive(Debug, Clone, Serialize, Deserialize)]
191pub struct ModerationRequest {
192    /// Moderation model.
193    #[serde(default)]
194    pub model: ModerationModel,
195    /// Content to moderate.
196    pub input: ModerationInput,
197}
198
199impl ModerationRequest {
200    /// Create a new moderation request with text content.
201    ///
202    /// The service accepts at most 2,000 Unicode characters.
203    pub fn new_text(text: impl Into<String>) -> Self {
204        Self {
205            model: ModerationModel::default(),
206            input: ModerationInput::Text(text.into()),
207        }
208    }
209
210    /// Create a new moderation request with multimedia content.
211    pub fn new_multimedia(content_type: MediaType, url: impl Into<String>) -> Self {
212        Self {
213            model: ModerationModel::default(),
214            input: ModerationInput::Multimedia(MultimediaInput {
215                content_type,
216                url: url.into(),
217            }),
218        }
219    }
220
221    /// Create a request containing multiple structured items.
222    pub fn new_items(items: Vec<ModerationItem>) -> Self {
223        Self {
224            model: ModerationModel::default(),
225            input: ModerationInput::Items(items),
226        }
227    }
228
229    /// Validate request constraints before dispatch.
230    pub fn validate(&self) -> Result<(), validator::ValidationErrors> {
231        let mut errors = validator::ValidationErrors::new();
232
233        match &self.input {
234            ModerationInput::Text(text) => validate_text(text, &mut errors),
235            ModerationInput::Multimedia(multimedia) => {
236                validate_multimedia(multimedia, &mut errors);
237            },
238            ModerationInput::Items(items) if items.is_empty() => {
239                errors.add("input", validator::ValidationError::new("items_required"));
240            },
241            ModerationInput::Items(items) => {
242                for item in items {
243                    match item {
244                        ModerationItem::Text(item) => validate_text(&item.text, &mut errors),
245                        ModerationItem::Multimedia(item) => {
246                            validate_multimedia(item, &mut errors);
247                        },
248                    }
249                }
250            },
251        }
252
253        if errors.is_empty() {
254            Ok(())
255        } else {
256            Err(errors)
257        }
258    }
259}
260
261#[cfg(test)]
262mod request_debug_tests {
263    use super::*;
264
265    #[test]
266    fn request_debug_redacts_text_and_media_urls() {
267        let text = format!(
268            "{:?}",
269            ModerationRequest::new_text("private moderation text")
270        );
271        assert!(!text.contains("private moderation text"));
272
273        let media = format!(
274            "{:?}",
275            ModerationRequest::new_multimedia(
276                MediaType::Image,
277                "https://private.example/image.png"
278            )
279        );
280        assert!(!media.contains("private.example"));
281
282        let items = format!(
283            "{:?}",
284            ModerationRequest::new_items(vec![ModerationItem::text("private item")])
285        );
286        assert!(!items.contains("private item"));
287        assert!(items.contains("len: 1"));
288    }
289}
290
291fn validate_text(text: &str, errors: &mut validator::ValidationErrors) {
292    if text.trim().is_empty() {
293        errors.add("input", validator::ValidationError::new("text_required"));
294    } else if text.chars().count() > 2000 {
295        errors.add(
296            "input",
297            validator::ValidationError::new("text_length_exceeded"),
298        );
299    }
300}
301
302fn validate_multimedia(multimedia: &MultimediaInput, errors: &mut validator::ValidationErrors) {
303    let valid_url = multimedia
304        .url
305        .parse::<url::Url>()
306        .is_ok_and(|url| matches!(url.scheme(), "http" | "https"));
307    if !valid_url {
308        errors.add("input", validator::ValidationError::new("invalid_url"));
309    }
310}
311
312/// Risk level for moderated content.
313#[derive(Debug, Clone, Default, Serialize, Deserialize)]
314#[non_exhaustive]
315pub enum RiskLevel {
316    /// No risk detected.
317    #[default]
318    #[serde(rename = "PASS")]
319    Pass,
320    /// Suspicious content that requires review.
321    #[serde(rename = "REVIEW")]
322    Review,
323    /// Policy-violating content that should be rejected.
324    #[serde(rename = "REJECT")]
325    Reject,
326    /// A value introduced by a newer service revision.
327    #[serde(other)]
328    Unknown,
329}
330
331/// Moderation result for a single content item.
332///
333/// The frozen OpenAPI schema does not mark any item property as required.
334#[derive(Debug, Clone, Serialize, Deserialize)]
335pub struct ModerationResult {
336    /// Type of content that was moderated.
337    #[serde(rename = "content_type", skip_serializing_if = "Option::is_none")]
338    pub content_type: Option<String>,
339    /// Assessed risk level.
340    #[serde(rename = "risk_level", skip_serializing_if = "Option::is_none")]
341    pub risk_level: Option<RiskLevel>,
342    /// Detected risk types.
343    #[serde(rename = "risk_type", skip_serializing_if = "Option::is_none")]
344    pub risk_types: Option<Vec<String>>,
345}
346
347/// Usage statistics for moderation API.
348#[derive(Debug, Clone, Serialize, Deserialize)]
349pub struct ModerationUsage {
350    /// Text-moderation usage statistics.
351    #[serde(rename = "moderation_text", skip_serializing_if = "Option::is_none")]
352    pub moderation_text: Option<ModerationTextUsage>,
353}
354
355/// Text moderation usage statistics.
356#[derive(Debug, Clone, Serialize, Deserialize)]
357pub struct ModerationTextUsage {
358    /// Number of text-moderation calls.
359    #[serde(rename = "call_count", skip_serializing_if = "Option::is_none")]
360    pub call_count: Option<f64>,
361}
362
363/// Content moderation response.
364#[derive(Debug, Clone, Serialize)]
365pub struct ModerationResponse {
366    /// Task identifier.
367    #[serde(skip_serializing_if = "Option::is_none")]
368    pub id: Option<String>,
369    /// Request creation time as a Unix timestamp in seconds.
370    #[serde(skip_serializing_if = "Option::is_none")]
371    pub created: Option<u64>,
372    /// Request identifier.
373    #[serde(
374        rename = "request_id",
375        default,
376        skip_serializing_if = "Option::is_none",
377        deserialize_with = "super::super::serde_helpers::optional_string_from_number_or_string"
378    )]
379    pub request_id: Option<String>,
380    /// Moderation results.
381    #[serde(rename = "result_list", skip_serializing_if = "Option::is_none")]
382    pub result_list: Option<Vec<ModerationResult>>,
383    /// Usage statistics.
384    #[serde(skip_serializing_if = "Option::is_none")]
385    pub usage: Option<ModerationUsage>,
386}
387
388#[derive(Deserialize)]
389struct ModerationResponseWire {
390    id: Option<String>,
391    created: Option<u64>,
392    #[serde(
393        default,
394        deserialize_with = "super::super::serde_helpers::optional_string_from_number_or_string"
395    )]
396    request_id: Option<String>,
397    result_list: Option<Vec<ModerationResult>>,
398    usage: Option<ModerationUsage>,
399}
400
401impl<'de> Deserialize<'de> for ModerationResponse {
402    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
403    where
404        D: Deserializer<'de>,
405    {
406        let wire = ModerationResponseWire::deserialize(deserializer)?;
407        if wire.id.is_none()
408            && wire.created.is_none()
409            && wire.request_id.is_none()
410            && wire.result_list.is_none()
411            && wire.usage.is_none()
412        {
413            return Err(D::Error::custom(
414                "moderation response contained no documented fields",
415            ));
416        }
417        Ok(Self {
418            id: wire.id,
419            created: wire.created,
420            request_id: wire.request_id,
421            result_list: wire.result_list,
422            usage: wire.usage,
423        })
424    }
425}
426
427#[cfg(test)]
428mod tests {
429    use super::*;
430
431    #[test]
432    fn multimedia_uses_the_nested_wire_shape() {
433        let request =
434            ModerationRequest::new_multimedia(MediaType::Image, "https://example.com/image.png");
435        let value = serde_json::to_value(request).unwrap();
436        assert_eq!(value["input"]["type"], "image_url");
437        assert_eq!(
438            value["input"]["image_url"]["url"],
439            "https://example.com/image.png"
440        );
441    }
442
443    #[test]
444    fn batch_items_use_structured_text_and_media_shapes() {
445        let request = ModerationRequest::new_items(vec![
446            ModerationItem::text("hello"),
447            ModerationItem::multimedia(MediaType::Audio, "https://example.com/audio.mp3"),
448        ]);
449        assert!(request.validate().is_ok());
450        let value = serde_json::to_value(request).unwrap();
451        assert_eq!(value["input"][0]["type"], "text");
452        assert_eq!(value["input"][1]["type"], "audio_url");
453        assert_eq!(
454            value["input"][1]["audio_url"]["url"],
455            "https://example.com/audio.mp3"
456        );
457    }
458
459    #[test]
460    fn text_limit_counts_characters_instead_of_utf8_bytes() {
461        assert!(ModerationRequest::new_text(" ").validate().is_err());
462        assert!(
463            ModerationRequest::new_text("测".repeat(2_000))
464                .validate()
465                .is_ok()
466        );
467        assert!(
468            ModerationRequest::new_text("测".repeat(2_001))
469                .validate()
470                .is_err()
471        );
472    }
473
474    #[test]
475    fn response_requires_one_documented_non_null_field() {
476        assert!(serde_json::from_str::<ModerationResponse>("{}").is_err());
477        assert!(serde_json::from_str::<ModerationResponse>(r#"{"id":null}"#).is_err());
478        let response: ModerationResponse = serde_json::from_str(r#"{"id":"mod-1"}"#).unwrap();
479        assert!(response.request_id.is_none());
480    }
481
482    #[test]
483    fn result_fields_are_optional_without_defaulting_risk_to_pass() {
484        let result: ModerationResult = serde_json::from_str("{}").unwrap();
485        assert!(result.content_type.is_none());
486        assert!(result.risk_level.is_none());
487        assert!(result.risk_types.is_none());
488    }
489}