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
21use crate::catalog::ModelSpec;
22#[cfg(feature = "reqwest")]
23use crate::client::env::{self, EnvError};
24use crate::completion::ModelRef;
25use crate::driver::DynModel;
26use crate::http_client::{DynHttpClient, HttpClientExt};
27use crate::operation::Completion;
28use crate::providers::{anthropic, gemini, openai};
29use crate::serve::ErasedHandler;
30use crate::serve::adapters::ModelAdapter;
31use crate::wire::Secret;
32
33pub(crate) const OPENAI_DIALECTS: &[&openai::wire::Dialect] = &[
40 &openai::wire::OPENAI,
41 &openai::wire::AZURE,
42 &openai::wire::DEEPSEEK,
43 &openai::wire::GROQ,
44 &openai::wire::HYPERBOLIC,
45 &openai::wire::MIRA,
46 &openai::wire::PERPLEXITY,
47 &openai::wire::TOGETHER,
48 &openai::wire::HUGGINGFACE,
49 &openai::wire::LLAMACPP,
50 &openai::wire::MISTRAL,
51 &openai::wire::OPENROUTER,
52 &openai::wire::VENICE,
53 &openai::wire::DOUBLEWORD,
54 &openai::wire::ZAI,
55 &openai::wire::MINIMAX,
56 &openai::wire::MOONSHOT,
57 &openai::wire::XIAOMIMIMO,
58 &openai::wire::COHERE,
59 &openai::wire::OLLAMA,
60 &crate::providers::xai::DIALECT,
61 &crate::providers::chatgpt::DIALECT,
62 &crate::providers::copilot::wire::DIALECT,
63];
64
65#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
69pub enum Format {
70 OpenAi,
72 Anthropic,
74 Gemini,
76}
77
78impl Format {
79 pub const ALL: [Format; 3] = [Format::OpenAi, Format::Anthropic, Format::Gemini];
81
82 pub fn as_str(self) -> &'static str {
84 match self {
85 Format::OpenAi => "openai",
86 Format::Anthropic => "anthropic",
87 Format::Gemini => "gemini",
88 }
89 }
90
91 pub fn named(name: &str) -> Option<Self> {
93 Self::ALL.into_iter().find(|format| format.as_str() == name)
94 }
95
96 fn names() -> String {
98 Self::ALL
99 .iter()
100 .map(|format| format.as_str())
101 .collect::<Vec<_>>()
102 .join(", ")
103 }
104}
105
106impl fmt::Display for Format {
107 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108 f.write_str(self.as_str())
109 }
110}
111
112impl Serialize for Format {
113 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
114 serializer.serialize_str(self.as_str())
115 }
116}
117
118impl<'de> Deserialize<'de> for Format {
119 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
120 let name = String::deserialize(deserializer)?;
121 Self::named(&name).ok_or_else(|| {
122 de::Error::custom(format!(
123 "`{name}` is not a protocol family ({})",
124 Format::names()
125 ))
126 })
127 }
128}
129
130#[derive(Debug, Clone, Copy)]
132enum Registered {
133 OpenAi(openai::wire::Dialect),
134 Anthropic(anthropic::wire::Dialect),
135 Gemini,
136}
137
138impl PartialEq for Registered {
140 fn eq(&self, other: &Self) -> bool {
141 self.vendor() == other.vendor() && self.format() == other.format()
142 }
143}
144
145impl Registered {
146 fn vendor(&self) -> &'static str {
147 match self {
148 Self::OpenAi(dialect) => dialect.name,
149 Self::Anthropic(dialect) => dialect.name,
150 Self::Gemini => gemini::PROVIDER_NAME,
151 }
152 }
153
154 fn format(&self) -> Format {
155 match self {
156 Self::OpenAi(_) => Format::OpenAi,
157 Self::Anthropic(_) => Format::Anthropic,
158 Self::Gemini => Format::Gemini,
159 }
160 }
161
162 fn config(&self, api_key: impl Into<Secret>) -> ProviderConfig {
163 match self {
164 Self::OpenAi(dialect) => {
165 ProviderConfig::OpenAi(openai::wire::OpenAIConfig::with_key(dialect, api_key))
166 }
167 Self::Anthropic(dialect) => ProviderConfig::Anthropic(
168 anthropic::wire::AnthropicConfig::with_key(dialect, api_key),
169 ),
170 Self::Gemini => ProviderConfig::Gemini(gemini::GeminiConfig::new(api_key)),
171 }
172 }
173
174 #[cfg(feature = "reqwest")]
177 fn config_from_env(&self) -> Result<ProviderConfig, EnvError> {
178 Ok(match self {
179 Self::OpenAi(dialect) => {
180 let (api_key, auth) = openai_credential_from_env(dialect)?;
181 ProviderConfig::OpenAi(openai::wire::OpenAIConfig::from_env_with_credential(
182 dialect, api_key, auth,
183 )?)
184 }
185 Self::Anthropic(dialect) => {
186 ProviderConfig::Anthropic(anthropic::wire::AnthropicConfig::from_env_with(dialect)?)
187 }
188 Self::Gemini => ProviderConfig::Gemini(gemini::GeminiConfig::from_env()?),
189 })
190 }
191}
192
193#[derive(Debug)]
197struct CatalogOnly {
198 vendor: &'static str,
200 home: &'static str,
202 credential: bool,
204}
205
206static CATALOG_ONLY: [CatalogOnly; 5] = [
208 CatalogOnly {
209 vendor: "aws_bedrock",
210 home: "the `rig-bedrock` crate",
211 credential: true,
212 },
213 CatalogOnly {
214 vendor: "vertexai",
215 home: "the `rig-vertexai` crate",
216 credential: true,
217 },
218 CatalogOnly {
219 vendor: "gemini-grpc",
220 home: "the `rig-gemini-grpc` crate",
221 credential: true,
222 },
223 CatalogOnly {
224 vendor: "candle",
225 home: "the `rig-candle` crate",
226 credential: false,
227 },
228 CatalogOnly {
229 vendor: "voyageai",
230 home: "`rig_core::providers::voyageai`, which serves embeddings and reranking only",
231 credential: true,
232 },
233];
234
235#[derive(Debug, Clone, Copy)]
237enum Kind {
238 Registered(Registered),
240 CatalogOnly(&'static CatalogOnly),
242}
243
244#[derive(Debug, Clone, Copy)]
252pub struct ProviderId(Kind);
253
254impl PartialEq for ProviderId {
257 fn eq(&self, other: &Self) -> bool {
258 self.vendor() == other.vendor() && self.format() == other.format()
259 }
260}
261
262impl Eq for ProviderId {}
263
264impl Hash for ProviderId {
265 fn hash<H: Hasher>(&self, state: &mut H) {
266 self.vendor().hash(state);
267 self.format().hash(state);
268 }
269}
270
271impl ProviderId {
272 pub fn all() -> impl Iterator<Item = ProviderId> {
276 openai::wire::all()
277 .map(|dialect| Registered::OpenAi(*dialect))
278 .chain(anthropic::wire::all().map(|dialect| Registered::Anthropic(*dialect)))
279 .chain(std::iter::once(Registered::Gemini))
280 .map(|registered| ProviderId(Kind::Registered(registered)))
281 }
282
283 pub fn new(vendor: &str, format: Format) -> Option<Self> {
286 let registered = match format {
287 Format::OpenAi => {
288 openai::wire::by_name(vendor).map(|dialect| Registered::OpenAi(*dialect))
289 }
290 Format::Anthropic => {
291 anthropic::wire::Dialect::by_name(vendor).map(Registered::Anthropic)
292 }
293 Format::Gemini => (vendor == gemini::PROVIDER_NAME).then_some(Registered::Gemini),
294 };
295 registered.map(|registered| Self(Kind::Registered(registered)))
296 }
297
298 pub fn catalog(vendor: &str) -> Option<Self> {
303 Self::vendor_selections(vendor).next().or_else(|| {
304 CATALOG_ONLY
305 .iter()
306 .find(|provider| provider.vendor == vendor)
307 .map(|provider| Self(Kind::CatalogOnly(provider)))
308 })
309 }
310
311 pub fn vendor(&self) -> &'static str {
313 match &self.0 {
314 Kind::Registered(registered) => registered.vendor(),
315 Kind::CatalogOnly(provider) => provider.vendor,
316 }
317 }
318
319 pub fn format(&self) -> Option<Format> {
321 match &self.0 {
322 Kind::Registered(registered) => Some(registered.format()),
323 Kind::CatalogOnly(_) => None,
324 }
325 }
326
327 pub fn is_registered(&self) -> bool {
330 matches!(self.0, Kind::Registered(_))
331 }
332
333 pub(crate) fn served_by(&self) -> Option<&'static str> {
336 match &self.0 {
337 Kind::Registered(_) => None,
338 Kind::CatalogOnly(provider) => Some(provider.home),
339 }
340 }
341
342 pub fn vendor_selections(vendor: &str) -> impl Iterator<Item = ProviderId> + '_ {
344 Self::all().filter(move |id| id.vendor() == vendor)
345 }
346
347 pub fn resolve(selection: &str) -> Result<Self, SelectionError> {
354 let malformed = || SelectionError::Malformed {
355 selection: selection.to_owned(),
356 };
357 let (vendor, format) = match selection.split_once('/') {
358 Some((_, rest)) if rest.contains('/') => return Err(malformed()),
359 Some((vendor, format)) => (vendor, Some(format)),
360 None => (selection, None),
361 };
362 if vendor.is_empty() {
363 return Err(malformed());
364 }
365 let Some(format) = format else {
366 let mut registered = Self::vendor_selections(vendor);
367 let first = registered.next().ok_or_else(|| SelectionError::Unknown {
368 vendor: vendor.to_owned(),
369 })?;
370 return match registered.next() {
371 None => Ok(first),
372 Some(_) => Err(SelectionError::Ambiguous {
373 vendor: vendor.to_owned(),
374 alternatives: alternatives(vendor),
375 }),
376 };
377 };
378 let Some(family) = Format::named(format) else {
379 return Err(SelectionError::UnknownFormat {
380 format: format.to_owned(),
381 });
382 };
383 Self::new(vendor, family).ok_or_else(|| {
384 let alternatives = alternatives(vendor);
385 match alternatives.is_empty() {
386 true => SelectionError::Unknown {
387 vendor: vendor.to_owned(),
388 },
389 false => SelectionError::Unregistered {
390 vendor: vendor.to_owned(),
391 format: family,
392 alternatives,
393 },
394 }
395 })
396 }
397
398 pub fn config(&self, api_key: impl Into<Secret>) -> Option<ProviderConfig> {
402 match &self.0 {
403 Kind::Registered(registered) => Some(registered.config(api_key)),
404 Kind::CatalogOnly(_) => None,
405 }
406 }
407
408 pub fn api_key_env(&self) -> Option<&'static str> {
412 match &self.0 {
413 Kind::Registered(Registered::OpenAi(dialect)) => Some(dialect.api_key_env),
414 Kind::Registered(Registered::Anthropic(dialect)) => Some(dialect.api_key_env),
415 Kind::Registered(Registered::Gemini) => Some(gemini::API_KEY_ENV),
416 Kind::CatalogOnly(_) => None,
417 }
418 }
419
420 pub fn requires_credential(&self) -> bool {
425 match &self.0 {
426 Kind::Registered(Registered::OpenAi(dialect)) => {
427 !matches!(dialect.quirks.auth, openai::wire::Auth::OptionalBearer)
428 }
429 Kind::Registered(Registered::Anthropic(_) | Registered::Gemini) => true,
430 Kind::CatalogOnly(provider) => provider.credential,
431 }
432 }
433}
434
435fn alternatives(vendor: &str) -> Vec<String> {
437 ProviderId::vendor_selections(vendor)
438 .map(|id| id.to_string())
439 .collect()
440}
441
442impl fmt::Display for ProviderId {
443 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
444 match self.format() {
445 Some(format) => write!(f, "{}/{format}", self.vendor()),
446 None => f.write_str(self.vendor()),
447 }
448 }
449}
450
451impl Serialize for ProviderId {
452 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
453 serializer.collect_str(self)
454 }
455}
456
457impl<'de> Deserialize<'de> for ProviderId {
460 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
461 let selection = String::deserialize(deserializer)?;
462 Self::resolve(&selection)
463 .or_else(|error| {
464 Self::catalog(&selection)
465 .filter(|id| !id.is_registered())
466 .ok_or(error)
467 })
468 .map_err(de::Error::custom)
469 }
470}
471
472#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
474pub enum SelectionError {
475 #[error("no registered provider is named `{vendor}`")]
477 Unknown {
478 vendor: String,
480 },
481 #[error(
483 "`{vendor}` speaks no {format} endpoint in this build (it speaks {})",
484 alternatives.join(", ")
485 )]
486 Unregistered {
487 vendor: String,
489 format: Format,
491 alternatives: Vec<String>,
493 },
494 #[error(
496 "`{vendor}` names more than one registered selection: name one of {}",
497 alternatives.join(", ")
498 )]
499 Ambiguous {
500 vendor: String,
502 alternatives: Vec<String>,
504 },
505 #[error("`{format}` is not a protocol family ({})", Format::names())]
507 UnknownFormat {
508 format: String,
510 },
511 #[error("`{selection}` is not a provider selection: expected `vendor` or `vendor/format`")]
513 Malformed {
514 selection: String,
516 },
517}
518
519#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
524pub enum ProviderConfig {
525 #[serde(rename = "openai")]
527 OpenAi(openai::wire::OpenAIConfig),
528 #[serde(rename = "anthropic")]
530 Anthropic(anthropic::wire::AnthropicConfig),
531 #[serde(rename = "gemini")]
533 Gemini(gemini::GeminiConfig),
534}
535
536impl ProviderConfig {
537 pub fn id(&self) -> Option<ProviderId> {
541 ProviderId::new(self.vendor(), self.format())
542 }
543
544 pub fn vendor(&self) -> &'static str {
546 match self {
547 Self::OpenAi(provider) => provider.dialect.name,
548 Self::Anthropic(provider) => provider.dialect.name,
549 Self::Gemini(_) => gemini::PROVIDER_NAME,
550 }
551 }
552
553 pub fn format(&self) -> Format {
555 match self {
556 Self::OpenAi(_) => Format::OpenAi,
557 Self::Anthropic(_) => Format::Anthropic,
558 Self::Gemini(_) => Format::Gemini,
559 }
560 }
561
562 pub fn is_unauthenticated(&self) -> bool {
565 match self {
566 Self::OpenAi(provider) => provider.api_key.is_empty(),
567 Self::Anthropic(provider) => provider.api_key.is_empty(),
568 Self::Gemini(provider) => provider.api_key.is_empty(),
569 }
570 }
571
572 pub fn with_credential(mut self, api_key: impl Into<Secret>) -> Self {
575 let api_key = api_key.into();
576 match &mut self {
577 Self::OpenAi(provider) => provider.api_key = api_key,
578 Self::Anthropic(provider) => provider.api_key = api_key,
579 Self::Gemini(provider) => provider.api_key = api_key,
580 }
581 self
582 }
583
584 pub fn base_url(&self) -> &str {
593 match self {
594 Self::OpenAi(provider) => &provider.base_url,
595 Self::Anthropic(provider) => &provider.base_url,
596 Self::Gemini(provider) => &provider.base_url,
597 }
598 }
599
600 pub fn with_base_url(self, base_url: impl Into<String>) -> Self {
604 let base_url = base_url.into();
605 match self {
606 Self::OpenAi(provider) => Self::OpenAi(provider.with_base_url(base_url)),
607 Self::Anthropic(provider) => Self::Anthropic(provider.with_base_url(base_url)),
608 Self::Gemini(provider) => Self::Gemini(provider.with_base_url(base_url)),
609 }
610 }
611
612 pub fn completion_handler(
619 &self,
620 label: &str,
621 model: &str,
622 http: DynHttpClient,
623 ) -> ErasedHandler {
624 ErasedHandler::new(ModelAdapter::new(label, self.completion_model(model, http)))
625 }
626
627 fn completion_model(&self, model: &str, http: DynHttpClient) -> DynModel<Completion> {
629 match self {
630 Self::OpenAi(provider) => provider.clone().connect(http).completion(model).erase(),
631 Self::Anthropic(provider) => provider.clone().connect(http).completion(model).erase(),
632 Self::Gemini(provider) => provider.clone().connect(http).completion(model).erase(),
633 }
634 }
635
636 #[cfg(feature = "reqwest")]
640 fn with_credential_from_env(self) -> Result<Self, EnvError> {
641 Ok(match self {
642 Self::OpenAi(mut provider) => {
643 let (api_key, auth) = openai_credential_from_env(&provider.dialect)?;
644 provider.api_key = api_key.into();
645 if auth != provider.dialect.quirks.auth {
647 provider.auth = auth;
648 }
649 Self::OpenAi(provider)
650 }
651 Self::Anthropic(mut provider) => {
652 provider.api_key = env::required(provider.dialect.api_key_env)?.into();
653 Self::Anthropic(provider)
654 }
655 Self::Gemini(mut provider) => {
656 provider.api_key = env::required(gemini::API_KEY_ENV)?.into();
657 Self::Gemini(provider)
658 }
659 })
660 }
661}
662
663#[cfg(feature = "reqwest")]
667fn openai_credential_from_env(
668 dialect: &openai::wire::Dialect,
669) -> Result<(String, openai::wire::Auth), EnvError> {
670 if matches!(dialect.quirks.auth, openai::wire::Auth::OptionalBearer)
671 && dialect.alternate_auth.is_none()
672 {
673 let api_key = env::optional(dialect.api_key_env)?.unwrap_or_default();
674 return Ok((api_key, dialect.quirks.auth));
675 }
676 openai::wire::OpenAIConfig::credential_from_env(dialect)
677}
678
679#[derive(Clone, Debug, PartialEq)]
682pub enum Provider {
683 Registered(ProviderId),
685 Configured(ProviderConfig),
687}
688
689#[derive(Clone, Debug, PartialEq)]
706pub struct ProviderRef {
707 recipe: Recipe,
709 model: String,
713}
714
715#[derive(Clone, Debug, PartialEq)]
718enum Recipe {
719 Registered(Registered),
720 Configured(ProviderConfig),
721}
722
723impl ProviderRef {
724 pub fn registered(id: ProviderId, model: impl Into<String>) -> Result<Self, RefError> {
729 match id.0 {
730 Kind::Registered(registered) => Self::new(Recipe::Registered(registered), model.into()),
731 Kind::CatalogOnly(provider) => Err(RefError::Selection(SelectionError::Unknown {
732 vendor: provider.vendor.to_owned(),
733 })),
734 }
735 }
736
737 pub fn configured(config: ProviderConfig, model: impl Into<String>) -> Result<Self, RefError> {
745 Self::new(Recipe::Configured(config.with_credential("")), model.into())
746 }
747
748 pub fn provider(&self) -> Provider {
750 match &self.recipe {
751 Recipe::Registered(registered) => {
752 Provider::Registered(ProviderId(Kind::Registered(*registered)))
753 }
754 Recipe::Configured(config) => Provider::Configured(config.clone()),
755 }
756 }
757
758 fn new(recipe: Recipe, model: String) -> Result<Self, RefError> {
759 if model.is_empty() {
760 return Err(RefError::EmptyModel);
761 }
762 Ok(Self { recipe, model })
763 }
764
765 pub fn model(&self) -> &str {
767 &self.model
768 }
769
770 pub fn parse(text: &str) -> Result<Self, RefError> {
775 let Some((selection, model)) = text.split_once(':') else {
776 return Err(RefError::NoModel {
777 reference: text.to_owned(),
778 });
779 };
780 if model.is_empty() {
781 return Err(RefError::NoModel {
782 reference: text.to_owned(),
783 });
784 }
785 Self::registered(ProviderId::resolve(selection)?, model)
786 }
787
788 pub fn id(&self) -> Option<ProviderId> {
790 match &self.recipe {
791 Recipe::Registered(registered) => Some(ProviderId(Kind::Registered(*registered))),
792 Recipe::Configured(config) => config.id(),
793 }
794 }
795
796 pub fn config(&self, api_key: impl Into<Secret>) -> ProviderConfig {
799 match &self.recipe {
800 Recipe::Registered(registered) => registered.config(api_key),
801 Recipe::Configured(config) => config.clone().with_credential(api_key),
802 }
803 }
804}
805
806impl ProviderRef {
807 #[cfg(feature = "reqwest")]
812 #[cfg_attr(docsrs, doc(cfg(feature = "reqwest")))]
813 pub fn completion_model(&self) -> Result<DynModel<Completion>, EnvError> {
814 let config = match &self.recipe {
815 Recipe::Registered(registered) => registered.config_from_env()?,
816 Recipe::Configured(config) => config.clone().with_credential_from_env()?,
817 };
818 Ok(config.completion_model(&self.model, rig_reqwest::shared()))
819 }
820
821 pub fn completion_model_with(
824 &self,
825 api_key: impl Into<Secret>,
826 http: impl HttpClientExt + 'static,
827 ) -> DynModel<Completion> {
828 self.config(api_key)
829 .completion_model(&self.model, DynHttpClient::new(http))
830 }
831}
832
833impl std::str::FromStr for ProviderRef {
834 type Err = RefError;
835
836 fn from_str(text: &str) -> Result<Self, Self::Err> {
837 Self::parse(text)
838 }
839}
840
841impl fmt::Display for ProviderRef {
844 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
845 match &self.recipe {
846 Recipe::Registered(registered) => {
847 write!(
848 f,
849 "{}/{}:{}",
850 registered.vendor(),
851 registered.format(),
852 self.model
853 )
854 }
855 Recipe::Configured(config) => {
856 write!(f, "{}/{}:{}", config.vendor(), config.format(), self.model)
857 }
858 }
859 }
860}
861
862#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
864pub enum RefError {
865 #[error("model identifier must not be empty")]
867 EmptyModel,
868 #[error("`{reference}` names no model: expected `vendor[/format]:model`")]
870 NoModel {
871 reference: String,
873 },
874 #[error(transparent)]
876 Selection(#[from] SelectionError),
877}
878
879#[non_exhaustive]
883#[derive(Clone, Copy, Debug)]
884pub enum ModelSelector<'a> {
885 Spec(&'a ModelSpec),
887 Reference(&'a str),
890}
891
892impl<'a> From<&'a ModelSpec> for ModelSelector<'a> {
893 fn from(spec: &'a ModelSpec) -> Self {
894 Self::Spec(spec)
895 }
896}
897
898impl<'a> From<&'a str> for ModelSelector<'a> {
899 fn from(reference: &'a str) -> Self {
900 Self::Reference(reference)
901 }
902}
903
904impl<'a> From<&'a String> for ModelSelector<'a> {
905 fn from(reference: &'a String) -> Self {
906 Self::Reference(reference)
907 }
908}
909
910impl<'a> From<&'a ModelRef> for ModelSelector<'a> {
911 fn from(reference: &'a ModelRef) -> Self {
912 Self::Reference(reference.as_str())
913 }
914}
915
916impl ModelSelector<'_> {
917 pub fn provider_ref(self) -> Result<ProviderRef, ConnectError> {
920 let (id, model) = match self {
921 Self::Spec(spec) => (spec.provider, spec.id.as_str()),
922 Self::Reference(reference) => {
923 let (vendor, model) =
924 crate::catalog::split_reference(reference).ok_or_else(|| {
925 ConnectError::Malformed {
926 reference: reference.to_owned(),
927 }
928 })?;
929 if let Some(id) = ProviderId::catalog(vendor).filter(|id| !id.is_registered()) {
930 (id, model)
931 } else if selection_grammar(reference) {
932 return Ok(ProviderRef::parse(reference)?);
933 } else {
934 let id = ProviderId::catalog(vendor).ok_or_else(|| {
935 RefError::Selection(SelectionError::Unknown {
936 vendor: vendor.to_owned(),
937 })
938 })?;
939 (id, model)
940 }
941 }
942 };
943 match id.served_by() {
944 Some(served_by) => Err(ConnectError::CatalogOnly {
945 vendor: id.vendor().to_owned(),
946 served_by,
947 }),
948 None => Ok(ProviderRef::registered(id, model)?),
949 }
950 }
951}
952
953fn selection_grammar(reference: &str) -> bool {
955 reference.split_once(':').is_some_and(|(selection, _)| {
956 selection
957 .split_once('/')
958 .is_none_or(|(_, format)| Format::named(format).is_some())
959 })
960}
961
962#[non_exhaustive]
964#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
965pub enum ConnectError {
966 #[error("`{reference}` names no model: expected `vendor/model` or `vendor[/format]:model`")]
968 Malformed {
969 reference: String,
971 },
972 #[error("the registry cannot connect to `{vendor}`: its models are served by {served_by}")]
975 CatalogOnly {
976 vendor: String,
978 served_by: &'static str,
980 },
981 #[error(transparent)]
983 Reference(#[from] RefError),
984}
985
986#[cfg(feature = "reqwest")]
1000#[cfg_attr(docsrs, doc(cfg(feature = "reqwest")))]
1001pub fn connect<'a>(
1002 model: impl Into<ModelSelector<'a>>,
1003 api_key: impl Into<Secret>,
1004) -> Result<DynModel<Completion>, ConnectError> {
1005 let reference = model.into().provider_ref()?;
1006 Ok(reference
1007 .config(api_key)
1008 .completion_model(reference.model(), rig_reqwest::shared()))
1009}
1010
1011pub fn connect_with<'a>(
1014 model: impl Into<ModelSelector<'a>>,
1015 api_key: impl Into<Secret>,
1016 http: impl HttpClientExt + 'static,
1017) -> Result<DynModel<Completion>, ConnectError> {
1018 Ok(model
1019 .into()
1020 .provider_ref()?
1021 .completion_model_with(api_key, http))
1022}
1023
1024const REF_FIELDS: &[&str] = &["config", "model"];
1027
1028impl Serialize for ProviderRef {
1031 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
1032 match &self.recipe {
1033 Recipe::Registered(_) => serializer.collect_str(self),
1034 Recipe::Configured(config) => {
1035 let mut object = serializer.serialize_struct("ProviderRef", 2)?;
1036 object.serialize_field("config", config)?;
1037 object.serialize_field("model", &self.model)?;
1038 object.end()
1039 }
1040 }
1041 }
1042}
1043
1044impl<'de> Deserialize<'de> for ProviderRef {
1047 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
1048 deserializer.deserialize_any(RefVisitor)
1049 }
1050}
1051
1052struct RefVisitor;
1053
1054impl<'de> Visitor<'de> for RefVisitor {
1055 type Value = ProviderRef;
1056
1057 fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1058 f.write_str(
1059 "a provider reference `vendor[/format]:model`, or an object with `config` and `model`",
1060 )
1061 }
1062
1063 fn visit_str<E: de::Error>(self, text: &str) -> Result<Self::Value, E> {
1064 ProviderRef::parse(text).map_err(de::Error::custom)
1065 }
1066
1067 fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
1068 let mut config: Option<ProviderConfig> = None;
1069 let mut model: Option<String> = None;
1070 while let Some(field) = map.next_key::<String>()? {
1071 match field.as_str() {
1072 "config" => {
1073 if config.is_some() {
1074 return Err(de::Error::duplicate_field("config"));
1075 }
1076 config = Some(map.next_value()?);
1077 }
1078 "model" => {
1079 if model.is_some() {
1080 return Err(de::Error::duplicate_field("model"));
1081 }
1082 model = Some(map.next_value()?);
1083 }
1084 unknown => return Err(de::Error::unknown_field(unknown, REF_FIELDS)),
1085 }
1086 }
1087 ProviderRef::configured(
1088 config.ok_or_else(|| de::Error::missing_field("config"))?,
1089 model.ok_or_else(|| de::Error::missing_field("model"))?,
1090 )
1091 .map_err(de::Error::custom)
1092 }
1093}
1094
1095#[cfg(test)]
1096mod tests;