use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct ChatOptions {
#[serde(skip_serializing_if = "BTreeMap::is_empty")]
pub logit_bias: BTreeMap<String, i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prediction: Option<Prediction>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_logprobs: Option<u8>,
#[serde(skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f64>,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub modalities: Vec<Modality>,
#[serde(skip_serializing_if = "Option::is_none")]
pub audio: Option<AudioOutput>,
#[serde(skip_serializing_if = "Option::is_none")]
pub web_search_options: Option<WebSearchOptions>,
}
impl ChatOptions {
pub fn new() -> Self {
Self::default()
}
pub fn logit_bias(mut self, token: u32, bias: i32) -> Self {
self.logit_bias.insert(token.to_string(), bias);
self
}
pub fn prediction(mut self, content: impl Into<String>) -> Self {
self.prediction = Some(Prediction::content(content));
self
}
pub fn logprobs(mut self, logprobs: bool) -> Self {
self.logprobs = Some(logprobs);
self
}
pub fn top_logprobs(mut self, count: u8) -> Self {
self.top_logprobs = Some(count);
self
}
pub fn frequency_penalty(mut self, penalty: f64) -> Self {
self.frequency_penalty = Some(penalty);
self
}
pub fn presence_penalty(mut self, penalty: f64) -> Self {
self.presence_penalty = Some(penalty);
self
}
pub fn modalities(mut self, modalities: impl IntoIterator<Item = Modality>) -> Self {
self.modalities = modalities.into_iter().collect();
self
}
pub fn audio(mut self, voice: impl Into<String>, format: AudioFormat) -> Self {
self.audio = Some(AudioOutput {
voice: voice.into(),
format,
});
self
}
pub fn web_search_options(mut self, options: WebSearchOptions) -> Self {
self.web_search_options = Some(options);
self
}
}
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum Prediction {
Content {
content: String,
},
}
impl Prediction {
pub fn content(content: impl Into<String>) -> Self {
Self::Content {
content: content.into(),
}
}
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum Modality {
Text,
Audio,
}
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct AudioOutput {
pub voice: String,
pub format: AudioFormat,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum AudioFormat {
Wav,
Aac,
Mp3,
Flac,
Opus,
Pcm16,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct WebSearchOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub search_context_size: Option<SearchContextSize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_location: Option<UserLocation>,
}
impl WebSearchOptions {
pub fn new() -> Self {
Self::default()
}
pub fn search_context_size(mut self, size: SearchContextSize) -> Self {
self.search_context_size = Some(size);
self
}
pub fn user_location(mut self, location: ApproximateLocation) -> Self {
self.user_location = Some(UserLocation::Approximate {
approximate: location,
});
self
}
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum SearchContextSize {
Low,
Medium,
High,
}
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum UserLocation {
Approximate {
approximate: ApproximateLocation,
},
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct ApproximateLocation {
#[serde(skip_serializing_if = "Option::is_none")]
pub city: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub country: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub region: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timezone: Option<String>,
}
impl ApproximateLocation {
pub fn new() -> Self {
Self::default()
}
pub fn city(mut self, city: impl Into<String>) -> Self {
self.city = Some(city.into());
self
}
pub fn country(mut self, country: impl Into<String>) -> Self {
self.country = Some(country.into());
self
}
pub fn region(mut self, region: impl Into<String>) -> Self {
self.region = Some(region.into());
self
}
pub fn timezone(mut self, timezone: impl Into<String>) -> Self {
self.timezone = Some(timezone.into());
self
}
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
pub struct PromptTokensDetails {
#[serde(default)]
pub cached_tokens: Option<u64>,
#[serde(default)]
pub audio_tokens: Option<u64>,
#[serde(default)]
pub cache_write_tokens: Option<u64>,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
pub struct CompletionTokensDetails {
#[serde(default)]
pub reasoning_tokens: Option<u64>,
#[serde(default)]
pub audio_tokens: Option<u64>,
#[serde(default)]
pub accepted_prediction_tokens: Option<u64>,
#[serde(default)]
pub rejected_prediction_tokens: Option<u64>,
}