use crate::client::env::{self, EnvError};
use crate::completion::{CompletionRequest, ReplayTarget};
use crate::error::{EncodeError, ProviderError};
use crate::operation::Completion;
use crate::providers::openai::wire::{COHERE, Chat, OpenAIConfig};
use crate::wire::{
Decoder, Descriptor, Encoded, Flow, Mode, Out, Secret, Wire, WireEvent, WireFrame,
};
use serde::{Deserialize, Serialize};
use super::NativeChat;
const BASE_URL: &str = "https://api.cohere.ai";
const COMPATIBILITY_PATH: &str = "/compatibility/v1";
const API_KEY_ENV: &str = "COHERE_API_KEY";
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CohereConfig {
pub api_key: Secret,
pub base_url: String,
}
impl CohereConfig {
pub fn new(api_key: impl Into<Secret>) -> Self {
Self {
api_key: api_key.into(),
base_url: BASE_URL.to_owned(),
}
}
pub fn from_env() -> Result<Self, EnvError> {
Ok(Self::new(env::required(API_KEY_ENV)?))
}
pub fn with_base_url(mut self, base_url: impl AsRef<str>) -> Self {
self.base_url = base_url.as_ref().trim_end_matches('/').to_owned();
self
}
pub fn completion(&self, model: impl Into<String>) -> CohereChat {
let model = model.into();
CohereChat {
compatibility_api: OpenAIConfig::with_key(&COHERE, self.api_key.clone())
.with_base_url(format!("{}{COMPATIBILITY_PATH}", self.base_url))
.chat(model.clone()),
native_api: NativeChat::new(self.clone(), model),
route: ChatRoute::Compatibility,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ChatRoute {
Auto,
Native,
#[default]
Compatibility,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CohereChat {
pub compatibility_api: Chat,
pub native_api: NativeChat,
pub route: ChatRoute,
}
impl CohereChat {
pub fn with_route(mut self, route: ChatRoute) -> Self {
self.route = route;
self
}
pub fn with_strict_tools(mut self) -> Self {
self.compatibility_api = self.compatibility_api.with_strict_tools();
self.native_api = self.native_api.with_strict_tools();
self
}
fn native_for(&self, request: &CompletionRequest) -> bool {
match self.route {
ChatRoute::Auto => !request.documents.is_empty(),
ChatRoute::Native => true,
ChatRoute::Compatibility => false,
}
}
}
impl Wire for CohereChat {
type Op = Completion;
type Payload = Encoded;
type Frame = WireFrame;
type Decoder<'id> = CohereDecoder;
type Reassembler = document::Routed;
fn describe(&self) -> Descriptor<'_> {
self.compatibility_api.describe().replay(self)
}
fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
if self.native_for(&request) {
self.native_api.encode(request, mode)
} else {
self.compatibility_api.encode(request, mode)
}
}
fn decoder<'id>(&self) -> Self::Decoder<'id> {
CohereDecoder {
native_api: Wire::decoder(&self.native_api),
compatibility_api: Wire::decoder(&self.compatibility_api),
native_seen: false,
}
}
fn reassembler(&self) -> Self::Reassembler {
document::Routed::new(
Wire::reassembler(&self.native_api),
Wire::reassembler(&self.compatibility_api),
)
}
}
impl ReplayTarget for CohereChat {
fn map_options(
&self,
request: &CompletionRequest,
fields: crate::completion::options::OptionFields<'_>,
) -> crate::completion::options::OptionMap {
if self.native_for(request) {
self.native_api.map_options(request, fields)
} else {
self.compatibility_api.map_options(request, fields)
}
}
fn api(&self) -> crate::message::Api {
self.compatibility_api.api()
}
fn provider(&self) -> &str {
ReplayTarget::provider(&self.compatibility_api)
}
fn model(&self) -> &str {
ReplayTarget::model(&self.compatibility_api)
}
fn accepts(&self, model: &str) -> crate::completion::Accepts {
self.compatibility_api.accepts(model)
}
fn route(&self, request: &CompletionRequest) -> Option<&dyn ReplayTarget> {
Some(if self.native_for(request) {
&self.native_api
} else {
&self.compatibility_api
})
}
}
pub enum CohereEvent {
Native(super::streaming::ChatEvent),
Compatibility(crate::providers::openai::wire::chat::ChatEvent),
Failure(ProviderError),
}
pub struct CohereDecoder {
native_api: super::streaming::ChatDecoder,
compatibility_api: crate::providers::openai::wire::chat::ChatDecoder,
native_seen: bool,
}
#[derive(Debug, Clone, Copy)]
enum Shape {
Native,
Compatibility,
Error,
}
fn shape(data: &str) -> Shape {
let Ok(serde_json::Value::Object(frame)) = serde_json::from_str(data) else {
return Shape::Compatibility;
};
if frame.get("type").is_some_and(serde_json::Value::is_string) {
return Shape::Native;
}
if frame.contains_key("choices") {
return Shape::Compatibility;
}
match frame.get("message") {
Some(serde_json::Value::Object(_)) => Shape::Native,
Some(serde_json::Value::String(_)) => Shape::Error,
_ => Shape::Compatibility,
}
}
impl<'id> Decoder<'id, Completion> for CohereDecoder {
type Event = CohereEvent;
fn classify(&self, frame: WireFrame) -> WireEvent<CohereEvent> {
let data = frame.as_str();
match shape(&data) {
Shape::Error => WireEvent::Known(CohereEvent::Failure(
ProviderError::from_provider_body(data.into_owned()),
)),
Shape::Native => self.native_api.classify(frame).map(CohereEvent::Native),
Shape::Compatibility => self
.compatibility_api
.classify(frame)
.map(CohereEvent::Compatibility),
}
}
fn decode(
&mut self,
event: CohereEvent,
out: Out<'id, Completion>,
) -> Result<Flow, ProviderError> {
match event {
CohereEvent::Native(event) => {
self.native_seen = true;
self.native_api.decode(event, out)
}
CohereEvent::Compatibility(event) => self.compatibility_api.decode(event, out),
CohereEvent::Failure(error) => Err(error),
}
}
fn eof(&mut self, out: Out<'id, Completion>) -> Result<Flow, ProviderError> {
if self.native_seen {
self.native_api.eof(out)
} else {
self.compatibility_api.eof(out)
}
}
}
mod document;
#[cfg(test)]
mod tests;