llm/providers/openai/
responses_provider.rs1use 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}