use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize, Serializer};
use serde_json::Value;
use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
use crate::message::Api;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct CohereExt;
impl ProviderExtension for CohereExt {
const PROVIDER: &'static str = super::PROVIDER_NAME;
type Options = CohereOptions;
type Extras = CohereExtras;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct CohereOptions {
#[serde(rename = "*")]
pub shared: CohereShared,
#[serde(rename = "cohere.chat")]
pub chat: CohereNative,
}
impl ExtensionOptions for CohereOptions {
type Ext = CohereExt;
}
impl CohereOptions {
pub fn frequency_penalty(mut self, penalty: f64) -> Self {
self.shared.frequency_penalty = Some(penalty);
self
}
pub fn presence_penalty(mut self, penalty: f64) -> Self {
self.shared.presence_penalty = Some(penalty);
self
}
pub fn citation_mode(mut self, mode: CitationMode) -> Self {
self.chat.citation_mode = Some(mode);
self
}
pub fn safety_mode(mut self, mode: SafetyMode) -> Self {
self.chat.safety_mode = Some(mode);
self
}
pub fn priority(mut self, priority: u32) -> Self {
self.chat.priority = Some(priority);
self
}
pub fn top_k(mut self, top_k: u32) -> Self {
self.chat.top_k = Some(top_k);
self
}
pub fn logprobs(mut self, logprobs: bool) -> Self {
self.chat.logprobs = Some(logprobs);
self
}
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct CohereShared {
#[serde(skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f64>,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct CohereNative {
#[serde(
rename = "citation_options",
serialize_with = "citation_options",
skip_serializing_if = "Option::is_none"
)]
pub citation_mode: Option<CitationMode>,
#[serde(skip_serializing_if = "Option::is_none")]
pub safety_mode: Option<SafetyMode>,
#[serde(skip_serializing_if = "Option::is_none")]
pub priority: Option<u32>,
#[serde(rename = "k", skip_serializing_if = "Option::is_none")]
pub top_k: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<bool>,
}
fn citation_options<S: Serializer>(
mode: &Option<CitationMode>,
serializer: S,
) -> Result<S::Ok, S::Error> {
#[derive(Serialize)]
struct CitationOptions<'a> {
mode: &'a Option<CitationMode>,
}
CitationOptions { mode }.serialize(serializer)
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
pub enum CitationMode {
Enabled,
Disabled,
Fast,
Accurate,
Off,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
pub enum SafetyMode {
Contextual,
Strict,
Off,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq)]
pub struct CohereExtras {
pub id: Option<String>,
pub finish_reason: Option<String>,
pub billed_units: Option<BilledUnits>,
pub tokens: Option<Tokens>,
pub cached_tokens: Option<f64>,
pub tool_plan: Option<String>,
pub logprobs: Option<Vec<Logprob>>,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, Default, PartialEq, Deserialize)]
#[serde(default)]
pub struct BilledUnits {
pub input_tokens: Option<f64>,
pub output_tokens: Option<f64>,
pub search_units: Option<f64>,
pub classifications: Option<f64>,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, Default, PartialEq, Deserialize)]
#[serde(default)]
pub struct Tokens {
pub input_tokens: Option<f64>,
pub output_tokens: Option<f64>,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
#[serde(default)]
pub struct Logprob {
pub token_ids: Vec<u32>,
pub text: Option<String>,
pub logprobs: Vec<f64>,
}
fn at<T: DeserializeOwned>(raw: &Value, pointer: &str) -> Result<Option<T>, serde_json::Error> {
match raw.pointer(pointer) {
None | Some(Value::Null) => Ok(None),
Some(value) => T::deserialize(value).map(Some),
}
}
impl ReplyExtras for CohereExtras {
fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
Ok(Self {
id: at(raw, "/id")?,
finish_reason: at(raw, "/finish_reason")?,
billed_units: at(raw, "/usage/billed_units")?,
tokens: at(raw, "/usage/tokens")?,
cached_tokens: at(raw, "/usage/cached_tokens")?,
tool_plan: at(raw, "/message/tool_plan")?,
logprobs: at(raw, "/logprobs")?,
})
}
}
#[cfg(test)]
mod tests;