rig_core/providers/cohere/
extension.rs1use serde::de::DeserializeOwned;
17use serde::{Deserialize, Serialize, Serializer};
18use serde_json::Value;
19
20use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
21use crate::message::Api;
22
23#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
25pub struct CohereExt;
26
27impl ProviderExtension for CohereExt {
28 const PROVIDER: &'static str = super::PROVIDER_NAME;
29 type Options = CohereOptions;
30 type Extras = CohereExtras;
31}
32
33#[non_exhaustive]
35#[derive(Clone, Debug, Default, PartialEq, Serialize)]
36pub struct CohereOptions {
37 #[serde(rename = "*")]
39 pub shared: CohereShared,
40 #[serde(rename = "cohere.chat")]
42 pub chat: CohereNative,
43}
44
45impl ExtensionOptions for CohereOptions {
46 type Ext = CohereExt;
47}
48
49impl CohereOptions {
50 pub fn frequency_penalty(mut self, penalty: f64) -> Self {
53 self.shared.frequency_penalty = Some(penalty);
54 self
55 }
56
57 pub fn presence_penalty(mut self, penalty: f64) -> Self {
60 self.shared.presence_penalty = Some(penalty);
61 self
62 }
63
64 pub fn citation_mode(mut self, mode: CitationMode) -> Self {
66 self.chat.citation_mode = Some(mode);
67 self
68 }
69
70 pub fn safety_mode(mut self, mode: SafetyMode) -> Self {
72 self.chat.safety_mode = Some(mode);
73 self
74 }
75
76 pub fn priority(mut self, priority: u32) -> Self {
79 self.chat.priority = Some(priority);
80 self
81 }
82
83 pub fn top_k(mut self, top_k: u32) -> Self {
86 self.chat.top_k = Some(top_k);
87 self
88 }
89
90 pub fn logprobs(mut self, logprobs: bool) -> Self {
93 self.chat.logprobs = Some(logprobs);
94 self
95 }
96}
97
98#[non_exhaustive]
100#[derive(Clone, Debug, Default, PartialEq, Serialize)]
101pub struct CohereShared {
102 #[serde(skip_serializing_if = "Option::is_none")]
104 pub frequency_penalty: Option<f64>,
105 #[serde(skip_serializing_if = "Option::is_none")]
107 pub presence_penalty: Option<f64>,
108}
109
110#[non_exhaustive]
112#[derive(Clone, Debug, Default, PartialEq, Serialize)]
113pub struct CohereNative {
114 #[serde(
116 rename = "citation_options",
117 serialize_with = "citation_options",
118 skip_serializing_if = "Option::is_none"
119 )]
120 pub citation_mode: Option<CitationMode>,
121 #[serde(skip_serializing_if = "Option::is_none")]
123 pub safety_mode: Option<SafetyMode>,
124 #[serde(skip_serializing_if = "Option::is_none")]
126 pub priority: Option<u32>,
127 #[serde(rename = "k", skip_serializing_if = "Option::is_none")]
129 pub top_k: Option<u32>,
130 #[serde(skip_serializing_if = "Option::is_none")]
132 pub logprobs: Option<bool>,
133}
134
135fn citation_options<S: Serializer>(
137 mode: &Option<CitationMode>,
138 serializer: S,
139) -> Result<S::Ok, S::Error> {
140 #[derive(Serialize)]
141 struct CitationOptions<'a> {
142 mode: &'a Option<CitationMode>,
143 }
144 CitationOptions { mode }.serialize(serializer)
145}
146
147#[non_exhaustive]
149#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
150#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
151pub enum CitationMode {
152 Enabled,
154 Disabled,
156 Fast,
158 Accurate,
160 Off,
162}
163
164#[non_exhaustive]
166#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
167#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
168pub enum SafetyMode {
169 Contextual,
171 Strict,
173 Off,
175}
176
177#[non_exhaustive]
181#[derive(Clone, Debug, Default, PartialEq)]
182pub struct CohereExtras {
183 pub id: Option<String>,
185 pub finish_reason: Option<String>,
187 pub billed_units: Option<BilledUnits>,
189 pub tokens: Option<Tokens>,
191 pub cached_tokens: Option<f64>,
194 pub tool_plan: Option<String>,
196 pub logprobs: Option<Vec<Logprob>>,
198}
199
200#[non_exhaustive]
202#[derive(Clone, Copy, Debug, Default, PartialEq, Deserialize)]
203#[serde(default)]
204pub struct BilledUnits {
205 pub input_tokens: Option<f64>,
207 pub output_tokens: Option<f64>,
209 pub search_units: Option<f64>,
211 pub classifications: Option<f64>,
213}
214
215#[non_exhaustive]
217#[derive(Clone, Copy, Debug, Default, PartialEq, Deserialize)]
218#[serde(default)]
219pub struct Tokens {
220 pub input_tokens: Option<f64>,
222 pub output_tokens: Option<f64>,
224}
225
226#[non_exhaustive]
228#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
229#[serde(default)]
230pub struct Logprob {
231 pub token_ids: Vec<u32>,
233 pub text: Option<String>,
235 pub logprobs: Vec<f64>,
237}
238
239fn at<T: DeserializeOwned>(raw: &Value, pointer: &str) -> Result<Option<T>, serde_json::Error> {
241 match raw.pointer(pointer) {
242 None | Some(Value::Null) => Ok(None),
243 Some(value) => T::deserialize(value).map(Some),
244 }
245}
246
247impl ReplyExtras for CohereExtras {
248 fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
249 Ok(Self {
250 id: at(raw, "/id")?,
251 finish_reason: at(raw, "/finish_reason")?,
252 billed_units: at(raw, "/usage/billed_units")?,
253 tokens: at(raw, "/usage/tokens")?,
254 cached_tokens: at(raw, "/usage/cached_tokens")?,
255 tool_plan: at(raw, "/message/tool_plan")?,
256 logprobs: at(raw, "/logprobs")?,
257 })
258 }
259}
260
261#[cfg(test)]
262mod tests;