Skip to main content

llm/providers/openrouter/
provider.rs

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