Skip to main content

llm/providers/openai/
responses_provider.rs

1use async_openai::Client;
2use async_openai::config::OpenAIConfig;
3use serde_json::Value;
4use tokio_stream::StreamExt;
5use tracing::debug;
6
7use crate::provider::{error_stream, get_context_window, stream_from};
8use crate::providers::openai_compatible::AetherOpenAiConfig;
9use crate::providers::openai_responses::mappers::{ResponsesRequestPolicy, build_wire_request};
10use crate::providers::openai_responses::streaming::{ResponsesStreamEvent, process_response_stream};
11use crate::{
12    Context, LlmError, LlmModel, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, ProviderFactory,
13    Result, StreamingModelProvider,
14};
15
16pub struct OpenAiProvider {
17    client: Client<AetherOpenAiConfig>,
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    async fn from_env_with_connection(connection: ProviderConnectionConfig) -> Result<Self> {
27        let api_key = match connection.auth_mode {
28            ProviderAuthMode::Default => {
29                std::env::var("OPENAI_API_KEY").map_err(|_| LlmError::MissingApiKey("OPENAI_API_KEY".to_string()))?
30            }
31            ProviderAuthMode::None => String::new(),
32        };
33
34        let mut config = OpenAIConfig::new().with_api_key(api_key);
35        if let Some(base_url) = connection.base_url {
36            config = config.with_api_base(base_url);
37        }
38        let config = AetherOpenAiConfig::new(config, connection.auth_mode);
39
40        Ok(Self { client: Client::with_config(config), model: "gpt-4.1".to_string() })
41    }
42
43    fn with_model(mut self, model: &str) -> Self {
44        if !model.is_empty() {
45            self.model = model.to_string();
46        }
47        self
48    }
49}
50
51impl StreamingModelProvider for OpenAiProvider {
52    fn stream_response(&self, context: &Context) -> LlmResponseStream {
53        let client = self.client.clone();
54        let model = self.model.clone();
55        let request = match build_wire_request(&model, context, &ResponsesRequestPolicy::openai()) {
56            Ok(request) => request,
57            Err(e) => return error_stream(e),
58        };
59
60        stream_from(
61            async move {
62                debug!("Starting OpenAI Responses API stream for model: {model}");
63                client
64                    .responses()
65                    .create_stream_byot::<Value, ResponsesStreamEvent>(request)
66                    .await
67                    .map_err(|e| LlmError::ApiRequest(e.to_string()))
68            },
69            |stream| {
70                process_response_stream(Box::pin(
71                    stream.map(|result| result.map_err(|e| LlmError::StreamInterrupted(e.to_string()))),
72                ))
73            },
74        )
75    }
76
77    fn display_name(&self) -> String {
78        format!("OpenAI ({})", self.model)
79    }
80
81    fn context_window(&self) -> Option<u32> {
82        get_context_window("openai", &self.model)
83    }
84
85    fn model(&self) -> Option<LlmModel> {
86        format!("openai:{}", self.model).parse().ok()
87    }
88}
89
90#[cfg(test)]
91mod tests {
92    use super::*;
93    use crate::providers::test_capture_server::CaptureServer;
94    use crate::{ChatMessage, ReasoningEffort};
95
96    #[tokio::test]
97    async fn stream_response_sends_max_effort_on_the_wire() {
98        let mut server = CaptureServer::start().await;
99        let connection = ProviderConnectionConfig {
100            base_url: Some(server.base_url.clone()),
101            auth_mode: ProviderAuthMode::None,
102            ..Default::default()
103        };
104        let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap().with_model("gpt-5.6");
105        let mut context = Context::new(vec![ChatMessage::user("Think harder")], vec![]);
106        context.set_reasoning_effort(Some(ReasoningEffort::Max));
107        context.set_prompt_cache_key(Some("cache-key".to_string()));
108
109        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
110        let captured = server.captured().await;
111
112        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
113        assert_eq!(captured.body["reasoning"]["effort"], "max");
114        assert_eq!(captured.body["model"], "gpt-5.6");
115        assert_eq!(captured.body["prompt_cache_key"], "cache-key");
116        assert_eq!(captured.body["stream"], true);
117    }
118
119    #[tokio::test]
120    async fn stream_response_surfaces_a_mapping_failure_as_the_only_item() {
121        let connection = ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() };
122        let provider = OpenAiProvider::from_env_with_connection(connection).await.unwrap();
123        let context = Context::new(
124            vec![ChatMessage::User {
125                content: vec![crate::ContentBlock::Audio {
126                    data: "YXVkaW8=".to_string(),
127                    mime_type: "audio/wav".to_string(),
128                }],
129                timestamp: crate::types::IsoString::now(),
130            }],
131            vec![],
132        );
133
134        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
135
136        assert_eq!(responses.len(), 1);
137        assert!(matches!(responses[0], Err(LlmError::UnsupportedContent(_))), "{responses:?}");
138    }
139
140    #[test]
141    fn test_provider_display_name() {
142        let config = AetherOpenAiConfig::new(OpenAIConfig::new().with_api_key("test"), ProviderAuthMode::Default);
143        let provider = OpenAiProvider { client: Client::with_config(config), model: "gpt-4.1".to_string() };
144        assert_eq!(provider.display_name(), "OpenAI (gpt-4.1)");
145    }
146}