Skip to main content

rig_core/client/
openai.rs

1//! The OpenAI client: an [`OpenAIConfig`] on a transport, and the models it
2//! builds.
3
4use crate::client::macros::http_client;
5use crate::driver::Model;
6use crate::error::ProviderError;
7use crate::model::ModelList;
8use crate::providers::chatgpt::auth::{AuthError, Authenticator};
9use crate::providers::openai::responses_api::wire::Responses;
10
11#[cfg(feature = "image")]
12use crate::providers::openai::wire::Images;
13#[cfg(feature = "audio")]
14use crate::providers::openai::wire::Speech;
15use crate::providers::openai::wire::{
16    Chat, Embeddings, OpenAIConfig, OpenAiWire, Rerank, Transcriptions,
17};
18
19http_client!(
20    /// An OpenAI-shaped provider: its [`OpenAIConfig`] on a transport. Every
21    /// model it builds sends through that transport.
22    ///
23    /// ```no_run
24    /// use rig_core::completion::CompletionRequest;
25    /// use rig_core::providers::openai::{self, OpenAIConfig};
26    ///
27    /// # async fn run(http: rig_core::http_client::DynHttpClient) -> Result<(), Box<dyn std::error::Error>> {
28    /// let openai = OpenAIConfig::from_env()?.connect(http);
29    /// let model = openai.completion(openai::GPT_5_2);
30    /// let response = model.call(CompletionRequest::new("Capital of France?")).await?;
31    /// # let _ = response;
32    /// # Ok(())
33    /// # }
34    /// ```
35    OpenAI,
36    OpenAIConfig
37);
38
39impl OpenAI {
40    /// Official OpenAI with `api_key`, on the shared reqwest client.
41    #[cfg(feature = "reqwest")]
42    #[cfg_attr(docsrs, doc(cfg(feature = "reqwest")))]
43    pub fn new(api_key: impl Into<crate::wire::Secret>) -> Self {
44        OpenAIConfig::new(api_key).client()
45    }
46
47    /// Official OpenAI from `OPENAI_API_KEY` and the optional
48    /// `OPENAI_BASE_URL`, on the shared reqwest client.
49    #[cfg(feature = "reqwest")]
50    #[cfg_attr(docsrs, doc(cfg(feature = "reqwest")))]
51    pub fn from_env() -> Result<Self, crate::client::env::EnvError> {
52        Ok(OpenAIConfig::from_env()?.client())
53    }
54
55    /// The completion model for `model` on the configuration's
56    /// [`completion_route`](OpenAIConfig::completion_route).
57    pub fn completion(&self, model: impl Into<String>) -> Model<OpenAiWire> {
58        self.model(self.config.completion(model))
59    }
60
61    /// The Responses model for `model`: `POST /responses`.
62    pub fn responses(&self, model: impl Into<String>) -> Model<Responses> {
63        self.model(self.config.responses(model))
64    }
65
66    /// The Chat Completions model for `model`.
67    pub fn chat(&self, model: impl Into<String>) -> Model<Chat> {
68        self.model(self.config.chat(model))
69    }
70
71    /// The embedding model for `model`, `ndims` wide when set.
72    pub fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Model<Embeddings> {
73        self.model(self.config.embedding(model, ndims))
74    }
75
76    /// The rerank model for `model`.
77    pub fn rerank(&self, model: impl Into<String>) -> Model<Rerank> {
78        self.model(self.config.rerank(model))
79    }
80
81    /// The transcription model for `model`.
82    pub fn transcription(&self, model: impl Into<String>) -> Model<Transcriptions> {
83        self.model(self.config.transcription(model))
84    }
85
86    /// The image-generation model for `model`.
87    #[cfg(feature = "image")]
88    #[cfg_attr(docsrs, doc(cfg(feature = "image")))]
89    pub fn image_generation(&self, model: impl Into<String>) -> Model<Images> {
90        self.model(self.config.image_generation(model))
91    }
92
93    /// The speech model for `model`.
94    #[cfg(feature = "audio")]
95    #[cfg_attr(docsrs, doc(cfg(feature = "audio")))]
96    pub fn audio_generation(&self, model: impl Into<String>) -> Model<Speech> {
97        self.model(self.config.audio_generation(model))
98    }
99
100    /// The models this provider serves, every page followed.
101    pub async fn list_models(&self) -> Result<ModelList, ProviderError> {
102        self.model(self.config.models()).list().await
103    }
104
105    /// Check that the provider accepts the configured credential. A 401 or
106    /// 403 reply is [`ProviderError::InvalidAuthentication`].
107    pub async fn verify(&self) -> Result<(), ProviderError> {
108        self.model(self.config.verify()).verify().await
109    }
110
111    /// This client with the ChatGPT subscription credential `authenticator`
112    /// resolves, reading or refreshing it through this client's transport:
113    /// the access token, and the account it belongs to when one is named.
114    /// Every other setting is kept.
115    ///
116    /// ```no_run
117    /// use rig_core::providers::chatgpt::{self, auth::{AuthSource, Authenticator, DeviceCodeHandler}};
118    /// use rig_core::providers::openai::OpenAIConfig;
119    ///
120    /// # async fn run(http: rig_core::http_client::DynHttpClient) -> Result<(), Box<dyn std::error::Error>> {
121    /// let authenticator = Authenticator::new(AuthSource::OAuth, None, DeviceCodeHandler::default(), true);
122    /// let chatgpt = OpenAIConfig::with_key(&chatgpt::DIALECT, "")
123    ///     .connect(http)
124    ///     .authenticate(&authenticator)
125    ///     .await?;
126    /// # let _ = chatgpt;
127    /// # Ok(())
128    /// # }
129    /// ```
130    pub async fn authenticate(self, authenticator: &Authenticator) -> Result<Self, AuthError> {
131        let context = authenticator.auth_context(&self.http).await?;
132        let mut config = self.config;
133        config.api_key = context.access_token;
134        if let Some(account_id) = context.account_id {
135            config.account_id = Some(account_id);
136        }
137        Ok(Self {
138            config,
139            http: self.http,
140        })
141    }
142}