llm/providers/openai/
responses_provider.rs1use async_openai::config::{Config, OpenAIConfig};
2use std::future::ready;
3use tracing::debug;
4
5use crate::provider::{error_stream, get_context_window, stream_from};
6use crate::providers::openai_compatible::AetherOpenAiConfig;
7use crate::providers::openai_responses::mappers::{ResponsesRequestPolicy, build_wire_request};
8use crate::providers::openai_responses::transport::{process_connection, send};
9use crate::{
10 Context, LlmError, LlmModel, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, ProviderFactory,
11 Result, StreamingModelProvider,
12};
13use reqwest::Url;
14
15pub struct OpenAiProvider {
16 config: AetherOpenAiConfig,
17 http: reqwest::Client,
18 model: String,
19}
20
21impl ProviderFactory for OpenAiProvider {
22 async fn from_env() -> Result<Self> {
23 Self::from_env_with_connection(ProviderConnectionConfig::default()).await
24 }
25
26 fn from_env_with_connection(connection: ProviderConnectionConfig) -> impl Future<Output = Result<Self>> + Send {
27 ready(provider_from_connection(connection))
28 }
29
30 fn with_model(mut self, model: &str) -> Self {
31 if !model.is_empty() {
32 self.model = model.to_string();
33 }
34 self
35 }
36}
37
38impl StreamingModelProvider for OpenAiProvider {
39 fn stream_response(&self, context: &Context) -> LlmResponseStream {
40 let http = self.http.clone();
41 let mut url = match Url::parse(&self.config.url("/responses")) {
42 Ok(url) => url,
43 Err(error) => return error_stream(LlmError::ProviderRequest(error.to_string())),
44 };
45 url.query_pairs_mut().extend_pairs(self.config.query());
46 let url = url.to_string();
47 let headers = self.config.headers();
48 let model = self.model.clone();
49 let request = match build_wire_request(&model, context, &ResponsesRequestPolicy::openai()) {
50 Ok(request) => request,
51 Err(e) => return error_stream(e),
52 };
53
54 stream_from(
55 async move {
56 debug!("Starting OpenAI Responses API stream for model: {model}");
57 send(&http, &url, headers, request).await
58 },
59 process_connection,
60 )
61 }
62
63 fn display_name(&self) -> String {
64 format!("OpenAI ({})", self.model)
65 }
66
67 fn context_window(&self) -> Option<u32> {
68 get_context_window("openai", &self.model)
69 }
70
71 fn model(&self) -> Option<LlmModel> {
72 format!("openai:{}", self.model).parse().ok()
73 }
74}
75
76fn provider_from_connection(connection: ProviderConnectionConfig) -> Result<OpenAiProvider> {
77 let api_key = match connection.auth_mode {
78 ProviderAuthMode::Default => {
79 std::env::var("OPENAI_API_KEY").map_err(|_| LlmError::MissingApiKey("OPENAI_API_KEY".to_string()))?
80 }
81 ProviderAuthMode::None => String::new(),
82 };
83
84 let mut config = OpenAIConfig::new().with_api_key(api_key);
85 if let Some(base_url) = connection.base_url {
86 config = config.with_api_base(base_url);
87 }
88 let config = AetherOpenAiConfig::new(config, connection.auth_mode);
89 let http = reqwest::Client::new();
90
91 Ok(OpenAiProvider { config, http, model: "gpt-4.1".to_string() })
92}
93
94#[cfg(test)]
95mod tests {
96 use super::*;
97 use crate::providers::test_capture_server::CaptureServer;
98 use crate::{ChatMessage, ContentBlock, MessageId, ReasoningEffort};
99 use tokio_stream::StreamExt;
100
101 #[tokio::test]
102 async fn stream_response_distinguishes_default_disabled_and_low() {
103 for (effort, expected) in [
104 (ReasoningEffort::Default, None),
105 (ReasoningEffort::Disabled, Some("none")),
106 (ReasoningEffort::Low, Some("low")),
107 ] {
108 let mut server = CaptureServer::start_responses().await;
109 let provider = OpenAiProvider::from_env_with_connection(ProviderConnectionConfig {
110 base_url: Some(server.base_url.clone()),
111 auth_mode: ProviderAuthMode::None,
112 ..Default::default()
113 })
114 .await
115 .unwrap()
116 .with_model("gpt-5.4");
117 let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
118 context.set_reasoning_effort(effort);
119 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
120 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
121 let body = server.captured().await.body;
122 assert_eq!(body["reasoning"]["effort"].as_str(), expected);
123 if effort == ReasoningEffort::Disabled {
124 assert!(body["reasoning"]["summary"].is_null());
125 }
126 }
127 }
128
129 #[tokio::test]
130 async fn stream_response_sends_max_effort_on_the_wire() {
131 let mut server = CaptureServer::start_responses().await;
132 let connection = ProviderConnectionConfig {
133 base_url: Some(server.base_url.clone()),
134 auth_mode: ProviderAuthMode::None,
135 ..Default::default()
136 };
137 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap().with_model("gpt-5.6");
138 let mut context = Context::new(vec![ChatMessage::user("Think harder")], vec![]);
139 context.set_reasoning_effort(ReasoningEffort::Max);
140 context.set_prompt_cache_key(Some("cache-key".to_string()));
141
142 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
143 let captured = server.captured().await;
144
145 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
146 assert_eq!(captured.body["reasoning"]["effort"], "max");
147 assert_eq!(captured.body["model"], "gpt-5.6");
148 assert_eq!(captured.body["prompt_cache_key"], "cache-key");
149 assert_eq!(captured.body["stream"], true);
150 }
151
152 #[tokio::test]
153 async fn http_200_failed_server_error_is_retryable_with_request_id() {
154 use crate::providers::test_capture_server::ResponseSpec;
155 let spec = ResponseSpec::sse(include_str!("../../../tests/fixtures/openai_responses/04_failed_server.sse"))
156 .with_header("x-request-id", "req-openai-1");
157 let mut server = CaptureServer::start_with_spec(spec).await;
158 let connection = ProviderConnectionConfig {
159 base_url: Some(server.base_url.clone()),
160 auth_mode: ProviderAuthMode::None,
161 ..Default::default()
162 };
163 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap();
164 let context = Context::new(vec![ChatMessage::user("hi")], vec![]);
165
166 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
167 let _ = server.captured().await;
168
169 assert!(!responses.iter().any(|r| matches!(r, Ok(crate::LlmResponse::Done { .. }))));
170 let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
171 assert!(err.is_retryable(), "server_error must be retryable: {err:?}");
172 let provider_error = err.provider().expect("expected provider error");
173 assert_eq!(provider_error.kind, crate::ProviderErrorKind::Server);
174 assert_eq!(provider_error.http_status, Some(200));
175 assert_eq!(provider_error.request_id.as_deref(), Some("req-openai-1"));
176 assert_eq!(provider_error.code.as_deref(), Some("server_error"));
177 }
178
179 #[tokio::test]
180 async fn stream_response_surfaces_a_mapping_failure_as_the_only_item() {
181 let connection = ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() };
182 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap();
183 let context = Context::new(
184 vec![ChatMessage::User {
185 message_id: MessageId::new(),
186 content: vec![ContentBlock::Audio { data: "YXVkaW8=".to_string(), mime_type: "audio/wav".to_string() }],
187 timestamp: crate::types::IsoString::now(),
188 }],
189 vec![],
190 );
191
192 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
193
194 assert_eq!(responses.len(), 1);
195 assert!(matches!(responses[0], Err(LlmError::UnsupportedContent(_))), "{responses:?}");
196 }
197
198 #[test]
199 fn test_provider_display_name() {
200 let config = AetherOpenAiConfig::new(OpenAIConfig::new().with_api_key("test"), ProviderAuthMode::Default);
201 let provider = OpenAiProvider { config, http: reqwest::Client::new(), model: "gpt-4.1".to_string() };
202 assert_eq!(provider.display_name(), "OpenAI (gpt-4.1)");
203 }
204}