Skip to main content

openai_interface/moderations/
mod.rs

1//! Given text and/or image inputs, classifies if those inputs are potentially
2//! harmful.
3//!
4//! > ![warn] This module is untested!
5//! > OpenAI-compatible providers (DeepSeek, Qwen) are not known to implement
6//! > this endpoint, and no OpenAI API key was available for testing. If you
7//! > encounter any issues, please report them on the repository.
8
9pub mod request {
10    use serde::Serialize;
11    use url::Url;
12
13    use crate::{
14        errors::OapiError,
15        rest::post::{Post, PostNoStream},
16    };
17
18    /// Classifies if text and/or image inputs are potentially harmful.
19    #[derive(Debug, Serialize, Default, Clone)]
20    pub struct ModerationRequest {
21        /// Input (or inputs) to classify. Can be a single string, an array of
22        /// strings, or an array of multi-modal input objects. Multi-modal
23        /// input objects are not covered by this type yet; pass them through
24        /// [`Self::extra_body`] by overriding `input` there.
25        pub input: ModerationInput,
26        /// The content moderation model you would like to use, e.g.
27        /// `omni-moderation-latest`. Defaults to `omni-moderation-latest`
28        /// when not set.
29        #[serde(skip_serializing_if = "Option::is_none")]
30        pub model: Option<String>,
31        /// Add additional JSON properties to the request.
32        #[serde(flatten, skip_serializing_if = "Option::is_none")]
33        pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
34    }
35
36    /// The input (or inputs) to classify.
37    #[derive(Debug, Serialize, Clone)]
38    #[serde(untagged)]
39    pub enum ModerationInput {
40        /// A single string to classify.
41        String(String),
42        /// An array of strings to classify.
43        StringArray(Vec<String>),
44    }
45
46    impl Default for ModerationInput {
47        fn default() -> Self {
48            Self::String(String::new())
49        }
50    }
51
52    impl ModerationRequest {
53        pub fn is_streaming(&self) -> bool {
54            false
55        }
56    }
57
58    impl Post for ModerationRequest {
59        fn is_streaming(&self) -> bool {
60            false
61        }
62
63        /// Builds the URL for the request.
64        ///
65        /// `base_url` should be like <https://api.openai.com/v1>
66        fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
67            let mut url =
68                Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
69            url.path_segments_mut()
70                .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
71                .push("moderations");
72
73            Ok(url.to_string())
74        }
75    }
76
77    impl PostNoStream for ModerationRequest {
78        type Response = super::response::ModerationCreateResponse;
79    }
80}
81
82pub mod response {
83    use serde::Deserialize;
84
85    /// Represents if a given text input is potentially harmful.
86    #[derive(Debug, Deserialize, Clone)]
87    pub struct ModerationCreateResponse {
88        /// The unique identifier for the moderation request.
89        pub id: String,
90        /// The model used to generate the moderation results.
91        pub model: String,
92        /// A list of moderation objects.
93        pub results: Vec<Moderation>,
94    }
95
96    /// A moderation object.
97    #[derive(Debug, Deserialize, Clone)]
98    pub struct Moderation {
99        /// A list of the categories, and whether they are flagged or not.
100        pub categories: Categories,
101        /// A list of the categories along with the input type(s) that the
102        /// score applies to.
103        #[serde(alias = "category_applied_input_types")]
104        pub category_applied_input_types: Option<CategoryAppliedInputTypes>,
105        /// A list of the categories along with their scores as predicted by
106        /// the model.
107        #[serde(alias = "category_scores")]
108        pub category_scores: CategoryScores,
109        /// Whether any of the categories are flagged.
110        pub flagged: bool,
111    }
112
113    /// A list of the categories, and whether they are flagged or not.
114    ///
115    /// Fields are optional so that responses of different moderation model
116    /// generations (which add or remove categories over time) can be
117    /// deserialized.
118    #[derive(Debug, Deserialize, Clone, Default)]
119    pub struct Categories {
120        pub harassment: Option<bool>,
121        #[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
122        pub harassment_threatening: Option<bool>,
123        pub hate: Option<bool>,
124        #[serde(rename = "hate/threatening", alias = "hate_threatening")]
125        pub hate_threatening: Option<bool>,
126        pub illicit: Option<bool>,
127        #[serde(rename = "illicit/violent", alias = "illicit_violent")]
128        pub illicit_violent: Option<bool>,
129        #[serde(rename = "self-harm", alias = "self_harm")]
130        pub self_harm: Option<bool>,
131        #[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
132        pub self_harm_instructions: Option<bool>,
133        #[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
134        pub self_harm_intent: Option<bool>,
135        pub sexual: Option<bool>,
136        #[serde(rename = "sexual/minors", alias = "sexual_minors")]
137        pub sexual_minors: Option<bool>,
138        pub violence: Option<bool>,
139        #[serde(rename = "violence/graphic", alias = "violence_graphic")]
140        pub violence_graphic: Option<bool>,
141    }
142
143    /// A list of the categories along with the input type(s) that the score
144    /// applies to, e.g. `["text"]` or `["text", "image"]`.
145    #[derive(Debug, Deserialize, Clone, Default)]
146    pub struct CategoryAppliedInputTypes {
147        pub harassment: Option<Vec<String>>,
148        #[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
149        pub harassment_threatening: Option<Vec<String>>,
150        pub hate: Option<Vec<String>>,
151        #[serde(rename = "hate/threatening", alias = "hate_threatening")]
152        pub hate_threatening: Option<Vec<String>>,
153        pub illicit: Option<Vec<String>>,
154        #[serde(rename = "illicit/violent", alias = "illicit_violent")]
155        pub illicit_violent: Option<Vec<String>>,
156        #[serde(rename = "self-harm", alias = "self_harm")]
157        pub self_harm: Option<Vec<String>>,
158        #[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
159        pub self_harm_instructions: Option<Vec<String>>,
160        #[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
161        pub self_harm_intent: Option<Vec<String>>,
162        pub sexual: Option<Vec<String>>,
163        #[serde(rename = "sexual/minors", alias = "sexual_minors")]
164        pub sexual_minors: Option<Vec<String>>,
165        pub violence: Option<Vec<String>>,
166        #[serde(rename = "violence/graphic", alias = "violence_graphic")]
167        pub violence_graphic: Option<Vec<String>>,
168    }
169
170    /// A list of the categories along with their scores as predicted by the
171    /// model.
172    #[derive(Debug, Deserialize, Clone, Default)]
173    pub struct CategoryScores {
174        pub harassment: Option<f32>,
175        #[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
176        pub harassment_threatening: Option<f32>,
177        pub hate: Option<f32>,
178        #[serde(rename = "hate/threatening", alias = "hate_threatening")]
179        pub hate_threatening: Option<f32>,
180        pub illicit: Option<f32>,
181        #[serde(rename = "illicit/violent", alias = "illicit_violent")]
182        pub illicit_violent: Option<f32>,
183        #[serde(rename = "self-harm", alias = "self_harm")]
184        pub self_harm: Option<f32>,
185        #[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
186        pub self_harm_instructions: Option<f32>,
187        #[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
188        pub self_harm_intent: Option<f32>,
189        pub sexual: Option<f32>,
190        #[serde(rename = "sexual/minors", alias = "sexual_minors")]
191        pub sexual_minors: Option<f32>,
192        pub violence: Option<f32>,
193        #[serde(rename = "violence/graphic", alias = "violence_graphic")]
194        pub violence_graphic: Option<f32>,
195    }
196
197    crate::impl_from_str!(ModerationCreateResponse);
198}
199
200#[cfg(test)]
201mod tests {
202    use super::request::{ModerationInput, ModerationRequest};
203    use crate::rest::post::Post;
204
205    /// Serializes a simple text moderation request.
206    #[test]
207    fn request_serialization() {
208        let request = ModerationRequest {
209            input: ModerationInput::String("I want to harm someone.".to_string()),
210            model: Some("omni-moderation-latest".to_string()),
211            ..Default::default()
212        };
213
214        let json = serde_json::to_string(&request).unwrap();
215        assert!(
216            json.contains(r#""input":"I want to harm someone.""#),
217            "json: {json}"
218        );
219        assert!(
220            json.contains(r#""model":"omni-moderation-latest""#),
221            "json: {json}"
222        );
223    }
224
225    #[test]
226    fn test_build_url() {
227        let request = ModerationRequest::default();
228        let url = request.build_url("https://api.openai.com/v1/").unwrap();
229        assert_eq!(url, "https://api.openai.com/v1/moderations");
230    }
231
232    /// Deserializes a moderation response.
233    ///
234    /// No accessible provider implements this endpoint, so this fixture is
235    /// NOT captured from a live response. The structure (category names,
236    /// slash-separated keys, aliases, `flagged`) follows the schema of
237    /// openai-python `types/moderation.py`; the values are constructed for
238    /// the test. Fields absent from the fixture exercise the optional-field
239    /// tolerance of the types.
240    #[test]
241    fn parse_response() {
242        let content = r#"{
243            "id": "modr-0a1b2c3d4e5f6a7b8c9d0e1f",
244            "model": "omni-moderation-latest",
245            "results": [
246                {
247                    "flagged": true,
248                    "categories": {
249                        "harassment": false,
250                        "harassment/threatening": false,
251                        "sexual": false,
252                        "self-harm/instructions": true,
253                        "violence": true
254                    },
255                    "category_scores": {
256                        "harassment": 0.01,
257                        "harassment/threatening": 0.02,
258                        "sexual": 0.0001,
259                        "self-harm/instructions": 0.87,
260                        "violence": 0.93
261                    },
262                    "category_applied_input_types": {
263                        "sexual": ["text"],
264                        "violence": ["text", "image"]
265                    }
266                }
267            ]
268        }"#;
269
270        let response: super::response::ModerationCreateResponse = content.parse().unwrap();
271        assert_eq!(response.results.len(), 1);
272
273        let result = &response.results[0];
274        assert!(result.flagged);
275        assert_eq!(result.categories.violence, Some(true));
276        assert_eq!(result.categories.self_harm_instructions, Some(true));
277        assert_eq!(result.categories.harassment, Some(false));
278        assert_eq!(result.category_scores.violence, Some(0.93));
279        assert_eq!(result.category_scores.self_harm_instructions, Some(0.87));
280        let applied = result.category_applied_input_types.as_ref().unwrap();
281        assert_eq!(
282            applied.violence,
283            Some(vec!["text".to_string(), "image".to_string()])
284        );
285    }
286}