pub mod request {
use serde::Serialize;
use url::Url;
use crate::{
errors::OapiError,
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct ModerationRequest {
pub input: ModerationInput,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(flatten, skip_serializing_if = "Option::is_none")]
pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Debug, Serialize, Clone)]
#[serde(untagged)]
pub enum ModerationInput {
String(String),
StringArray(Vec<String>),
}
impl Default for ModerationInput {
fn default() -> Self {
Self::String(String::new())
}
}
impl ModerationRequest {
pub fn is_streaming(&self) -> bool {
false
}
}
impl Post for ModerationRequest {
fn is_streaming(&self) -> bool {
false
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
let mut url =
Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
url.path_segments_mut()
.map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
.push("moderations");
Ok(url.to_string())
}
}
impl PostNoStream for ModerationRequest {
type Response = super::response::ModerationCreateResponse;
}
}
pub mod response {
use serde::Deserialize;
#[derive(Debug, Deserialize, Clone)]
pub struct ModerationCreateResponse {
pub id: String,
pub model: String,
pub results: Vec<Moderation>,
}
#[derive(Debug, Deserialize, Clone)]
pub struct Moderation {
pub categories: Categories,
#[serde(alias = "category_applied_input_types")]
pub category_applied_input_types: Option<CategoryAppliedInputTypes>,
#[serde(alias = "category_scores")]
pub category_scores: CategoryScores,
pub flagged: bool,
}
#[derive(Debug, Deserialize, Clone, Default)]
pub struct Categories {
pub harassment: Option<bool>,
#[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
pub harassment_threatening: Option<bool>,
pub hate: Option<bool>,
#[serde(rename = "hate/threatening", alias = "hate_threatening")]
pub hate_threatening: Option<bool>,
pub illicit: Option<bool>,
#[serde(rename = "illicit/violent", alias = "illicit_violent")]
pub illicit_violent: Option<bool>,
#[serde(rename = "self-harm", alias = "self_harm")]
pub self_harm: Option<bool>,
#[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
pub self_harm_instructions: Option<bool>,
#[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
pub self_harm_intent: Option<bool>,
pub sexual: Option<bool>,
#[serde(rename = "sexual/minors", alias = "sexual_minors")]
pub sexual_minors: Option<bool>,
pub violence: Option<bool>,
#[serde(rename = "violence/graphic", alias = "violence_graphic")]
pub violence_graphic: Option<bool>,
}
#[derive(Debug, Deserialize, Clone, Default)]
pub struct CategoryAppliedInputTypes {
pub harassment: Option<Vec<String>>,
#[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
pub harassment_threatening: Option<Vec<String>>,
pub hate: Option<Vec<String>>,
#[serde(rename = "hate/threatening", alias = "hate_threatening")]
pub hate_threatening: Option<Vec<String>>,
pub illicit: Option<Vec<String>>,
#[serde(rename = "illicit/violent", alias = "illicit_violent")]
pub illicit_violent: Option<Vec<String>>,
#[serde(rename = "self-harm", alias = "self_harm")]
pub self_harm: Option<Vec<String>>,
#[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
pub self_harm_instructions: Option<Vec<String>>,
#[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
pub self_harm_intent: Option<Vec<String>>,
pub sexual: Option<Vec<String>>,
#[serde(rename = "sexual/minors", alias = "sexual_minors")]
pub sexual_minors: Option<Vec<String>>,
pub violence: Option<Vec<String>>,
#[serde(rename = "violence/graphic", alias = "violence_graphic")]
pub violence_graphic: Option<Vec<String>>,
}
#[derive(Debug, Deserialize, Clone, Default)]
pub struct CategoryScores {
pub harassment: Option<f32>,
#[serde(rename = "harassment/threatening", alias = "harassment_threatening")]
pub harassment_threatening: Option<f32>,
pub hate: Option<f32>,
#[serde(rename = "hate/threatening", alias = "hate_threatening")]
pub hate_threatening: Option<f32>,
pub illicit: Option<f32>,
#[serde(rename = "illicit/violent", alias = "illicit_violent")]
pub illicit_violent: Option<f32>,
#[serde(rename = "self-harm", alias = "self_harm")]
pub self_harm: Option<f32>,
#[serde(rename = "self-harm/instructions", alias = "self_harm_instructions")]
pub self_harm_instructions: Option<f32>,
#[serde(rename = "self-harm/intent", alias = "self_harm_intent")]
pub self_harm_intent: Option<f32>,
pub sexual: Option<f32>,
#[serde(rename = "sexual/minors", alias = "sexual_minors")]
pub sexual_minors: Option<f32>,
pub violence: Option<f32>,
#[serde(rename = "violence/graphic", alias = "violence_graphic")]
pub violence_graphic: Option<f32>,
}
crate::impl_from_str!(ModerationCreateResponse);
}
#[cfg(test)]
mod tests {
use super::request::{ModerationInput, ModerationRequest};
use crate::rest::post::Post;
#[test]
fn request_serialization() {
let request = ModerationRequest {
input: ModerationInput::String("I want to harm someone.".to_string()),
model: Some("omni-moderation-latest".to_string()),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""input":"I want to harm someone.""#),
"json: {json}"
);
assert!(
json.contains(r#""model":"omni-moderation-latest""#),
"json: {json}"
);
}
#[test]
fn test_build_url() {
let request = ModerationRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/moderations");
}
#[test]
fn parse_response() {
let content = r#"{
"id": "modr-0a1b2c3d4e5f6a7b8c9d0e1f",
"model": "omni-moderation-latest",
"results": [
{
"flagged": true,
"categories": {
"harassment": false,
"harassment/threatening": false,
"sexual": false,
"self-harm/instructions": true,
"violence": true
},
"category_scores": {
"harassment": 0.01,
"harassment/threatening": 0.02,
"sexual": 0.0001,
"self-harm/instructions": 0.87,
"violence": 0.93
},
"category_applied_input_types": {
"sexual": ["text"],
"violence": ["text", "image"]
}
}
]
}"#;
let response: super::response::ModerationCreateResponse = content.parse().unwrap();
assert_eq!(response.results.len(), 1);
let result = &response.results[0];
assert!(result.flagged);
assert_eq!(result.categories.violence, Some(true));
assert_eq!(result.categories.self_harm_instructions, Some(true));
assert_eq!(result.categories.harassment, Some(false));
assert_eq!(result.category_scores.violence, Some(0.93));
assert_eq!(result.category_scores.self_harm_instructions, Some(0.87));
let applied = result.category_applied_input_types.as_ref().unwrap();
assert_eq!(
applied.violence,
Some(vec!["text".to_string(), "image".to_string()])
);
}
}