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}