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
16const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1";
20
21#[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#[derive(Debug, Default, Clone, Copy)]
36pub struct OpenAICompletionsExt {
37 pub(crate) system_instructions_placement: SystemInstructionsPlacement,
41}
42
43#[derive(Debug, Default, Clone, Copy)]
44pub struct OpenAICompletionsExtBuilder;
45
46type OpenAIApiKey = BearerAuth;
47
48pub type Client<H = reqwest::Client> = client::Client<OpenAIResponsesExt, H>;
50pub type ClientBuilder<H = crate::markers::Missing> =
51 client::ClientBuilder<OpenAIResponsesExtBuilder, OpenAIApiKey, H>;
52
53pub 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 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 pub fn with_system_instructions_as_messages(self) -> Self {
171 self.with_system_instructions_placement(SystemInstructionsPlacement::InputSystemMessages)
172 }
173
174 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 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 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 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 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 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#[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 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}