Skip to main content

rig_core/providers/openai/
client.rs

1use super::responses_api::{ResponsesProviderExt, SystemInstructionsPlacement};
2use crate::{
3    client::{
4        self, BearerAuth, Capabilities, Capable, DebugExt, Nothing, Provider, ProviderBuilder,
5        ProviderClient,
6    },
7    http_client::{self, HttpClientExt},
8    wasm_compat::{WasmCompatSend, WasmCompatSync},
9};
10use serde::Deserialize;
11use std::fmt::Debug;
12
13#[cfg(all(not(target_family = "wasm"), feature = "websocket"))]
14use crate::client::completion::CompletionClient;
15
16// ================================================================
17// Main OpenAI Client
18// ================================================================
19const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1";
20
21// ================================================================
22// OpenAI Responses API Extension
23// ================================================================
24#[derive(Debug, Default, Clone, Copy)]
25pub struct OpenAIResponsesExt {
26    pub(crate) system_instructions_placement: SystemInstructionsPlacement,
27}
28
29#[derive(Debug, Default, Clone, Copy)]
30pub struct OpenAIResponsesExtBuilder;
31
32// ================================================================
33// OpenAI Completions API Extension
34// ================================================================
35#[derive(Debug, Default, Clone, Copy)]
36pub struct OpenAICompletionsExt {
37    /// Carried through API switches so that a placement configured on a
38    /// Responses client survives `completions_api()` → `responses_api()`
39    /// round trips. Not used by Chat Completions requests themselves.
40    pub(crate) system_instructions_placement: SystemInstructionsPlacement,
41}
42
43#[derive(Debug, Default, Clone, Copy)]
44pub struct OpenAICompletionsExtBuilder;
45
46type OpenAIApiKey = BearerAuth;
47
48// Responses API client (default)
49pub type Client<H = reqwest::Client> = client::Client<OpenAIResponsesExt, H>;
50pub type ClientBuilder<H = crate::markers::Missing> =
51    client::ClientBuilder<OpenAIResponsesExtBuilder, OpenAIApiKey, H>;
52
53// Completions API client
54pub type CompletionsClient<H = reqwest::Client> = client::Client<OpenAICompletionsExt, H>;
55pub type CompletionsClientBuilder<H = crate::markers::Missing> =
56    client::ClientBuilder<OpenAICompletionsExtBuilder, OpenAIApiKey, H>;
57
58impl Provider for OpenAIResponsesExt {
59    type Builder = OpenAIResponsesExtBuilder;
60    const VERIFY_PATH: &'static str = "/models";
61}
62
63impl ResponsesProviderExt for OpenAIResponsesExt {
64    fn system_instructions_placement(&self) -> SystemInstructionsPlacement {
65        self.system_instructions_placement
66    }
67}
68
69impl Provider for OpenAICompletionsExt {
70    type Builder = OpenAICompletionsExtBuilder;
71    const VERIFY_PATH: &'static str = "/models";
72}
73
74impl<H> Capabilities<H> for OpenAIResponsesExt {
75    type Completion = Capable<super::responses_api::ResponsesCompletionModel<H>>;
76    type Embeddings = Capable<super::EmbeddingModel<H>>;
77    type Transcription = Capable<super::TranscriptionModel<H>>;
78    type ModelListing = Capable<super::OpenAIModelLister<H>>;
79    #[cfg(feature = "image")]
80    type ImageGeneration = Capable<super::ImageGenerationModel<H>>;
81    #[cfg(feature = "audio")]
82    type AudioGeneration = Capable<super::audio_generation::AudioGenerationModel<H>>;
83    type Rerank = Nothing;
84}
85
86impl<H> Capabilities<H> for OpenAICompletionsExt {
87    type Completion = Capable<super::completion::CompletionModel<H>>;
88    type Embeddings = Capable<super::GenericEmbeddingModel<OpenAICompletionsExt, H>>;
89    type Transcription = Capable<super::TranscriptionModel<H>>;
90    type ModelListing = Capable<super::OpenAIModelLister<H>>;
91    #[cfg(feature = "image")]
92    type ImageGeneration = Capable<super::ImageGenerationModel<H>>;
93    #[cfg(feature = "audio")]
94    type AudioGeneration = Capable<super::audio_generation::AudioGenerationModel<H>>;
95    type Rerank = Nothing;
96}
97
98impl DebugExt for OpenAIResponsesExt {}
99
100impl DebugExt for OpenAICompletionsExt {}
101
102impl ProviderBuilder for OpenAIResponsesExtBuilder {
103    type Extension<H>
104        = OpenAIResponsesExt
105    where
106        H: HttpClientExt;
107    type ApiKey = OpenAIApiKey;
108
109    const BASE_URL: &'static str = OPENAI_API_BASE_URL;
110
111    fn build<H>(
112        _builder: &client::ClientBuilder<Self, Self::ApiKey, H>,
113    ) -> http_client::Result<Self::Extension<H>>
114    where
115        H: HttpClientExt,
116    {
117        Ok(OpenAIResponsesExt::default())
118    }
119}
120
121impl ProviderBuilder for OpenAICompletionsExtBuilder {
122    type Extension<H>
123        = OpenAICompletionsExt
124    where
125        H: HttpClientExt;
126    type ApiKey = OpenAIApiKey;
127
128    const BASE_URL: &'static str = OPENAI_API_BASE_URL;
129
130    fn build<H>(
131        _builder: &client::ClientBuilder<Self, Self::ApiKey, H>,
132    ) -> http_client::Result<Self::Extension<H>>
133    where
134        H: HttpClientExt,
135    {
136        Ok(OpenAICompletionsExt::default())
137    }
138}
139
140impl<H> Client<H>
141where
142    H: HttpClientExt
143        + Clone
144        + std::fmt::Debug
145        + Default
146        + WasmCompatSend
147        + WasmCompatSync
148        + 'static,
149{
150    /// Sets where Rig system instructions are placed in Responses requests for
151    /// every completion model created from this client. Models capture the
152    /// placement when they are created, so models built before this call are
153    /// unaffected. See [`SystemInstructionsPlacement`] for when each placement applies.
154    pub fn with_system_instructions_placement(
155        self,
156        placement: SystemInstructionsPlacement,
157    ) -> Self {
158        let mut ext = *self.ext();
159        ext.system_instructions_placement = placement;
160        self.with_ext(ext)
161    }
162
163    /// Sends Rig system instructions as `system` messages in `input` instead of
164    /// as top-level Responses API `instructions` for every completion model
165    /// created from this client. Models built before this call are unaffected.
166    ///
167    /// OpenAI's Responses API supports `instructions`, and Rig uses it by
168    /// default. Use this compatibility fallback for OpenAI-compatible providers
169    /// that reject or ignore top-level `instructions`.
170    pub fn with_system_instructions_as_messages(self) -> Self {
171        self.with_system_instructions_placement(SystemInstructionsPlacement::InputSystemMessages)
172    }
173
174    /// Create a Completions API client from this Responses API client.
175    /// Useful for switching to the traditional Chat Completions API.
176    pub fn completions_api(self) -> CompletionsClient<H> {
177        let system_instructions_placement = self.ext().system_instructions_placement;
178        self.with_ext(OpenAICompletionsExt {
179            system_instructions_placement,
180        })
181    }
182}
183
184#[cfg(all(not(target_family = "wasm"), feature = "websocket"))]
185impl Client<reqwest::Client> {
186    /// WebSocket mode currently uses a native `tokio-tungstenite` transport and does
187    /// not reuse custom `HttpClientExt` backends, so this API is only exposed for the
188    /// default `reqwest::Client` transport.
189    pub fn responses_websocket_builder(
190        &self,
191        model: impl Into<String>,
192    ) -> super::responses_api::websocket::ResponsesWebSocketSessionBuilder {
193        super::responses_api::websocket::ResponsesWebSocketSessionBuilder::new(
194            self.completion_model(model),
195        )
196    }
197
198    /// This API is OpenAI-specific and only available on non-wasm targets in `rig-core`.
199    pub async fn responses_websocket(
200        &self,
201        model: impl Into<String>,
202    ) -> Result<
203        super::responses_api::websocket::ResponsesWebSocketSession,
204        crate::completion::CompletionError,
205    > {
206        self.responses_websocket_builder(model).connect().await
207    }
208}
209
210impl<H> CompletionsClient<H>
211where
212    H: HttpClientExt
213        + Clone
214        + std::fmt::Debug
215        + Default
216        + WasmCompatSend
217        + WasmCompatSync
218        + 'static,
219{
220    /// Create a Responses API client from this Completions API client.
221    /// Useful for switching to the newer Responses API. A system-instructions
222    /// placement configured before switching to the Completions API is
223    /// restored.
224    pub fn responses_api(self) -> Client<H> {
225        let system_instructions_placement = self.ext().system_instructions_placement;
226        self.with_ext(OpenAIResponsesExt {
227            system_instructions_placement,
228        })
229    }
230}
231
232impl ProviderClient for Client {
233    type Input = OpenAIApiKey;
234    type Error = crate::client::ProviderClientError;
235
236    /// Create a new OpenAI Responses API client from the `OPENAI_API_KEY` environment variable.
237    fn from_env() -> Result<Self, Self::Error> {
238        let base_url = crate::client::optional_env_var("OPENAI_BASE_URL")?;
239        let api_key = crate::client::required_env_var("OPENAI_API_KEY")?;
240
241        let mut builder = Client::builder().api_key(&api_key);
242
243        if let Some(base) = base_url {
244            builder = builder.base_url(&base);
245        }
246
247        builder.build().map_err(Into::into)
248    }
249
250    fn from_val(input: Self::Input) -> Result<Self, Self::Error> {
251        Self::new(input).map_err(Into::into)
252    }
253}
254
255impl ProviderClient for CompletionsClient {
256    type Input = OpenAIApiKey;
257    type Error = crate::client::ProviderClientError;
258
259    /// Create a new OpenAI Completions API client from the `OPENAI_API_KEY` environment variable.
260    fn from_env() -> Result<Self, Self::Error> {
261        let base_url = crate::client::optional_env_var("OPENAI_BASE_URL")?;
262        let api_key = crate::client::required_env_var("OPENAI_API_KEY")?;
263
264        let mut builder = CompletionsClient::builder().api_key(&api_key);
265
266        if let Some(base) = base_url {
267            builder = builder.base_url(&base);
268        }
269
270        builder.build().map_err(Into::into)
271    }
272
273    fn from_val(input: Self::Input) -> Result<Self, Self::Error> {
274        Self::new(input).map_err(Into::into)
275    }
276}
277
278/// Error envelope returned by OpenAI-compatible providers alongside 2xx
279/// statuses. Providers spell the message field differently (`message`,
280/// `error`, nested objects), so anything that isn't a valid success payload
281/// is treated as an error envelope and the raw body is preserved for the
282/// caller; `message` is only used for logging.
283#[derive(Debug, Deserialize)]
284pub struct ApiErrorResponse {
285    #[serde(default, alias = "error", deserialize_with = "error_message_or_value")]
286    pub(crate) message: String,
287}
288
289fn error_message_or_value<'de, D>(deserializer: D) -> Result<String, D::Error>
290where
291    D: serde::Deserializer<'de>,
292{
293    let value = serde_json::Value::deserialize(deserializer)?;
294    Ok(match value {
295        serde_json::Value::String(message) => message,
296        other => other.to_string(),
297    })
298}
299
300#[derive(Debug, Deserialize)]
301#[serde(untagged)]
302pub(crate) enum ApiResponse<T> {
303    Ok(T),
304    Err(ApiErrorResponse),
305}
306
307#[cfg(test)]
308mod tests {
309    use crate::client::{CompletionClient, EmbeddingsClient};
310    use crate::message::ImageDetail;
311    use crate::providers::openai::{
312        AssistantContent, Function, ImageUrl, Message, ToolCall, ToolType, UserContent,
313    };
314    use crate::{OneOrMany, message};
315    use serde_path_to_error::deserialize;
316
317    #[test]
318    fn test_deserialize_message() {
319        let assistant_message_json = r#"
320        {
321            "role": "assistant",
322            "content": "\n\nHello there, how may I assist you today?"
323        }
324        "#;
325
326        let assistant_message_json2 = r#"
327        {
328            "role": "assistant",
329            "content": [
330                {
331                    "type": "text",
332                    "text": "\n\nHello there, how may I assist you today?"
333                }
334            ],
335            "tool_calls": null
336        }
337        "#;
338
339        let assistant_message_json3 = r#"
340        {
341            "role": "assistant",
342            "tool_calls": [
343                {
344                    "id": "call_h89ipqYUjEpCPI6SxspMnoUU",
345                    "type": "function",
346                    "function": {
347                        "name": "subtract",
348                        "arguments": "{\"x\": 2, \"y\": 5}"
349                    }
350                }
351            ],
352            "content": null,
353            "refusal": null
354        }
355        "#;
356
357        let user_message_json = r#"
358        {
359            "role": "user",
360            "content": [
361                {
362                    "type": "text",
363                    "text": "What's in this image?"
364                },
365                {
366                    "type": "image_url",
367                    "image_url": {
368                        "url": "https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg"
369                    }
370                },
371                {
372                    "type": "audio",
373                    "input_audio": {
374                        "data": "...",
375                        "format": "mp3"
376                    }
377                }
378            ]
379        }
380        "#;
381
382        let assistant_message: Message = {
383            let jd = &mut serde_json::Deserializer::from_str(assistant_message_json);
384            deserialize(jd).unwrap_or_else(|err| {
385                panic!(
386                    "Deserialization error at {} ({}:{}): {}",
387                    err.path(),
388                    err.inner().line(),
389                    err.inner().column(),
390                    err
391                );
392            })
393        };
394
395        let assistant_message2: Message = {
396            let jd = &mut serde_json::Deserializer::from_str(assistant_message_json2);
397            deserialize(jd).unwrap_or_else(|err| {
398                panic!(
399                    "Deserialization error at {} ({}:{}): {}",
400                    err.path(),
401                    err.inner().line(),
402                    err.inner().column(),
403                    err
404                );
405            })
406        };
407
408        let assistant_message3: Message = {
409            let jd: &mut serde_json::Deserializer<serde_json::de::StrRead<'_>> =
410                &mut serde_json::Deserializer::from_str(assistant_message_json3);
411            deserialize(jd).unwrap_or_else(|err| {
412                panic!(
413                    "Deserialization error at {} ({}:{}): {}",
414                    err.path(),
415                    err.inner().line(),
416                    err.inner().column(),
417                    err
418                );
419            })
420        };
421
422        let user_message: Message = {
423            let jd = &mut serde_json::Deserializer::from_str(user_message_json);
424            deserialize(jd).unwrap_or_else(|err| {
425                panic!(
426                    "Deserialization error at {} ({}:{}): {}",
427                    err.path(),
428                    err.inner().line(),
429                    err.inner().column(),
430                    err
431                );
432            })
433        };
434
435        match assistant_message {
436            Message::Assistant { content, .. } => {
437                assert_eq!(
438                    content[0],
439                    AssistantContent::Text {
440                        text: "\n\nHello there, how may I assist you today?".to_string()
441                    }
442                );
443            }
444            _ => panic!("Expected assistant message"),
445        }
446
447        match assistant_message2 {
448            Message::Assistant {
449                content,
450                tool_calls,
451                ..
452            } => {
453                assert_eq!(
454                    content[0],
455                    AssistantContent::Text {
456                        text: "\n\nHello there, how may I assist you today?".to_string()
457                    }
458                );
459
460                assert_eq!(tool_calls, vec![]);
461            }
462            _ => panic!("Expected assistant message"),
463        }
464
465        match assistant_message3 {
466            Message::Assistant {
467                content,
468                tool_calls,
469                refusal,
470                ..
471            } => {
472                assert!(content.is_empty());
473                assert!(refusal.is_none());
474                assert_eq!(
475                    tool_calls[0],
476                    ToolCall {
477                        id: "call_h89ipqYUjEpCPI6SxspMnoUU".to_string(),
478                        r#type: ToolType::Function,
479                        function: Function {
480                            name: "subtract".to_string(),
481                            arguments: serde_json::json!({"x": 2, "y": 5}),
482                        },
483                    }
484                );
485            }
486            _ => panic!("Expected assistant message"),
487        }
488
489        match user_message {
490            Message::User { content, .. } => {
491                let (first, second) = {
492                    let mut iter = content.into_iter();
493                    (iter.next().unwrap(), iter.next().unwrap())
494                };
495                assert_eq!(
496                    first,
497                    UserContent::Text {
498                        text: "What's in this image?".to_string()
499                    }
500                );
501                assert_eq!(second, UserContent::Image { image_url: ImageUrl { url: "https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg".to_string(), detail: None } });
502            }
503            _ => panic!("Expected user message"),
504        }
505    }
506
507    #[test]
508    fn test_message_to_message_conversion() {
509        let user_message = message::Message::User {
510            content: OneOrMany::one(message::UserContent::text("Hello")),
511        };
512
513        let assistant_message = message::Message::Assistant {
514            id: None,
515            content: OneOrMany::one(message::AssistantContent::text("Hi there!")),
516        };
517
518        let converted_user_message: Vec<Message> = user_message.clone().try_into().unwrap();
519        let converted_assistant_message: Vec<Message> =
520            assistant_message.clone().try_into().unwrap();
521
522        match converted_user_message[0].clone() {
523            Message::User { content, .. } => {
524                assert_eq!(
525                    content.first(),
526                    UserContent::Text {
527                        text: "Hello".to_string()
528                    }
529                );
530            }
531            _ => panic!("Expected user message"),
532        }
533
534        match converted_assistant_message[0].clone() {
535            Message::Assistant { content, .. } => {
536                assert_eq!(
537                    content[0].clone(),
538                    AssistantContent::Text {
539                        text: "Hi there!".to_string()
540                    }
541                );
542            }
543            _ => panic!("Expected assistant message"),
544        }
545
546        let original_user_message: message::Message =
547            converted_user_message[0].clone().try_into().unwrap();
548        let original_assistant_message: message::Message =
549            converted_assistant_message[0].clone().try_into().unwrap();
550
551        assert_eq!(original_user_message, user_message);
552        assert_eq!(original_assistant_message, assistant_message);
553    }
554
555    #[test]
556    fn test_message_from_message_conversion() {
557        let user_message = Message::User {
558            content: OneOrMany::one(UserContent::Text {
559                text: "Hello".to_string(),
560            }),
561            name: None,
562        };
563
564        let assistant_message = Message::Assistant {
565            content: vec![AssistantContent::Text {
566                text: "Hi there!".to_string(),
567            }],
568            reasoning: None,
569            refusal: None,
570            audio: None,
571            name: None,
572            tool_calls: vec![],
573            reasoning_details: vec![],
574            images: vec![],
575        };
576
577        let converted_user_message: message::Message = user_message.clone().try_into().unwrap();
578        let converted_assistant_message: message::Message =
579            assistant_message.clone().try_into().unwrap();
580
581        match converted_user_message.clone() {
582            message::Message::User { content } => {
583                assert_eq!(content.first(), message::UserContent::text("Hello"));
584            }
585            _ => panic!("Expected user message"),
586        }
587
588        match converted_assistant_message.clone() {
589            message::Message::Assistant { content, .. } => {
590                assert_eq!(
591                    content.first(),
592                    message::AssistantContent::text("Hi there!")
593                );
594            }
595            _ => panic!("Expected assistant message"),
596        }
597
598        let original_user_message: Vec<Message> = converted_user_message.try_into().unwrap();
599        let original_assistant_message: Vec<Message> =
600            converted_assistant_message.try_into().unwrap();
601
602        assert_eq!(original_user_message[0], user_message);
603        assert_eq!(original_assistant_message[0], assistant_message);
604    }
605
606    #[test]
607    fn test_user_message_single_text_serializes_as_string() {
608        let user_message = Message::User {
609            content: OneOrMany::one(UserContent::Text {
610                text: "Hello world".to_string(),
611            }),
612            name: None,
613        };
614
615        let serialized = serde_json::to_value(&user_message).unwrap();
616
617        assert_eq!(serialized["role"], "user");
618        assert_eq!(serialized["content"], "Hello world");
619    }
620
621    #[test]
622    fn test_user_message_multiple_parts_serializes_as_array() {
623        let user_message = Message::User {
624            content: OneOrMany::many(vec![
625                UserContent::Text {
626                    text: "What's in this image?".to_string(),
627                },
628                UserContent::Image {
629                    image_url: ImageUrl {
630                        url: "https://example.com/image.jpg".to_string(),
631                        detail: Some(ImageDetail::default()),
632                    },
633                },
634            ])
635            .unwrap(),
636            name: None,
637        };
638
639        let serialized = serde_json::to_value(&user_message).unwrap();
640
641        assert_eq!(serialized["role"], "user");
642        assert!(serialized["content"].is_array());
643        assert_eq!(serialized["content"].as_array().unwrap().len(), 2);
644    }
645
646    #[test]
647    fn test_user_message_single_image_serializes_as_array() {
648        let user_message = Message::User {
649            content: OneOrMany::one(UserContent::Image {
650                image_url: ImageUrl {
651                    url: "https://example.com/image.jpg".to_string(),
652                    detail: Some(ImageDetail::default()),
653                },
654            }),
655            name: None,
656        };
657
658        let serialized = serde_json::to_value(&user_message).unwrap();
659
660        assert_eq!(serialized["role"], "user");
661        // Single non-text content should still serialize as array
662        assert!(serialized["content"].is_array());
663    }
664    #[test]
665    fn test_client_initialization() {
666        let _client =
667            crate::providers::openai::Client::new("dummy-key").expect("Client::new() failed");
668        let _client_from_builder = crate::providers::openai::Client::builder()
669            .api_key("dummy-key")
670            .build()
671            .expect("Client::builder() failed");
672    }
673
674    #[test]
675    fn test_legacy_chat_completion_model_type_annotation_still_compiles() {
676        let client = crate::providers::openai::Client::new("dummy-key")
677            .expect("Client::new() failed")
678            .completions_api();
679
680        let _model: crate::providers::openai::completion::CompletionModel<reqwest::Client> =
681            client.completion_model("gpt-4o");
682    }
683
684    #[test]
685    fn test_legacy_embedding_model_type_annotation_still_compiles() {
686        let client =
687            crate::providers::openai::Client::new("dummy-key").expect("Client::new() failed");
688
689        let _model: crate::providers::openai::EmbeddingModel<reqwest::Client> =
690            client.embedding_model(crate::providers::openai::TEXT_EMBEDDING_3_SMALL);
691    }
692}