1use std::fmt;
15use std::hash::{Hash, Hasher};
16
17use serde::de::{self, MapAccess, Visitor};
18use serde::ser::SerializeStruct;
19use serde::{Deserialize, Deserializer, Serialize, Serializer};
20
21#[cfg(feature = "reqwest")]
22use crate::client::env::{self, EnvError};
23use crate::driver::DynModel;
24use crate::http_client::{DynHttpClient, HttpClientExt};
25use crate::operation::Completion;
26use crate::providers::{anthropic, gemini, openai};
27use crate::serve::ErasedHandler;
28use crate::serve::adapters::ModelAdapter;
29use crate::wire::Secret;
30
31pub(crate) const OPENAI_DIALECTS: &[&openai::wire::Dialect] = &[
38 &openai::wire::OPENAI,
39 &openai::wire::AZURE,
40 &openai::wire::DEEPSEEK,
41 &openai::wire::GROQ,
42 &openai::wire::HYPERBOLIC,
43 &openai::wire::MIRA,
44 &openai::wire::PERPLEXITY,
45 &openai::wire::TOGETHER,
46 &openai::wire::HUGGINGFACE,
47 &openai::wire::LLAMACPP,
48 &openai::wire::MISTRAL,
49 &openai::wire::OPENROUTER,
50 &openai::wire::VENICE,
51 &openai::wire::DOUBLEWORD,
52 &openai::wire::ZAI,
53 &openai::wire::MINIMAX,
54 &openai::wire::MOONSHOT,
55 &openai::wire::XIAOMIMIMO,
56 &crate::providers::xai::DIALECT,
57 &crate::providers::chatgpt::DIALECT,
58 &crate::providers::copilot::wire::DIALECT,
59];
60
61#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
65pub enum Format {
66 OpenAi,
68 Anthropic,
70 Gemini,
72}
73
74impl Format {
75 pub const ALL: [Format; 3] = [Format::OpenAi, Format::Anthropic, Format::Gemini];
77
78 pub fn as_str(self) -> &'static str {
80 match self {
81 Format::OpenAi => "openai",
82 Format::Anthropic => "anthropic",
83 Format::Gemini => "gemini",
84 }
85 }
86
87 pub fn named(name: &str) -> Option<Self> {
89 Self::ALL.into_iter().find(|format| format.as_str() == name)
90 }
91
92 fn names() -> String {
94 Self::ALL
95 .iter()
96 .map(|format| format.as_str())
97 .collect::<Vec<_>>()
98 .join(", ")
99 }
100}
101
102impl fmt::Display for Format {
103 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
104 f.write_str(self.as_str())
105 }
106}
107
108impl Serialize for Format {
109 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
110 serializer.serialize_str(self.as_str())
111 }
112}
113
114impl<'de> Deserialize<'de> for Format {
115 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
116 let name = String::deserialize(deserializer)?;
117 Self::named(&name).ok_or_else(|| {
118 de::Error::custom(format!(
119 "`{name}` is not a protocol family ({})",
120 Format::names()
121 ))
122 })
123 }
124}
125
126#[derive(Debug, Clone, Copy)]
128enum Registered {
129 OpenAi(openai::wire::Dialect),
130 Anthropic(anthropic::wire::Dialect),
131 Gemini,
132}
133
134#[derive(Debug, Clone, Copy)]
138pub struct ProviderId(Registered);
139
140impl PartialEq for ProviderId {
143 fn eq(&self, other: &Self) -> bool {
144 self.vendor() == other.vendor() && self.format() == other.format()
145 }
146}
147
148impl Eq for ProviderId {}
149
150impl Hash for ProviderId {
151 fn hash<H: Hasher>(&self, state: &mut H) {
152 self.vendor().hash(state);
153 self.format().hash(state);
154 }
155}
156
157impl ProviderId {
158 pub fn all() -> impl Iterator<Item = ProviderId> {
161 openai::wire::all()
162 .map(|dialect| ProviderId(Registered::OpenAi(*dialect)))
163 .chain(
164 anthropic::wire::all().map(|dialect| ProviderId(Registered::Anthropic(*dialect))),
165 )
166 .chain(std::iter::once(ProviderId(Registered::Gemini)))
167 }
168
169 pub fn new(vendor: &str, format: Format) -> Option<Self> {
172 match format {
173 Format::OpenAi => {
174 openai::wire::by_name(vendor).map(|dialect| Self(Registered::OpenAi(*dialect)))
175 }
176 Format::Anthropic => anthropic::wire::Dialect::by_name(vendor)
177 .map(|dialect| Self(Registered::Anthropic(dialect))),
178 Format::Gemini => (vendor == gemini::PROVIDER_NAME).then_some(Self(Registered::Gemini)),
179 }
180 }
181
182 pub fn vendor(&self) -> &'static str {
184 match &self.0 {
185 Registered::OpenAi(dialect) => dialect.name,
186 Registered::Anthropic(dialect) => dialect.name,
187 Registered::Gemini => gemini::PROVIDER_NAME,
188 }
189 }
190
191 pub fn format(&self) -> Format {
193 match &self.0 {
194 Registered::OpenAi(_) => Format::OpenAi,
195 Registered::Anthropic(_) => Format::Anthropic,
196 Registered::Gemini => Format::Gemini,
197 }
198 }
199
200 pub fn vendor_selections(vendor: &str) -> impl Iterator<Item = ProviderId> + '_ {
202 Self::all().filter(move |id| id.vendor() == vendor)
203 }
204
205 pub fn resolve(selection: &str) -> Result<Self, SelectionError> {
211 let malformed = || SelectionError::Malformed {
212 selection: selection.to_owned(),
213 };
214 let (vendor, format) = match selection.split_once('/') {
215 Some((_, rest)) if rest.contains('/') => return Err(malformed()),
216 Some((vendor, format)) => (vendor, Some(format)),
217 None => (selection, None),
218 };
219 if vendor.is_empty() {
220 return Err(malformed());
221 }
222 let Some(format) = format else {
223 let mut registered = Self::vendor_selections(vendor);
224 let first = registered.next().ok_or_else(|| SelectionError::Unknown {
225 vendor: vendor.to_owned(),
226 })?;
227 return match registered.next() {
228 None => Ok(first),
229 Some(_) => Err(SelectionError::Ambiguous {
230 vendor: vendor.to_owned(),
231 alternatives: alternatives(vendor),
232 }),
233 };
234 };
235 let Some(family) = Format::named(format) else {
236 return Err(SelectionError::UnknownFormat {
237 format: format.to_owned(),
238 });
239 };
240 Self::new(vendor, family).ok_or_else(|| {
241 let alternatives = alternatives(vendor);
242 match alternatives.is_empty() {
243 true => SelectionError::Unknown {
244 vendor: vendor.to_owned(),
245 },
246 false => SelectionError::Unregistered {
247 vendor: vendor.to_owned(),
248 format: family,
249 alternatives,
250 },
251 }
252 })
253 }
254
255 pub fn config(&self, api_key: impl Into<Secret>) -> ProviderConfig {
258 match &self.0 {
259 Registered::OpenAi(dialect) => {
260 ProviderConfig::OpenAi(openai::wire::OpenAIConfig::with_key(dialect, api_key))
261 }
262 Registered::Anthropic(dialect) => ProviderConfig::Anthropic(
263 anthropic::wire::AnthropicConfig::with_dialect(api_key, dialect),
264 ),
265 Registered::Gemini => ProviderConfig::Gemini(gemini::GeminiConfig::new(api_key)),
266 }
267 }
268
269 #[cfg(feature = "reqwest")]
272 fn config_from_env(&self) -> Result<ProviderConfig, EnvError> {
273 Ok(match &self.0 {
274 Registered::OpenAi(dialect) => {
275 let (api_key, auth) = openai_credential_from_env(dialect)?;
276 ProviderConfig::OpenAi(openai::wire::OpenAIConfig::from_env_with_credential(
277 dialect, api_key, auth,
278 )?)
279 }
280 Registered::Anthropic(dialect) => {
281 ProviderConfig::Anthropic(anthropic::wire::AnthropicConfig::from_env_with(dialect)?)
282 }
283 Registered::Gemini => ProviderConfig::Gemini(gemini::GeminiConfig::from_env()?),
284 })
285 }
286
287 pub fn api_key_env(&self) -> &'static str {
289 match &self.0 {
290 Registered::OpenAi(dialect) => dialect.api_key_env,
291 Registered::Anthropic(dialect) => dialect.api_key_env,
292 Registered::Gemini => gemini::API_KEY_ENV,
293 }
294 }
295
296 pub fn requires_credential(&self) -> bool {
301 match &self.0 {
302 Registered::OpenAi(dialect) => {
303 !matches!(dialect.quirks.auth, openai::wire::Auth::OptionalBearer)
304 }
305 Registered::Anthropic(_) | Registered::Gemini => true,
306 }
307 }
308}
309
310fn alternatives(vendor: &str) -> Vec<String> {
312 ProviderId::vendor_selections(vendor)
313 .map(|id| id.to_string())
314 .collect()
315}
316
317impl fmt::Display for ProviderId {
318 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
319 write!(f, "{}/{}", self.vendor(), self.format())
320 }
321}
322
323impl Serialize for ProviderId {
324 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
325 serializer.collect_str(self)
326 }
327}
328
329impl<'de> Deserialize<'de> for ProviderId {
330 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
331 let selection = String::deserialize(deserializer)?;
332 Self::resolve(&selection).map_err(de::Error::custom)
333 }
334}
335
336#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
338pub enum SelectionError {
339 #[error("no registered provider is named `{vendor}`")]
341 Unknown {
342 vendor: String,
344 },
345 #[error(
347 "`{vendor}` speaks no {format} endpoint in this build (it speaks {})",
348 alternatives.join(", ")
349 )]
350 Unregistered {
351 vendor: String,
353 format: Format,
355 alternatives: Vec<String>,
357 },
358 #[error(
360 "`{vendor}` names more than one registered selection: name one of {}",
361 alternatives.join(", ")
362 )]
363 Ambiguous {
364 vendor: String,
366 alternatives: Vec<String>,
368 },
369 #[error("`{format}` is not a protocol family ({})", Format::names())]
371 UnknownFormat {
372 format: String,
374 },
375 #[error("`{selection}` is not a provider selection: expected `vendor` or `vendor/format`")]
377 Malformed {
378 selection: String,
380 },
381}
382
383#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
388pub enum ProviderConfig {
389 #[serde(rename = "openai")]
391 OpenAi(openai::wire::OpenAIConfig),
392 #[serde(rename = "anthropic")]
394 Anthropic(anthropic::wire::AnthropicConfig),
395 #[serde(rename = "gemini")]
397 Gemini(gemini::GeminiConfig),
398}
399
400impl ProviderConfig {
401 pub fn id(&self) -> Option<ProviderId> {
405 ProviderId::new(self.vendor(), self.format())
406 }
407
408 pub fn vendor(&self) -> &'static str {
410 match self {
411 Self::OpenAi(provider) => provider.dialect.name,
412 Self::Anthropic(provider) => provider.dialect.name,
413 Self::Gemini(_) => gemini::PROVIDER_NAME,
414 }
415 }
416
417 pub fn format(&self) -> Format {
419 match self {
420 Self::OpenAi(_) => Format::OpenAi,
421 Self::Anthropic(_) => Format::Anthropic,
422 Self::Gemini(_) => Format::Gemini,
423 }
424 }
425
426 pub fn is_unauthenticated(&self) -> bool {
429 match self {
430 Self::OpenAi(provider) => provider.api_key.is_empty(),
431 Self::Anthropic(provider) => provider.api_key.is_empty(),
432 Self::Gemini(provider) => provider.api_key.is_empty(),
433 }
434 }
435
436 pub fn with_credential(mut self, api_key: impl Into<Secret>) -> Self {
439 let api_key = api_key.into();
440 match &mut self {
441 Self::OpenAi(provider) => provider.api_key = api_key,
442 Self::Anthropic(provider) => provider.api_key = api_key,
443 Self::Gemini(provider) => provider.api_key = api_key,
444 }
445 self
446 }
447
448 pub fn completion_handler(
455 &self,
456 label: &str,
457 model: &str,
458 http: DynHttpClient,
459 ) -> ErasedHandler {
460 ErasedHandler::new(ModelAdapter::new(label, self.completion_model(model, http)))
461 }
462
463 fn completion_model(&self, model: &str, http: DynHttpClient) -> DynModel<Completion> {
465 match self {
466 Self::OpenAi(provider) => provider.clone().connect(http).completion(model).erase(),
467 Self::Anthropic(provider) => provider.clone().connect(http).completion(model).erase(),
468 Self::Gemini(provider) => provider.clone().connect(http).completion(model).erase(),
469 }
470 }
471
472 #[cfg(feature = "reqwest")]
476 fn with_credential_from_env(self) -> Result<Self, EnvError> {
477 Ok(match self {
478 Self::OpenAi(mut provider) => {
479 let (api_key, auth) = openai_credential_from_env(&provider.dialect)?;
480 provider.api_key = api_key.into();
481 if auth != provider.dialect.quirks.auth {
483 provider.auth = auth;
484 }
485 Self::OpenAi(provider)
486 }
487 Self::Anthropic(mut provider) => {
488 provider.api_key = env::required(provider.dialect.api_key_env)?.into();
489 Self::Anthropic(provider)
490 }
491 Self::Gemini(mut provider) => {
492 provider.api_key = env::required(gemini::API_KEY_ENV)?.into();
493 Self::Gemini(provider)
494 }
495 })
496 }
497}
498
499#[cfg(feature = "reqwest")]
503fn openai_credential_from_env(
504 dialect: &openai::wire::Dialect,
505) -> Result<(String, openai::wire::Auth), EnvError> {
506 if matches!(dialect.quirks.auth, openai::wire::Auth::OptionalBearer)
507 && dialect.alternate_auth.is_none()
508 {
509 let api_key = env::optional(dialect.api_key_env)?.unwrap_or_default();
510 return Ok((api_key, dialect.quirks.auth));
511 }
512 openai::wire::OpenAIConfig::credential_from_env(dialect)
513}
514
515#[derive(Clone, Debug, PartialEq)]
518pub enum Provider {
519 Registered(ProviderId),
521 Configured(ProviderConfig),
523}
524
525#[derive(Clone, Debug, PartialEq)]
542pub struct ProviderRef {
543 provider: Provider,
545 model: String,
549}
550
551impl ProviderRef {
552 pub fn registered(id: ProviderId, model: impl Into<String>) -> Result<Self, RefError> {
555 Self::new(Provider::Registered(id), model.into())
556 }
557
558 pub fn configured(config: ProviderConfig, model: impl Into<String>) -> Result<Self, RefError> {
566 Self::new(
567 Provider::Configured(config.with_credential("")),
568 model.into(),
569 )
570 }
571
572 pub fn provider(&self) -> &Provider {
574 &self.provider
575 }
576
577 fn new(provider: Provider, model: String) -> Result<Self, RefError> {
578 if model.is_empty() {
579 return Err(RefError::EmptyModel);
580 }
581 Ok(Self { provider, model })
582 }
583
584 pub fn model(&self) -> &str {
586 &self.model
587 }
588
589 pub fn parse(text: &str) -> Result<Self, RefError> {
594 let Some((selection, model)) = text.split_once(':') else {
595 return Err(RefError::NoModel {
596 reference: text.to_owned(),
597 });
598 };
599 if model.is_empty() {
600 return Err(RefError::NoModel {
601 reference: text.to_owned(),
602 });
603 }
604 Self::registered(ProviderId::resolve(selection)?, model)
605 }
606
607 pub fn id(&self) -> Option<ProviderId> {
609 match &self.provider {
610 Provider::Registered(id) => Some(*id),
611 Provider::Configured(config) => config.id(),
612 }
613 }
614
615 pub fn config(&self, api_key: impl Into<Secret>) -> ProviderConfig {
618 match &self.provider {
619 Provider::Registered(id) => id.config(api_key),
620 Provider::Configured(config) => config.clone().with_credential(api_key),
621 }
622 }
623}
624
625impl ProviderRef {
626 #[cfg(feature = "reqwest")]
631 #[cfg_attr(docsrs, doc(cfg(feature = "reqwest")))]
632 pub fn completion_model(&self) -> Result<DynModel<Completion>, EnvError> {
633 let config = match &self.provider {
634 Provider::Registered(id) => id.config_from_env()?,
635 Provider::Configured(config) => config.clone().with_credential_from_env()?,
636 };
637 Ok(config.completion_model(&self.model, rig_reqwest::shared()))
638 }
639
640 pub fn completion_model_with(
643 &self,
644 api_key: impl Into<Secret>,
645 http: impl HttpClientExt + 'static,
646 ) -> DynModel<Completion> {
647 self.config(api_key)
648 .completion_model(&self.model, DynHttpClient::new(http))
649 }
650}
651
652impl std::str::FromStr for ProviderRef {
653 type Err = RefError;
654
655 fn from_str(text: &str) -> Result<Self, Self::Err> {
656 Self::parse(text)
657 }
658}
659
660impl fmt::Display for ProviderRef {
663 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
664 match &self.provider {
665 Provider::Registered(id) => write!(f, "{id}:{}", self.model),
666 Provider::Configured(config) => {
667 write!(f, "{}/{}:{}", config.vendor(), config.format(), self.model)
668 }
669 }
670 }
671}
672
673#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
675pub enum RefError {
676 #[error("model identifier must not be empty")]
678 EmptyModel,
679 #[error("`{reference}` names no model: expected `vendor[/format]:model`")]
681 NoModel {
682 reference: String,
684 },
685 #[error(transparent)]
687 Selection(#[from] SelectionError),
688}
689
690const REF_FIELDS: &[&str] = &["config", "model"];
693
694impl Serialize for ProviderRef {
697 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
698 match &self.provider {
699 Provider::Registered(_) => serializer.collect_str(self),
700 Provider::Configured(config) => {
701 let mut object = serializer.serialize_struct("ProviderRef", 2)?;
702 object.serialize_field("config", config)?;
703 object.serialize_field("model", &self.model)?;
704 object.end()
705 }
706 }
707 }
708}
709
710impl<'de> Deserialize<'de> for ProviderRef {
713 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
714 deserializer.deserialize_any(RefVisitor)
715 }
716}
717
718struct RefVisitor;
719
720impl<'de> Visitor<'de> for RefVisitor {
721 type Value = ProviderRef;
722
723 fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
724 f.write_str(
725 "a provider reference `vendor[/format]:model`, or an object with `config` and `model`",
726 )
727 }
728
729 fn visit_str<E: de::Error>(self, text: &str) -> Result<Self::Value, E> {
730 ProviderRef::parse(text).map_err(de::Error::custom)
731 }
732
733 fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
734 let mut config: Option<ProviderConfig> = None;
735 let mut model: Option<String> = None;
736 while let Some(field) = map.next_key::<String>()? {
737 match field.as_str() {
738 "config" => {
739 if config.is_some() {
740 return Err(de::Error::duplicate_field("config"));
741 }
742 config = Some(map.next_value()?);
743 }
744 "model" => {
745 if model.is_some() {
746 return Err(de::Error::duplicate_field("model"));
747 }
748 model = Some(map.next_value()?);
749 }
750 unknown => return Err(de::Error::unknown_field(unknown, REF_FIELDS)),
751 }
752 }
753 ProviderRef::configured(
754 config.ok_or_else(|| de::Error::missing_field("config"))?,
755 model.ok_or_else(|| de::Error::missing_field("model"))?,
756 )
757 .map_err(de::Error::custom)
758 }
759}
760
761#[cfg(test)]
762mod tests;