Skip to main content

llm/providers/openai/
responses_provider.rs

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