use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::client::OpenAiClient;
use crate::error::OpenAiError;
pub struct Moderations<'a> {
pub(crate) client: &'a OpenAiClient,
}
impl Moderations<'_> {
pub async fn create(&self, request: &ModerationRequest) -> Result<Moderation, OpenAiError> {
self.client.post_json("/moderations", request).await
}
}
#[derive(Debug, Clone, Serialize)]
pub struct ModerationRequest {
input: ModerationInput,
#[serde(skip_serializing_if = "Option::is_none")]
model: Option<String>,
}
impl ModerationRequest {
pub fn new(input: impl Into<ModerationInput>) -> Self {
Self {
input: input.into(),
model: None,
}
}
pub fn model(mut self, model: impl Into<String>) -> Self {
self.model = Some(model.into());
self
}
}
#[derive(Debug, Clone, Serialize)]
#[serde(untagged)]
pub enum ModerationInput {
Text(String),
Texts(Vec<String>),
Items(Value),
}
impl From<&str> for ModerationInput {
fn from(text: &str) -> Self {
ModerationInput::Text(text.to_string())
}
}
impl From<String> for ModerationInput {
fn from(text: String) -> Self {
ModerationInput::Text(text)
}
}
impl From<Vec<String>> for ModerationInput {
fn from(texts: Vec<String>) -> Self {
ModerationInput::Texts(texts)
}
}
impl From<Vec<&str>> for ModerationInput {
fn from(texts: Vec<&str>) -> Self {
ModerationInput::Texts(texts.into_iter().map(str::to_string).collect())
}
}
#[derive(Debug, Clone, Deserialize)]
pub struct Moderation {
pub id: Option<String>,
pub model: Option<String>,
#[serde(default)]
pub results: Vec<ModerationResult>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct ModerationResult {
#[serde(default)]
pub flagged: bool,
#[serde(default)]
pub categories: HashMap<String, bool>,
#[serde(default)]
pub category_scores: HashMap<String, f64>,
pub category_applied_input_types: Option<Value>,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn serializes_request() {
let request = ModerationRequest::new("some text").model("omni-moderation-latest");
let value = serde_json::to_value(&request).unwrap();
assert_eq!(
value,
json!({"input": "some text", "model": "omni-moderation-latest"})
);
}
#[test]
fn deserializes_result() {
let body = json!({
"id": "modr-1",
"model": "omni-moderation-latest",
"results": [{
"flagged": true,
"categories": {"violence": true, "harassment/threatening": false},
"category_scores": {"violence": 0.98, "harassment/threatening": 0.01}
}]
});
let moderation: Moderation = serde_json::from_value(body).unwrap();
let result = &moderation.results[0];
assert!(result.flagged);
assert!(result.categories["violence"]);
assert!(result.category_scores["harassment/threatening"] < 0.5);
}
}