Skip to main content

rig_core/providers/cohere/
wire.rs

1//! Cohere's configuration, and its chat wire: the OpenAI Compatibility API
2//! by default, or the native chat API when the wire's route opts in.
3//!
4//! ```
5//! use rig_core::providers::cohere::{ChatRoute, CohereConfig};
6//! let config = CohereConfig::new("key").with_base_url("https://api.cohere.ai/");
7//! assert_eq!(config.base_url, "https://api.cohere.ai");
8//! let chat = config.completion("command-a-03-2025");
9//! assert_eq!(chat.route, ChatRoute::Compatibility);
10//! let chat = chat.with_route(ChatRoute::Auto);
11//! assert_eq!(chat.route, ChatRoute::Auto);
12//! ```
13
14use crate::client::env::{self, EnvError};
15use crate::completion::{CompletionRequest, ReplayTarget};
16use crate::error::{EncodeError, ProviderError};
17use crate::operation::Completion;
18use crate::providers::openai::wire::{COHERE, Chat, OpenAIConfig};
19use crate::wire::{
20    Decoder, Descriptor, Encoded, Flow, Mode, Out, Secret, Wire, WireEvent, WireFrame,
21};
22use serde::{Deserialize, Serialize};
23
24use super::NativeChat;
25
26/// Cohere's API root.
27const BASE_URL: &str = "https://api.cohere.ai";
28
29/// Where Cohere's OpenAI Compatibility API sits under the API root.
30const COMPATIBILITY_PATH: &str = "/compatibility/v1";
31
32/// The environment variable carrying the API key.
33const API_KEY_ENV: &str = "COHERE_API_KEY";
34
35/// The settings of a Cohere provider: serializable, and the credential is
36/// never serialized. [`connect`](Self::connect) puts it on a transport as a
37/// [`Cohere`](super::Cohere) client.
38#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
39pub struct CohereConfig {
40    /// The API key, sent as `Authorization: Bearer`.
41    pub api_key: Secret,
42    /// The API root, without a trailing slash.
43    pub base_url: String,
44}
45
46impl CohereConfig {
47    /// Cohere with default settings.
48    pub fn new(api_key: impl Into<Secret>) -> Self {
49        Self {
50            api_key: api_key.into(),
51            base_url: BASE_URL.to_owned(),
52        }
53    }
54
55    /// Cohere from `COHERE_API_KEY`.
56    pub fn from_env() -> Result<Self, EnvError> {
57        Ok(Self::new(env::required(API_KEY_ENV)?))
58    }
59
60    /// Point the wires at another API root.
61    pub fn with_base_url(mut self, base_url: impl AsRef<str>) -> Self {
62        self.base_url = base_url.as_ref().trim_end_matches('/').to_owned();
63        self
64    }
65
66    /// The chat wire for `model` under this API root, on the
67    /// [`ChatRoute::Compatibility`] route.
68    pub fn completion(&self, model: impl Into<String>) -> CohereChat {
69        let model = model.into();
70        CohereChat {
71            compatibility_api: OpenAIConfig::with_key(&COHERE, self.api_key.clone())
72                .with_base_url(format!("{}{COMPATIBILITY_PATH}", self.base_url))
73                .chat(model.clone()),
74            native_api: NativeChat::new(self.clone(), model),
75            route: ChatRoute::Compatibility,
76        }
77    }
78}
79
80/// Which Cohere API a chat request goes to. The default is
81/// [`Compatibility`](Self::Compatibility); the native API is opt-in.
82#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
83#[serde(rename_all = "snake_case")]
84pub enum ChatRoute {
85    /// The native API for a request that carries documents, so Cohere
86    /// grounds the answer in them and cites them, and the Compatibility
87    /// API for any other.
88    Auto,
89    /// The native API for every request.
90    Native,
91    /// The Compatibility API for every request, with documents sent as
92    /// text in the history.
93    #[default]
94    Compatibility,
95}
96
97/// Cohere chat: each request goes to the OpenAI Compatibility API or the
98/// native chat API, as [`ChatRoute`] says. Requests stay on the
99/// Compatibility API unless [`with_route`](Self::with_route) opts in to
100/// the native one.
101///
102/// The two are different APIs to history replay. A turn made on one route
103/// replays on the other from its canonical fields only, so its citations
104/// and tool plan stay behind. [`ChatRoute::Native`] keeps a whole
105/// conversation on the native API.
106#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
107pub struct CohereChat {
108    /// The Compatibility API's wire.
109    pub compatibility_api: Chat,
110    /// The native API's wire.
111    pub native_api: NativeChat,
112    /// Which API each request goes to.
113    pub route: ChatRoute,
114}
115
116impl CohereChat {
117    /// Send requests by `route`.
118    pub fn with_route(mut self, route: ChatRoute) -> Self {
119        self.route = route;
120        self
121    }
122
123    /// Ask both APIs to hold every tool call to its tool's schema.
124    pub fn with_strict_tools(mut self) -> Self {
125        self.compatibility_api = self.compatibility_api.with_strict_tools();
126        self.native_api = self.native_api.with_strict_tools();
127        self
128    }
129
130    /// Whether `request` goes to the native API.
131    fn native_for(&self, request: &CompletionRequest) -> bool {
132        match self.route {
133            ChatRoute::Auto => !request.documents.is_empty(),
134            ChatRoute::Native => true,
135            ChatRoute::Compatibility => false,
136        }
137    }
138}
139
140impl Wire for CohereChat {
141    type Op = Completion;
142    type Payload = Encoded;
143    type Frame = WireFrame;
144    type Decoder<'id> = CohereDecoder;
145    type Reassembler = document::Routed;
146
147    /// The Compatibility API's description, naming this wire as the replay
148    /// target that routes each request.
149    fn describe(&self) -> Descriptor<'_> {
150        self.compatibility_api.describe().replay(self)
151    }
152
153    fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
154        if self.native_for(&request) {
155            self.native_api.encode(request, mode)
156        } else {
157            self.compatibility_api.encode(request, mode)
158        }
159    }
160
161    fn decoder<'id>(&self) -> Self::Decoder<'id> {
162        CohereDecoder {
163            native_api: Wire::decoder(&self.native_api),
164            compatibility_api: Wire::decoder(&self.compatibility_api),
165            native_seen: false,
166        }
167    }
168
169    fn reassembler(&self) -> Self::Reassembler {
170        document::Routed::new(
171            Wire::reassembler(&self.native_api),
172            Wire::reassembler(&self.compatibility_api),
173        )
174    }
175}
176
177/// Every request names its route, so these are only what the replay of a
178/// history adapted to the wire as a whole reads: the Compatibility API's.
179impl ReplayTarget for CohereChat {
180    /// The answer of the API `request` goes to.
181    fn map_options(
182        &self,
183        request: &CompletionRequest,
184        fields: crate::completion::options::OptionFields<'_>,
185    ) -> crate::completion::options::OptionMap {
186        if self.native_for(request) {
187            self.native_api.map_options(request, fields)
188        } else {
189            self.compatibility_api.map_options(request, fields)
190        }
191    }
192
193    fn api(&self) -> crate::message::Api {
194        self.compatibility_api.api()
195    }
196
197    fn provider(&self) -> &str {
198        ReplayTarget::provider(&self.compatibility_api)
199    }
200
201    fn model(&self) -> &str {
202        ReplayTarget::model(&self.compatibility_api)
203    }
204
205    fn accepts(&self, model: &str) -> crate::completion::Accepts {
206        self.compatibility_api.accepts(model)
207    }
208
209    fn route(&self, request: &CompletionRequest) -> Option<&dyn ReplayTarget> {
210        Some(if self.native_for(request) {
211            &self.native_api
212        } else {
213            &self.compatibility_api
214        })
215    }
216}
217
218/// One frame of a Cohere chat reply, from whichever API answered.
219pub enum CohereEvent {
220    /// A native chat event or reply.
221    Native(super::streaming::ChatEvent),
222    /// A Compatibility API chunk, reply or signal.
223    Compatibility(crate::providers::openai::wire::chat::ChatEvent),
224    /// Cohere's error envelope delivered with a success status, which fails
225    /// the turn.
226    Failure(ProviderError),
227}
228
229/// Decodes a Cohere chat reply from either API. The decoder is not told
230/// which API a request went to, so each frame is read by its shape: a
231/// native event carries a `type`, a native reply a `message` object, and
232/// Cohere's error envelope a `message` string.
233pub struct CohereDecoder {
234    native_api: super::streaming::ChatDecoder,
235    compatibility_api: crate::providers::openai::wire::chat::ChatDecoder,
236    /// Whether a native frame arrived, so the end of the frames is read as
237    /// the native API's.
238    native_seen: bool,
239}
240
241/// How a frame's shape names the API that sent it.
242#[derive(Debug, Clone, Copy)]
243enum Shape {
244    Native,
245    Compatibility,
246    /// `{"id": ..., "message": "..."}`, Cohere's error body.
247    Error,
248}
249
250/// Which API sent `data`, read off its top-level keys.
251fn shape(data: &str) -> Shape {
252    let Ok(serde_json::Value::Object(frame)) = serde_json::from_str(data) else {
253        return Shape::Compatibility;
254    };
255    if frame.get("type").is_some_and(serde_json::Value::is_string) {
256        return Shape::Native;
257    }
258    if frame.contains_key("choices") {
259        return Shape::Compatibility;
260    }
261    match frame.get("message") {
262        Some(serde_json::Value::Object(_)) => Shape::Native,
263        Some(serde_json::Value::String(_)) => Shape::Error,
264        _ => Shape::Compatibility,
265    }
266}
267
268impl<'id> Decoder<'id, Completion> for CohereDecoder {
269    type Event = CohereEvent;
270
271    fn classify(&self, frame: WireFrame) -> WireEvent<CohereEvent> {
272        let data = frame.as_str();
273        match shape(&data) {
274            Shape::Error => WireEvent::Known(CohereEvent::Failure(
275                ProviderError::from_provider_body(data.into_owned()),
276            )),
277            Shape::Native => self.native_api.classify(frame).map(CohereEvent::Native),
278            Shape::Compatibility => self
279                .compatibility_api
280                .classify(frame)
281                .map(CohereEvent::Compatibility),
282        }
283    }
284
285    fn decode(
286        &mut self,
287        event: CohereEvent,
288        out: Out<'id, Completion>,
289    ) -> Result<Flow, ProviderError> {
290        match event {
291            CohereEvent::Native(event) => {
292                self.native_seen = true;
293                self.native_api.decode(event, out)
294            }
295            CohereEvent::Compatibility(event) => self.compatibility_api.decode(event, out),
296            CohereEvent::Failure(error) => Err(error),
297        }
298    }
299
300    fn eof(&mut self, out: Out<'id, Completion>) -> Result<Flow, ProviderError> {
301        if self.native_seen {
302            self.native_api.eof(out)
303        } else {
304            self.compatibility_api.eof(out)
305        }
306    }
307}
308
309mod document;
310
311#[cfg(test)]
312mod tests;