Skip to main content

llm/providers/openai/
responses_provider.rs

1use 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}