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, 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}