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