llm/providers/openrouter/
provider.rs1use super::types::OpenRouterChatRequest;
2use crate::provider::{get_context_window, validate_reasoning};
3use crate::provider_connection::DEFAULT_STREAM_IDLE_TIMEOUT;
4use crate::providers::http::{http_client, openai_client};
5use crate::providers::openai_compatible::{
6 AetherOpenAiConfig, build_chat_request, streaming::create_custom_stream_generic,
7};
8use crate::providers::response_stream::error_stream;
9use crate::{
10 Context, LlmError, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, ProviderFactory, Result,
11 StreamingModelProvider,
12};
13use async_openai::{Client, config::OpenAIConfig};
14use std::future::ready;
15use std::time::Duration;
16
17pub struct OpenRouterProvider {
18 client: Client<AetherOpenAiConfig>,
19 model: String,
20 idle_timeout: Duration,
21}
22
23impl OpenRouterProvider {
24 pub fn new(api_key: String, model: String) -> Self {
25 let config = openai_config(Some(api_key), ProviderConnectionConfig::default());
26 Self { client: openai_client(config, http_client()), model, idle_timeout: DEFAULT_STREAM_IDLE_TIMEOUT }
27 }
28
29 pub fn default(model: &str) -> Result<Self> {
30 let api_key = std::env::var("OPENROUTER_API_KEY")
31 .map_err(|_| LlmError::MissingApiKey("OPENROUTER_API_KEY".to_string()))?;
32
33 Ok(Self::new(api_key, model.to_string()))
34 }
35
36 fn try_stream_response(&self, context: &Context) -> Result<LlmResponseStream> {
37 validate_reasoning(context, self.model().as_ref())?;
38 let mut request = build_chat_request(&self.model, context, None)?;
39 request.prompt_cache_key = context.prompt_cache_key().map(String::from);
40 let request = OpenRouterChatRequest::from_compatible(request, context.session_affinity_key());
41 Ok(create_custom_stream_generic(&self.client, request, self.idle_timeout))
42 }
43}
44
45impl ProviderFactory for OpenRouterProvider {
46 async fn from_env() -> Result<Self> {
47 Self::from_env_with_connection(ProviderConnectionConfig::default()).await
48 }
49
50 fn from_env_with_connection(connection: ProviderConnectionConfig) -> impl Future<Output = Result<Self>> + Send {
51 ready(provider_from_connection(connection))
52 }
53
54 fn with_model(mut self, model: &str) -> Self {
55 self.model = model.to_string();
56 self
57 }
58}
59
60impl StreamingModelProvider for OpenRouterProvider {
61 fn model(&self) -> Option<crate::LlmModel> {
62 format!("openrouter:{}", self.model).parse().ok()
63 }
64
65 fn context_window(&self) -> Option<u32> {
66 get_context_window("openrouter", &self.model)
67 }
68
69 fn stream_response(&self, context: &Context) -> LlmResponseStream {
70 self.try_stream_response(context).unwrap_or_else(error_stream)
71 }
72
73 fn display_name(&self) -> String {
74 format!("OpenRouter ({})", self.model)
75 }
76}
77
78fn openai_config(api_key: Option<String>, connection: ProviderConnectionConfig) -> AetherOpenAiConfig {
79 let api_key = api_key.unwrap_or_default();
80 let api_base = connection.base_url.unwrap_or_else(|| "https://openrouter.ai/api/v1".to_string());
81 let config = OpenAIConfig::new().with_api_key(api_key).with_api_base(api_base);
82 AetherOpenAiConfig::new(config, connection.auth_mode)
83}
84
85fn provider_from_connection(connection: ProviderConnectionConfig) -> Result<OpenRouterProvider> {
86 let api_key = match connection.auth_mode {
87 ProviderAuthMode::Default => Some(
88 std::env::var("OPENROUTER_API_KEY")
89 .map_err(|_| LlmError::MissingApiKey("OPENROUTER_API_KEY".to_string()))?,
90 ),
91 ProviderAuthMode::None => None,
92 };
93 let idle_timeout = connection.idle_timeout;
94 let config = openai_config(api_key, connection);
95 let client = openai_client(config, http_client());
96
97 Ok(OpenRouterProvider { client, model: String::new(), idle_timeout })
98}
99
100#[cfg(test)]
101mod tests {
102 use futures::{StreamExt, stream};
103 use reqwest::{Body, Method};
104
105 use super::*;
106 use crate::testing::FakeHttpService;
107 use crate::{ChatMessage, LlmResponse, ProviderErrorKind};
108
109 #[tokio::test]
110 async fn disabled_uses_unified_reasoning_without_conflicting_effort() {
111 let service = FakeHttpService::default();
112 service.route(Method::POST, OPENROUTER_URL, || response(200, OPENROUTER_FIXTURE));
113 let model = crate::LlmModel::all()
114 .iter()
115 .find(|model| {
116 model.provider_enum() == crate::catalog::Provider::OpenRouter && model.supports_reasoning_off()
117 })
118 .unwrap();
119 let provider = provider_with_service(&service).with_model(&model.model_id());
120 let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
121 context.set_reasoning_effort(crate::ReasoningEffort::Disabled);
122 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
123 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
124 let body = request_bodies(&service).pop().unwrap();
125 assert_eq!(body["reasoning"], serde_json::json!({"effort": "none"}));
126 assert!(body.get("reasoning_effort").is_none());
127 }
128
129 #[tokio::test]
130 async fn unknown_model_rejects_disabled_without_outbound_request() {
131 let service = FakeHttpService::default();
132 let provider = provider_with_service(&service).with_model("unknown-model");
133 let mut context = Context::new(vec![], vec![]);
134 context.set_reasoning_effort(crate::ReasoningEffort::Disabled);
135 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
136 assert_eq!(responses.len(), 1);
137 assert!(matches!(&responses[0], Err(LlmError::ReasoningValidation(_))));
138 assert!(!responses[0].as_ref().unwrap_err().is_retryable());
139 assert_eq!(request_bodies(&service).len(), 0);
140 }
141
142 #[tokio::test]
143 async fn stream_response_propagates_prompt_cache_key_and_keeps_cache_control() {
144 let service = FakeHttpService::default();
145 service.route(Method::POST, OPENROUTER_URL, || response(200, OPENROUTER_FIXTURE));
146 let provider = provider_with_service(&service).with_model("anthropic/claude-haiku-4.5");
147 let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
148 context.set_prompt_cache_key(Some("prefix-abc".to_string()));
149 context.set_session_affinity_key(Some("conversation-abc".to_string()));
150 context.set_reasoning_effort(crate::ReasoningEffort::High);
151
152 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
153 let body = request_bodies(&service).pop().unwrap();
154
155 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
156 assert!(responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))));
157 assert_eq!(body["model"], "anthropic/claude-haiku-4.5");
158 assert_eq!(body["prompt_cache_key"], "prefix-abc");
159 assert_eq!(body["session_id"], "conversation-abc");
160 assert_eq!(body["reasoning_effort"], "high");
161 assert_eq!(body["cache_control"]["type"], "ephemeral");
162 assert_eq!(body["usage"]["include"], true);
163 }
164
165 #[tokio::test]
166 async fn openrouter_rate_limit_response_is_a_retryable_provider_error() {
167 let body = r#"{"error":{"message":"Rate limit exceeded: free-models-per-day. Add 10 credits to unlock 1000 free model requests per day","code":429}}"#;
168 let error = rejected_request(429, body).await;
169 let provider_error = error.provider().expect("a 429 must classify as a provider error");
170
171 assert_eq!(provider_error.kind, ProviderErrorKind::RateLimit);
172 assert_eq!(provider_error.http_status, Some(429));
173 assert_eq!(provider_error.code.as_deref(), Some("429"));
174 assert_eq!(provider_error.request_id.as_deref(), Some("request-123"));
175 assert!(provider_error.message.contains("free-models-per-day"));
176 assert!(error.is_retryable(), "429 must be retryable, got {error}");
177 }
178
179 #[tokio::test]
180 async fn string_error_codes_are_preserved() {
181 let error = rejected_request(429, r#"{"error":{"code":"rate_limit_exceeded","message":"slow down"}}"#).await;
182 let provider_error = error.provider().unwrap();
183
184 assert_eq!(provider_error.code.as_deref(), Some("rate_limit_exceeded"));
185 assert!(provider_error.message.contains("slow down"));
186 assert!(error.is_retryable());
187 }
188
189 #[tokio::test]
190 async fn http_status_is_preserved_when_error_body_read_fails() {
191 let service = FakeHttpService::default();
192 service.route(Method::POST, OPENROUTER_URL, || {
193 let body = stream::iter([Err::<String, _>(std::io::Error::other("connection reset"))]);
194 response(429, Body::wrap_stream(body))
195 });
196 let error = collect_rejection(&provider_with_service(&service)).await;
197 let provider_error = error.provider().unwrap();
198
199 assert_eq!(provider_error.http_status, Some(429));
200 assert_eq!(provider_error.request_id.as_deref(), Some("request-123"));
201 assert_eq!(provider_error.kind, ProviderErrorKind::RateLimit);
202 assert!(error.is_retryable());
203 }
204
205 const OPENROUTER_URL: &str = "https://openrouter.ai/api/v1/chat/completions";
206 const OPENROUTER_FIXTURE: &str = include_str!("../../../tests/fixtures/openrouter/01_minimal.sse");
207
208 fn provider_with_service(service: &FakeHttpService) -> OpenRouterProvider {
209 let config = openai_config(Some("test-key".into()), ProviderConnectionConfig::default());
210 let client = openai_client(config, service.clone());
211 OpenRouterProvider { client, model: "test-model".into(), idle_timeout: DEFAULT_STREAM_IDLE_TIMEOUT }
212 }
213
214 fn response(status: u16, body: impl Into<Body>) -> reqwest::Response {
215 http::Response::builder()
216 .status(status)
217 .header("content-type", "text/event-stream")
218 .header("x-request-id", "request-123")
219 .body(body.into())
220 .unwrap()
221 .into()
222 }
223
224 fn request_bodies(service: &FakeHttpService) -> Vec<serde_json::Value> {
225 service
226 .take_requests()
227 .iter()
228 .map(|request| serde_json::from_slice(request.body().unwrap().as_bytes().unwrap()).unwrap())
229 .collect()
230 }
231
232 async fn rejected_request(status: u16, body: &str) -> LlmError {
233 let service = FakeHttpService::default();
234 let body = body.to_string();
235 service.route(Method::POST, OPENROUTER_URL, move || response(status, body.clone()));
236 collect_rejection(&provider_with_service(&service)).await
237 }
238
239 async fn collect_rejection(provider: &OpenRouterProvider) -> LlmError {
240 let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
241 let mut events = provider.stream_response(&context).collect::<Vec<_>>().await;
242 assert_eq!(events.len(), 1, "a rejected request must yield exactly one error: {events:?}");
243 events.pop().unwrap().expect_err("request must fail")
244 }
245}