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, ReasoningEffort};
99 use tokio_stream::StreamExt;
100
101 #[tokio::test]
102 async fn stream_response_sends_max_effort_on_the_wire() {
103 let mut server = CaptureServer::start_responses().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 http_200_failed_server_error_is_retryable_with_request_id() {
126 use crate::providers::test_capture_server::ResponseSpec;
127 let spec = ResponseSpec::sse(include_str!("../../../tests/fixtures/openai_responses/04_failed_server.sse"))
128 .with_header("x-request-id", "req-openai-1");
129 let mut server = CaptureServer::start_with_spec(spec).await;
130 let connection = ProviderConnectionConfig {
131 base_url: Some(server.base_url.clone()),
132 auth_mode: ProviderAuthMode::None,
133 ..Default::default()
134 };
135 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap();
136 let context = Context::new(vec![ChatMessage::user("hi")], vec![]);
137
138 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
139 let _ = server.captured().await;
140
141 assert!(!responses.iter().any(|r| matches!(r, Ok(crate::LlmResponse::Done { .. }))));
142 let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
143 assert!(err.is_retryable(), "server_error must be retryable: {err:?}");
144 let provider_error = err.provider().expect("expected provider error");
145 assert_eq!(provider_error.kind, crate::ProviderErrorKind::Server);
146 assert_eq!(provider_error.http_status, Some(200));
147 assert_eq!(provider_error.request_id.as_deref(), Some("req-openai-1"));
148 assert_eq!(provider_error.code.as_deref(), Some("server_error"));
149 }
150
151 #[tokio::test]
152 async fn stream_response_surfaces_a_mapping_failure_as_the_only_item() {
153 let connection = ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() };
154 let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap();
155 let context = Context::new(
156 vec![ChatMessage::User {
157 content: vec![crate::ContentBlock::Audio {
158 data: "YXVkaW8=".to_string(),
159 mime_type: "audio/wav".to_string(),
160 }],
161 timestamp: crate::types::IsoString::now(),
162 }],
163 vec![],
164 );
165
166 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
167
168 assert_eq!(responses.len(), 1);
169 assert!(matches!(responses[0], Err(LlmError::UnsupportedContent(_))), "{responses:?}");
170 }
171
172 #[test]
173 fn test_provider_display_name() {
174 let config = AetherOpenAiConfig::new(OpenAIConfig::new().with_api_key("test"), ProviderAuthMode::Default);
175 let provider = OpenAiProvider { config, http: reqwest::Client::new(), model: "gpt-4.1".to_string() };
176 assert_eq!(provider.display_name(), "OpenAI (gpt-4.1)");
177 }
178}