Skip to main content

rig_core/providers/cohere/
extension.rs

1//! Cohere's typed request options and reply extras. The shared section
2//! goes to both routes, the Compatibility API and the native chat API; the
3//! `"cohere.chat"` section goes to the native API only and is skipped on a
4//! request the Compatibility API takes.
5//!
6//! ```
7//! use rig_core::completion::CompletionRequest;
8//! use rig_core::providers::cohere::extension::{CitationMode, CohereOptions};
9//!
10//! let options = CohereOptions::default()
11//!     .frequency_penalty(0.2)
12//!     .citation_mode(CitationMode::Fast);
13//! let request = CompletionRequest::new("hi").provider_option(options);
14//! ```
15
16use 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/// The `cohere` provider: its key, [`CohereOptions`] and [`CohereExtras`].
24#[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/// Cohere's request options.
34#[non_exhaustive]
35#[derive(Clone, Debug, Default, PartialEq, Serialize)]
36pub struct CohereOptions {
37    /// The fields both routes read.
38    #[serde(rename = "*")]
39    pub shared: CohereShared,
40    /// The fields only the native chat API reads.
41    #[serde(rename = "cohere.chat")]
42    pub chat: CohereNative,
43}
44
45impl ExtensionOptions for CohereOptions {
46    type Ext = CohereExt;
47}
48
49impl CohereOptions {
50    /// Penalize tokens by how often they appeared (`frequency_penalty`,
51    /// 0.0 to 1.0).
52    pub fn frequency_penalty(mut self, penalty: f64) -> Self {
53        self.shared.frequency_penalty = Some(penalty);
54        self
55    }
56
57    /// Penalize tokens that appeared at all (`presence_penalty`, 0.0 to
58    /// 1.0).
59    pub fn presence_penalty(mut self, penalty: f64) -> Self {
60        self.shared.presence_penalty = Some(penalty);
61        self
62    }
63
64    /// How the native API cites the documents (`citation_options.mode`).
65    pub fn citation_mode(mut self, mode: CitationMode) -> Self {
66        self.chat.citation_mode = Some(mode);
67        self
68    }
69
70    /// The safety instruction the native API adds (`safety_mode`).
71    pub fn safety_mode(mut self, mode: SafetyMode) -> Self {
72        self.chat.safety_mode = Some(mode);
73        self
74    }
75
76    /// The request's queue priority on the native API (`priority`): lower
77    /// is served first, and 0 is the default.
78    pub fn priority(mut self, priority: u32) -> Self {
79        self.chat.priority = Some(priority);
80        self
81    }
82
83    /// Sample from the `top_k` most likely tokens on the native API (`k`,
84    /// 0 to 500).
85    pub fn top_k(mut self, top_k: u32) -> Self {
86        self.chat.top_k = Some(top_k);
87        self
88    }
89
90    /// Return each generated token's log probability on the native API
91    /// (`logprobs`), read back as [`CohereExtras::logprobs`].
92    pub fn logprobs(mut self, logprobs: bool) -> Self {
93        self.chat.logprobs = Some(logprobs);
94        self
95    }
96}
97
98/// The fields both Cohere routes read.
99#[non_exhaustive]
100#[derive(Clone, Debug, Default, PartialEq, Serialize)]
101pub struct CohereShared {
102    /// `frequency_penalty`.
103    #[serde(skip_serializing_if = "Option::is_none")]
104    pub frequency_penalty: Option<f64>,
105    /// `presence_penalty`.
106    #[serde(skip_serializing_if = "Option::is_none")]
107    pub presence_penalty: Option<f64>,
108}
109
110/// The fields only the native chat API reads.
111#[non_exhaustive]
112#[derive(Clone, Debug, Default, PartialEq, Serialize)]
113pub struct CohereNative {
114    /// `citation_options.mode`.
115    #[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    /// `safety_mode`.
122    #[serde(skip_serializing_if = "Option::is_none")]
123    pub safety_mode: Option<SafetyMode>,
124    /// `priority`.
125    #[serde(skip_serializing_if = "Option::is_none")]
126    pub priority: Option<u32>,
127    /// `k`.
128    #[serde(rename = "k", skip_serializing_if = "Option::is_none")]
129    pub top_k: Option<u32>,
130    /// `logprobs`.
131    #[serde(skip_serializing_if = "Option::is_none")]
132    pub logprobs: Option<bool>,
133}
134
135/// `{"mode": <mode>}`, the `citation_options` object.
136fn 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/// How the native API cites the documents a reply uses.
148#[non_exhaustive]
149#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
150#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
151pub enum CitationMode {
152    /// `ENABLED`.
153    Enabled,
154    /// `DISABLED`.
155    Disabled,
156    /// `FAST`.
157    Fast,
158    /// `ACCURATE`.
159    Accurate,
160    /// `OFF`.
161    Off,
162}
163
164/// The safety instruction the native API adds to the prompt.
165#[non_exhaustive]
166#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
167#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
168pub enum SafetyMode {
169    /// `CONTEXTUAL`, the default.
170    Contextual,
171    /// `STRICT`.
172    Strict,
173    /// `OFF`.
174    Off,
175}
176
177/// The fields of a Cohere reply Rig does not normalize. Every field but
178/// [`Self::id`] comes from a native chat API reply; a Compatibility API
179/// reply leaves it `None`.
180#[non_exhaustive]
181#[derive(Clone, Debug, Default, PartialEq)]
182pub struct CohereExtras {
183    /// `id`, on both routes.
184    pub id: Option<String>,
185    /// `finish_reason` as Cohere spells it (`COMPLETE`, `MAX_TOKENS`, ...).
186    pub finish_reason: Option<String>,
187    /// `usage.billed_units`.
188    pub billed_units: Option<BilledUnits>,
189    /// `usage.tokens`.
190    pub tokens: Option<Tokens>,
191    /// `usage.cached_tokens`: the part of `tokens.input_tokens` read from
192    /// the cache, which `billed_units` leaves out.
193    pub cached_tokens: Option<f64>,
194    /// `message.tool_plan`: the model's plan before its tool calls.
195    pub tool_plan: Option<String>,
196    /// `logprobs`, when [`CohereOptions::logprobs`] asked for them.
197    pub logprobs: Option<Vec<Logprob>>,
198}
199
200/// The units Cohere bills a request by.
201#[non_exhaustive]
202#[derive(Clone, Copy, Debug, Default, PartialEq, Deserialize)]
203#[serde(default)]
204pub struct BilledUnits {
205    /// `input_tokens`.
206    pub input_tokens: Option<f64>,
207    /// `output_tokens`.
208    pub output_tokens: Option<f64>,
209    /// `search_units`.
210    pub search_units: Option<f64>,
211    /// `classifications`.
212    pub classifications: Option<f64>,
213}
214
215/// The tokens a request read and wrote, system overhead included.
216#[non_exhaustive]
217#[derive(Clone, Copy, Debug, Default, PartialEq, Deserialize)]
218#[serde(default)]
219pub struct Tokens {
220    /// `input_tokens`.
221    pub input_tokens: Option<f64>,
222    /// `output_tokens`.
223    pub output_tokens: Option<f64>,
224}
225
226/// The log probabilities of one stretch of generated text.
227#[non_exhaustive]
228#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
229#[serde(default)]
230pub struct Logprob {
231    /// `token_ids`.
232    pub token_ids: Vec<u32>,
233    /// `text`, the tokens decoded.
234    pub text: Option<String>,
235    /// `logprobs`, one per token.
236    pub logprobs: Vec<f64>,
237}
238
239/// The value at `pointer` in `raw`, or `None` when it is absent or `null`.
240fn 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;