1pub mod request {
10 use serde::{Deserialize, Serialize};
11 use url::Url;
12
13 use crate::{
14 errors::OapiError,
15 rest::post::{Post, PostNoStream},
16 };
17
18 #[derive(Debug, Serialize, Deserialize, Default, Clone)]
20 pub struct ModerationRequest {
21 pub input: ModerationInput,
24 #[serde(skip_serializing_if = "Option::is_none")]
28 pub model: Option<String>,
29 #[serde(flatten, default, skip_serializing_if = "Option::is_none")]
31 pub extra_body_map: Option<serde_json::Map<String, serde_json::Value>>,
32 }
33
34 #[derive(Debug, Serialize, Deserialize, Clone)]
36 #[serde(untagged)]
37 pub enum ModerationInput {
38 String(String),
40 StringArray(Vec<String>),
42 MultiModal(Vec<ModerationInputContent>),
44 }
45
46 #[derive(Debug, Serialize, Deserialize, Clone)]
48 #[serde(tag = "type", rename_all = "snake_case")]
49 pub enum ModerationInputContent {
50 Text {
52 text: String,
54 },
55 ImageUrl {
57 image_url: ModerationImageUrl,
60 },
61 }
62
63 #[derive(Debug, Serialize, Deserialize, Clone)]
66 pub struct ModerationImageUrl {
67 pub url: String,
69 }
70
71 impl Default for ModerationInput {
72 fn default() -> Self {
73 Self::String(String::new())
74 }
75 }
76
77 impl ModerationRequest {
78 pub fn is_streaming(&self) -> bool {
79 false
80 }
81 }
82
83 impl Post for ModerationRequest {
84 fn is_streaming(&self) -> bool {
85 false
86 }
87
88 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
92 let mut url =
93 Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
94 url.path_segments_mut()
95 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
96 .push("moderations");
97
98 Ok(url.to_string())
99 }
100 }
101
102 impl PostNoStream for ModerationRequest {
103 type Response = super::response::ModerationCreateResponse;
104 }
105}
106
107pub mod response {
108 use serde::{Deserialize, Serialize};
109
110 #[derive(Debug, Deserialize, Serialize, Clone)]
112 pub struct ModerationCreateResponse {
113 pub id: String,
115 pub model: String,
117 pub results: Vec<Moderation>,
119 }
120
121 #[derive(Debug, Deserialize, Serialize, Clone)]
123 pub struct Moderation {
124 pub categories: Categories,
126 #[serde(alias = "category_applied_input_types")]
129 pub category_applied_input_types: Option<CategoryAppliedInputTypes>,
130 #[serde(alias = "category_scores")]
133 pub category_scores: CategoryScores,
134 pub flagged: bool,
136 }
137
138 #[derive(Debug, Deserialize, Serialize, Clone, Default)]
144 pub struct Categories {
145 pub harassment: Option<bool>,
146 #[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
147 pub harassment_threatening: Option<bool>,
148 pub hate: Option<bool>,
149 #[serde(rename = "hate/threatening", alias = "hate_threatening")]
150 pub hate_threatening: Option<bool>,
151 pub illicit: Option<bool>,
152 #[serde(rename = "illicit/violent", alias = "illicit_violent")]
153 pub illicit_violent: Option<bool>,
154 #[serde(rename = "self-harm", alias = "self_harm")]
155 pub self_harm: Option<bool>,
156 #[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
157 pub self_harm_instructions: Option<bool>,
158 #[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
159 pub self_harm_intent: Option<bool>,
160 pub sexual: Option<bool>,
161 #[serde(rename = "sexual/minors", alias = "sexual_minors")]
162 pub sexual_minors: Option<bool>,
163 pub violence: Option<bool>,
164 #[serde(rename = "violence/graphic", alias = "violence_graphic")]
165 pub violence_graphic: Option<bool>,
166 }
167
168 #[derive(Debug, Deserialize, Serialize, Clone, Default)]
171 pub struct CategoryAppliedInputTypes {
172 pub harassment: Option<Vec<String>>,
173 #[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
174 pub harassment_threatening: Option<Vec<String>>,
175 pub hate: Option<Vec<String>>,
176 #[serde(rename = "hate/threatening", alias = "hate_threatening")]
177 pub hate_threatening: Option<Vec<String>>,
178 pub illicit: Option<Vec<String>>,
179 #[serde(rename = "illicit/violent", alias = "illicit_violent")]
180 pub illicit_violent: Option<Vec<String>>,
181 #[serde(rename = "self-harm", alias = "self_harm")]
182 pub self_harm: Option<Vec<String>>,
183 #[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
184 pub self_harm_instructions: Option<Vec<String>>,
185 #[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
186 pub self_harm_intent: Option<Vec<String>>,
187 pub sexual: Option<Vec<String>>,
188 #[serde(rename = "sexual/minors", alias = "sexual_minors")]
189 pub sexual_minors: Option<Vec<String>>,
190 pub violence: Option<Vec<String>>,
191 #[serde(rename = "violence/graphic", alias = "violence_graphic")]
192 pub violence_graphic: Option<Vec<String>>,
193 }
194
195 #[derive(Debug, Deserialize, Serialize, Clone, Default)]
198 pub struct CategoryScores {
199 pub harassment: Option<f32>,
200 #[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
201 pub harassment_threatening: Option<f32>,
202 pub hate: Option<f32>,
203 #[serde(rename = "hate/threatening", alias = "hate_threatening")]
204 pub hate_threatening: Option<f32>,
205 pub illicit: Option<f32>,
206 #[serde(rename = "illicit/violent", alias = "illicit_violent")]
207 pub illicit_violent: Option<f32>,
208 #[serde(rename = "self-harm", alias = "self_harm")]
209 pub self_harm: Option<f32>,
210 #[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
211 pub self_harm_instructions: Option<f32>,
212 #[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
213 pub self_harm_intent: Option<f32>,
214 pub sexual: Option<f32>,
215 #[serde(rename = "sexual/minors", alias = "sexual_minors")]
216 pub sexual_minors: Option<f32>,
217 pub violence: Option<f32>,
218 #[serde(rename = "violence/graphic", alias = "violence_graphic")]
219 pub violence_graphic: Option<f32>,
220 }
221
222 crate::impl_from_str!(ModerationCreateResponse);
223}
224
225#[cfg(test)]
226mod tests {
227 use super::request::{
228 ModerationImageUrl, ModerationInput, ModerationInputContent, ModerationRequest,
229 };
230 use crate::rest::post::Post;
231
232 #[test]
234 fn request_serialization() {
235 let request = ModerationRequest {
236 input: ModerationInput::String("I want to harm someone.".to_string()),
237 model: Some("omni-moderation-latest".to_string()),
238 ..Default::default()
239 };
240
241 let json = serde_json::to_string(&request).unwrap();
242 assert!(
243 json.contains(r#""input":"I want to harm someone.""#),
244 "json: {json}"
245 );
246 assert!(
247 json.contains(r#""model":"omni-moderation-latest""#),
248 "json: {json}"
249 );
250 }
251
252 #[test]
255 fn multimodal_input_serialization() {
256 let request = ModerationRequest {
257 input: ModerationInput::MultiModal(vec![
258 ModerationInputContent::Text {
259 text: "hello".to_string(),
260 },
261 ModerationInputContent::ImageUrl {
262 image_url: ModerationImageUrl {
263 url: "https://example.com/a.png".to_string(),
264 },
265 },
266 ]),
267 ..Default::default()
268 };
269
270 let json = serde_json::to_string(&request).unwrap();
271 assert!(
272 json.contains(r#"{"type":"text","text":"hello"}"#),
273 "json: {json}"
274 );
275 assert!(
276 json.contains(
277 r#"{"type":"image_url","image_url":{"url":"https://example.com/a.png"}}"#
278 ),
279 "json: {json}"
280 );
281 }
282
283 #[test]
284 fn test_build_url() {
285 let request = ModerationRequest::default();
286 let url = request.build_url("https://api.openai.com/v1/").unwrap();
287 assert_eq!(url, "https://api.openai.com/v1/moderations");
288 }
289
290 #[test]
299 fn parse_response() {
300 let content = r#"{
301 "id": "modr-0a1b2c3d4e5f6a7b8c9d0e1f",
302 "model": "omni-moderation-latest",
303 "results": [
304 {
305 "flagged": true,
306 "categories": {
307 "harassment": false,
308 "harassment/threatening": false,
309 "sexual": false,
310 "self-harm/instructions": true,
311 "violence": true
312 },
313 "category_scores": {
314 "harassment": 0.01,
315 "harassment/threatening": 0.02,
316 "sexual": 0.0001,
317 "self-harm/instructions": 0.87,
318 "violence": 0.93
319 },
320 "category_applied_input_types": {
321 "sexual": ["text"],
322 "violence": ["text", "image"]
323 }
324 }
325 ]
326 }"#;
327
328 let response: super::response::ModerationCreateResponse = content.parse().unwrap();
329 assert_eq!(response.results.len(), 1);
330
331 let result = &response.results[0];
332 assert!(result.flagged);
333 assert_eq!(result.categories.violence, Some(true));
334 assert_eq!(result.categories.self_harm_instructions, Some(true));
335 assert_eq!(result.categories.harassment, Some(false));
336 assert_eq!(result.category_scores.violence, Some(0.93));
337 assert_eq!(result.category_scores.self_harm_instructions, Some(0.87));
338 let applied = result.category_applied_input_types.as_ref().unwrap();
339 assert_eq!(
340 applied.violence,
341 Some(vec!["text".to_string(), "image".to_string()])
342 );
343 }
344}