llm/providers/openai/
responses_provider.rs1use async_openai::Client;
2use async_openai::config::OpenAIConfig;
3use serde_json::Value;
4use std::future::ready;
5use tokio_stream::StreamExt;
6use tracing::debug;
7
8use crate::provider::{error_stream, get_context_window, stream_from};
9use crate::providers::openai_compatible::AetherOpenAiConfig;
10use crate::providers::openai_responses::mappers::{ResponsesRequestPolicy, build_wire_request};
11use crate::providers::openai_responses::streaming::{ResponsesStreamEvent, process_response_stream};
12use crate::{
13 Context, LlmError, LlmModel, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, ProviderFactory,
14 Result, StreamingModelProvider,
15};
16
17pub struct OpenAiProvider {
18 client: Client<AetherOpenAiConfig>,
19 model: String,
20}
21
22impl ProviderFactory for OpenAiProvider {
23 async fn from_env() -> Result<Self> {
24 Self::from_env_with_connection(ProviderConnectionConfig::default()).await
25 }
26
27 fn from_env_with_connection(connection: ProviderConnectionConfig) -> impl Future<Output = Result<Self>> + Send {
28 ready(provider_from_connection(connection))
29 }
30
31 fn with_model(mut self, model: &str) -> Self {
32 if !model.is_empty() {
33 self.model = model.to_string();
34 }
35 self
36 }
37}
38
39impl StreamingModelProvider for OpenAiProvider {
40 fn stream_response(&self, context: &Context) -> LlmResponseStream {
41 let client = self.client.clone();
42 let model = self.model.clone();
43 let request = match build_wire_request(&model, context, &ResponsesRequestPolicy::openai()) {
44 Ok(request) => request,
45 Err(e) => return error_stream(e),
46 };
47
48 stream_from(
49 async move {
50 debug!("Starting OpenAI Responses API stream for model: {model}");
51 client
52 .responses()
53 .create_stream_byot::<Value, ResponsesStreamEvent>(request)
54 .await
55 .map_err(|e| LlmError::ApiRequest(e.to_string()))
56 },
57 |stream| {
58 process_response_stream(Box::pin(
59 stream.map(|result| result.map_err(|e| LlmError::StreamInterrupted(e.to_string()))),
60 ))
61 },
62 )
63 }
64
65 fn display_name(&self) -> String {
66 format!("OpenAI ({})", self.model)
67 }
68
69 fn context_window(&self) -> Option<u32> {
70 get_context_window("openai", &self.model)
71 }
72
73 fn model(&self) -> Option<LlmModel> {
74 format!("openai:{}", self.model).parse().ok()
75 }
76}
77
78fn provider_from_connection(connection: ProviderConnectionConfig) -> Result<OpenAiProvider> {
79 let api_key = match connection.auth_mode {
80 ProviderAuthMode::Default => {
81 std::env::var("OPENAI_API_KEY").map_err(|_| LlmError::MissingApiKey("OPENAI_API_KEY".to_string()))?
82 }
83 ProviderAuthMode::None => String::new(),
84 };
85
86 let mut config = OpenAIConfig::new().with_api_key(api_key);
87 if let Some(base_url) = connection.base_url {
88 config = config.with_api_base(base_url);
89 }
90 let config = AetherOpenAiConfig::new(config, connection.auth_mode);
91
92 Ok(OpenAiProvider { client: Client::with_config(config), model: "gpt-4.1".to_string() })
93}
94
95#[cfg(test)]
96mod tests {
97 use super::*;
98 use crate::providers::test_capture_server::CaptureServer;
99 use crate::{ChatMessage, ReasoningEffort};
100
101 #[tokio::test]
102 async fn stream_response_sends_max_effort_on_the_wire() {
103 let mut server = CaptureServer::start().await;
104 let connection = ProviderConnectionConfig {
105 base_url: Some(server.base_url.clone()),
106 auth_mode: ProviderAuthMode::None,
107 ..Default::default()
108 };
109 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap().with_model("gpt-5.6");
110 let mut context = Context::new(vec![ChatMessage::user("Think harder")], vec![]);
111 context.set_reasoning_effort(Some(ReasoningEffort::Max));
112 context.set_prompt_cache_key(Some("cache-key".to_string()));
113
114 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
115 let captured = server.captured().await;
116
117 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
118 assert_eq!(captured.body["reasoning"]["effort"], "max");
119 assert_eq!(captured.body["model"], "gpt-5.6");
120 assert_eq!(captured.body["prompt_cache_key"], "cache-key");
121 assert_eq!(captured.body["stream"], true);
122 }
123
124 #[tokio::test]
125 async fn stream_response_surfaces_a_mapping_failure_as_the_only_item() {
126 let connection = ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() };
127 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap();
128 let context = Context::new(
129 vec![ChatMessage::User {
130 content: vec![crate::ContentBlock::Audio {
131 data: "YXVkaW8=".to_string(),
132 mime_type: "audio/wav".to_string(),
133 }],
134 timestamp: crate::types::IsoString::now(),
135 }],
136 vec![],
137 );
138
139 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
140
141 assert_eq!(responses.len(), 1);
142 assert!(matches!(responses[0], Err(LlmError::UnsupportedContent(_))), "{responses:?}");
143 }
144
145 #[test]
146 fn test_provider_display_name() {
147 let config = AetherOpenAiConfig::new(OpenAIConfig::new().with_api_key("test"), ProviderAuthMode::Default);
148 let provider = OpenAiProvider { client: Client::with_config(config), model: "gpt-4.1".to_string() };
149 assert_eq!(provider.display_name(), "OpenAI (gpt-4.1)");
150 }
151}