#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Provider {
Cohere,
OpenAiCompatible,
}
impl Provider {
pub fn as_str(self) -> &'static str {
match self {
Provider::Cohere => "cohere",
Provider::OpenAiCompatible => "openai-compatible",
}
}
pub const ALL: &'static [Provider] = &[Provider::Cohere, Provider::OpenAiCompatible];
pub const ACCEPTED_NAMES: &'static [&'static str] = &["cohere", "openai"];
}
impl std::fmt::Display for Provider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UnknownProvider(pub String);
impl std::fmt::Display for UnknownProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "unknown provider {:?} (expected one of: ", self.0)?;
for (i, name) in Provider::ACCEPTED_NAMES.iter().enumerate() {
if i > 0 {
f.write_str(", ")?;
}
write!(f, "{name}")?;
}
f.write_str(")")
}
}
impl std::error::Error for UnknownProvider {}
impl std::str::FromStr for Provider {
type Err = UnknownProvider;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"cohere" => Ok(Provider::Cohere),
"openai" | "openai-compatible" => Ok(Provider::OpenAiCompatible),
other => Err(UnknownProvider(other.to_string())),
}
}
}
#[cfg(test)]
mod tests {
use super::Provider;
#[test]
fn every_tag_round_trips_through_from_str() {
for &p in Provider::ALL {
assert_eq!(
p.as_str().parse::<Provider>(),
Ok(p),
"{p} did not round-trip"
);
}
}
#[test]
fn the_user_facing_openai_spelling_is_accepted() {
assert_eq!("openai".parse::<Provider>(), Ok(Provider::OpenAiCompatible));
}
#[test]
fn every_advertised_name_parses() {
for name in Provider::ACCEPTED_NAMES {
assert!(
name.parse::<Provider>().is_ok(),
"advertised {name:?} but it does not parse"
);
}
}
#[test]
fn every_variant_is_reachable_by_an_advertised_name() {
for &p in Provider::ALL {
let reachable = Provider::ACCEPTED_NAMES
.iter()
.any(|n| n.parse::<Provider>() == Ok(p));
assert!(
reachable,
"{p} has no entry in ACCEPTED_NAMES, so no documented spelling reaches it"
);
}
}
#[test]
fn an_unknown_name_lists_the_names_users_actually_type() {
let err = "nope".parse::<Provider>().unwrap_err();
let msg = err.to_string();
assert!(msg.contains("nope"), "{msg}");
assert!(msg.contains("cohere"), "{msg}");
assert!(msg.contains("openai"), "{msg}");
assert!(!msg.contains("openai-compatible"), "{msg}");
}
}