Skip to main content

rig_core/providers/openrouter/
extension.rs

1//! OpenRouter's typed request options and reply extras. Every option is
2//! in the shared section, so one entry serves the Chat and the Responses
3//! route alike.
4//!
5//! ```
6//! use rig_core::completion::CompletionRequest;
7//! use rig_core::providers::openrouter::extension::{
8//!     ModelFallbacks, OpenRouterOptions, ProviderPreferences,
9//! };
10//!
11//! # fn run() -> Result<(), Box<dyn std::error::Error>> {
12//! let options = OpenRouterOptions::new()
13//!     .provider(ProviderPreferences::new().only(["anthropic"]).zdr(true))
14//!     .models(ModelFallbacks::new(["openai/gpt-4o-mini"])?);
15//! let request = CompletionRequest::new("hi").provider_option(options);
16//! # let _ = request;
17//! # Ok(())
18//! # }
19//! ```
20
21use std::collections::BTreeMap;
22
23use serde::Serialize;
24use serde_json::{Map, Value};
25
26use crate::completion::provider_options::reply_field;
27use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
28use crate::message::Api;
29
30/// OpenRouter's extension marker.
31#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
32pub struct OpenRouterExt;
33
34impl ProviderExtension for OpenRouterExt {
35    const PROVIDER: &'static str = super::PROVIDER_NAME;
36    type Options = OpenRouterOptions;
37    type Extras = OpenRouterExtras;
38}
39
40/// OpenRouter's request options, all sent on every route.
41#[non_exhaustive]
42#[derive(Clone, Debug, Default, PartialEq, Serialize)]
43pub struct OpenRouterOptions {
44    /// The fields every route takes.
45    #[serde(rename = "*")]
46    pub shared: OpenRouterShared,
47}
48
49/// The fields OpenRouter takes on every route.
50#[non_exhaustive]
51#[derive(Clone, Debug, Default, PartialEq, Serialize)]
52pub struct OpenRouterShared {
53    /// Which upstream providers may serve the request, and in what order.
54    #[serde(skip_serializing_if = "Option::is_none")]
55    pub provider: Option<ProviderPreferences>,
56    /// The models tried after the request's own, in order.
57    #[serde(skip_serializing_if = "Option::is_none")]
58    pub models: Option<ModelFallbacks>,
59    /// Plugins such as `{"id": "web"}`.
60    #[serde(skip_serializing_if = "Vec::is_empty")]
61    pub plugins: Vec<Value>,
62    /// Groups requests into one session in OpenRouter's logs.
63    #[serde(skip_serializing_if = "Option::is_none")]
64    pub session_id: Option<String>,
65    /// String pairs attached to the request.
66    #[serde(skip_serializing_if = "BTreeMap::is_empty")]
67    pub metadata: BTreeMap<String, String>,
68    /// Reasoning fields sent beside the mapped effort.
69    #[serde(skip_serializing_if = "ReasoningExtra::is_empty")]
70    pub reasoning: ReasoningExtra,
71    /// Sample from the `k` most likely tokens.
72    #[serde(skip_serializing_if = "Option::is_none")]
73    pub top_k: Option<u32>,
74    /// The minimum probability of a token, relative to the most likely one.
75    #[serde(skip_serializing_if = "Option::is_none")]
76    pub min_p: Option<f64>,
77    /// Top-a sampling.
78    #[serde(skip_serializing_if = "Option::is_none")]
79    pub top_a: Option<f64>,
80    /// Penalizes tokens already in the prompt and the output.
81    #[serde(skip_serializing_if = "Option::is_none")]
82    pub repetition_penalty: Option<f64>,
83    /// A stable id of the end user.
84    #[serde(skip_serializing_if = "Option::is_none")]
85    pub user: Option<String>,
86}
87
88/// The `reasoning` fields OpenRouter takes beside the effort or budget the
89/// generation options map.
90#[non_exhaustive]
91#[derive(Clone, Debug, Default, PartialEq, Serialize)]
92pub struct ReasoningExtra {
93    /// Reason, but leave the reasoning out of the reply.
94    #[serde(skip_serializing_if = "Option::is_none")]
95    pub exclude: Option<bool>,
96    /// How much of the reasoning the reply summarizes.
97    #[serde(skip_serializing_if = "Option::is_none")]
98    pub summary: Option<ReasoningSummary>,
99}
100
101impl ReasoningExtra {
102    fn is_empty(&self) -> bool {
103        self.exclude.is_none() && self.summary.is_none()
104    }
105}
106
107/// How much of the reasoning a reply summarizes.
108#[non_exhaustive]
109#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
110#[serde(rename_all = "lowercase")]
111pub enum ReasoningSummary {
112    /// The provider decides.
113    Auto,
114    /// A short summary.
115    Concise,
116    /// A detailed summary.
117    Detailed,
118}
119
120impl OpenRouterOptions {
121    /// No option set.
122    pub fn new() -> Self {
123        Self::default()
124    }
125
126    /// Route by `preferences`.
127    pub fn provider(mut self, preferences: ProviderPreferences) -> Self {
128        self.shared.provider = Some(preferences);
129        self
130    }
131
132    /// Fall back to `models`, in order, after the request's own model.
133    pub fn models(mut self, models: ModelFallbacks) -> Self {
134        self.shared.models = Some(models);
135        self
136    }
137
138    /// Add `plugin`, such as `{"id": "web"}`.
139    pub fn plugin(mut self, plugin: Value) -> Self {
140        self.shared.plugins.push(plugin);
141        self
142    }
143
144    /// Group the request under session `id`.
145    pub fn session_id(mut self, id: impl Into<String>) -> Self {
146        self.shared.session_id = Some(id.into());
147        self
148    }
149
150    /// Attach the metadata pair `key`, `value`.
151    pub fn metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
152        self.shared.metadata.insert(key.into(), value.into());
153        self
154    }
155
156    /// Reason but leave the reasoning out of the reply when `exclude`.
157    pub fn reasoning_exclude(mut self, exclude: bool) -> Self {
158        self.shared.reasoning.exclude = Some(exclude);
159        self
160    }
161
162    /// Summarize the reasoning at `summary`.
163    pub fn reasoning_summary(mut self, summary: ReasoningSummary) -> Self {
164        self.shared.reasoning.summary = Some(summary);
165        self
166    }
167
168    /// Sample from the `k` most likely tokens.
169    pub fn top_k(mut self, k: u32) -> Self {
170        self.shared.top_k = Some(k);
171        self
172    }
173
174    /// Set the minimum relative token probability.
175    pub fn min_p(mut self, p: f64) -> Self {
176        self.shared.min_p = Some(p);
177        self
178    }
179
180    /// Set top-a sampling.
181    pub fn top_a(mut self, a: f64) -> Self {
182        self.shared.top_a = Some(a);
183        self
184    }
185
186    /// Set the repetition penalty.
187    pub fn repetition_penalty(mut self, penalty: f64) -> Self {
188        self.shared.repetition_penalty = Some(penalty);
189        self
190    }
191
192    /// Name the end user.
193    pub fn user(mut self, user: impl Into<String>) -> Self {
194        self.shared.user = Some(user.into());
195        self
196    }
197}
198
199impl ExtensionOptions for OpenRouterOptions {
200    type Ext = OpenRouterExt;
201}
202
203/// A non-empty list of fallback models.
204#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
205#[serde(transparent)]
206pub struct ModelFallbacks(Vec<String>);
207
208/// A fallback list with no model.
209#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
210#[error("model fallbacks need at least one model")]
211pub struct EmptyFallbacks;
212
213impl ModelFallbacks {
214    /// The models tried, in order, after the request's own.
215    ///
216    /// # Errors
217    ///
218    /// When `models` is empty.
219    pub fn new(
220        models: impl IntoIterator<Item = impl Into<String>>,
221    ) -> Result<Self, EmptyFallbacks> {
222        let models: Vec<String> = models.into_iter().map(Into::into).collect();
223        if models.is_empty() {
224            return Err(EmptyFallbacks);
225        }
226        Ok(Self(models))
227    }
228
229    /// The fallback models, in order.
230    pub fn models(&self) -> &[String] {
231        &self.0
232    }
233}
234
235/// Whether providers that may store request data are eligible.
236#[non_exhaustive]
237#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
238#[serde(rename_all = "lowercase")]
239pub enum DataCollection {
240    /// Any provider.
241    Allow,
242    /// Only providers that store no request data beyond the request.
243    Deny,
244}
245
246/// A quantization an upstream provider serves a model at.
247#[non_exhaustive]
248#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
249#[serde(rename_all = "lowercase")]
250pub enum Quantization {
251    /// 4-bit integers.
252    Int4,
253    /// 8-bit integers.
254    Int8,
255    /// 16-bit floats.
256    Fp16,
257    /// Brain floats.
258    Bf16,
259    /// 32-bit floats.
260    Fp32,
261    /// 8-bit floats.
262    Fp8,
263    /// Not stated by the provider.
264    Unknown,
265}
266
267/// What providers are ordered by.
268#[non_exhaustive]
269#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
270#[serde(rename_all = "lowercase")]
271pub enum ProviderSortStrategy {
272    /// Cheapest first.
273    Price,
274    /// Highest throughput first.
275    Throughput,
276    /// Lowest latency first.
277    Latency,
278}
279
280/// How a request with fallback models sorts its providers.
281#[non_exhaustive]
282#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
283#[serde(rename_all = "lowercase")]
284pub enum SortPartition {
285    /// Within each model.
286    Model,
287    /// Across every model.
288    None,
289}
290
291/// A sort with a partition.
292#[non_exhaustive]
293#[derive(Clone, Debug, PartialEq, Serialize)]
294pub struct ProviderSortConfig {
295    /// What providers are ordered by.
296    pub by: ProviderSortStrategy,
297    /// How fallback models are grouped.
298    #[serde(skip_serializing_if = "Option::is_none")]
299    pub partition: Option<SortPartition>,
300}
301
302impl ProviderSortConfig {
303    /// Sort by `by`.
304    pub fn new(by: ProviderSortStrategy) -> Self {
305        Self {
306            by,
307            partition: None,
308        }
309    }
310
311    /// Group fallback models by `partition`.
312    pub fn partition(mut self, partition: SortPartition) -> Self {
313        self.partition = Some(partition);
314        self
315    }
316}
317
318/// A provider sort: a strategy, or a strategy with a partition.
319#[non_exhaustive]
320#[derive(Clone, Debug, PartialEq, Serialize)]
321#[serde(untagged)]
322pub enum ProviderSort {
323    /// A strategy alone.
324    Simple(ProviderSortStrategy),
325    /// A strategy and a partition.
326    Complex(ProviderSortConfig),
327}
328
329impl From<ProviderSortStrategy> for ProviderSort {
330    fn from(strategy: ProviderSortStrategy) -> Self {
331        Self::Simple(strategy)
332    }
333}
334
335impl From<ProviderSortConfig> for ProviderSort {
336    fn from(config: ProviderSortConfig) -> Self {
337        Self::Complex(config)
338    }
339}
340
341/// A throughput floor in tokens per second, alone or per percentile.
342/// Providers below it are tried later, not excluded.
343#[non_exhaustive]
344#[derive(Clone, Debug, PartialEq, Serialize)]
345#[serde(untagged)]
346pub enum ThroughputThreshold {
347    /// One floor.
348    Simple(f64),
349    /// A floor per percentile.
350    Percentile(PercentileThresholds),
351}
352
353/// A latency ceiling in seconds, alone or per percentile. Providers above
354/// it are tried later, not excluded.
355#[non_exhaustive]
356#[derive(Clone, Debug, PartialEq, Serialize)]
357#[serde(untagged)]
358pub enum LatencyThreshold {
359    /// One ceiling.
360    Simple(f64),
361    /// A ceiling per percentile.
362    Percentile(PercentileThresholds),
363}
364
365/// A threshold per percentile.
366#[non_exhaustive]
367#[derive(Clone, Debug, Default, PartialEq, Serialize)]
368pub struct PercentileThresholds {
369    /// The median.
370    #[serde(skip_serializing_if = "Option::is_none")]
371    pub p50: Option<f64>,
372    /// The 75th percentile.
373    #[serde(skip_serializing_if = "Option::is_none")]
374    pub p75: Option<f64>,
375    /// The 90th percentile.
376    #[serde(skip_serializing_if = "Option::is_none")]
377    pub p90: Option<f64>,
378    /// The 99th percentile.
379    #[serde(skip_serializing_if = "Option::is_none")]
380    pub p99: Option<f64>,
381}
382
383impl PercentileThresholds {
384    /// No threshold.
385    pub fn new() -> Self {
386        Self::default()
387    }
388
389    /// Set the median threshold.
390    pub fn p50(mut self, value: f64) -> Self {
391        self.p50 = Some(value);
392        self
393    }
394
395    /// Set the 75th-percentile threshold.
396    pub fn p75(mut self, value: f64) -> Self {
397        self.p75 = Some(value);
398        self
399    }
400
401    /// Set the 90th-percentile threshold.
402    pub fn p90(mut self, value: f64) -> Self {
403        self.p90 = Some(value);
404        self
405    }
406
407    /// Set the 99th-percentile threshold.
408    pub fn p99(mut self, value: f64) -> Self {
409        self.p99 = Some(value);
410        self
411    }
412}
413
414/// Price ceilings, in USD per million tokens or per item. A request no
415/// eligible provider can serve under them fails.
416#[non_exhaustive]
417#[derive(Clone, Debug, Default, PartialEq, Serialize)]
418pub struct MaxPrice {
419    /// Per million prompt tokens.
420    #[serde(skip_serializing_if = "Option::is_none")]
421    pub prompt: Option<f64>,
422    /// Per million completion tokens.
423    #[serde(skip_serializing_if = "Option::is_none")]
424    pub completion: Option<f64>,
425    /// Per request.
426    #[serde(skip_serializing_if = "Option::is_none")]
427    pub request: Option<f64>,
428    /// Per image.
429    #[serde(skip_serializing_if = "Option::is_none")]
430    pub image: Option<f64>,
431}
432
433impl MaxPrice {
434    /// No ceiling.
435    pub fn new() -> Self {
436        Self::default()
437    }
438
439    /// Cap the prompt price.
440    pub fn prompt(mut self, price: f64) -> Self {
441        self.prompt = Some(price);
442        self
443    }
444
445    /// Cap the completion price.
446    pub fn completion(mut self, price: f64) -> Self {
447        self.completion = Some(price);
448        self
449    }
450
451    /// Cap the price per request.
452    pub fn request(mut self, price: f64) -> Self {
453        self.request = Some(price);
454        self
455    }
456
457    /// Cap the price per image.
458    pub fn image(mut self, price: f64) -> Self {
459        self.image = Some(price);
460        self
461    }
462}
463
464/// Which upstream providers may serve a request and in what order
465/// (<https://openrouter.ai/docs/guides/routing/provider-selection>). Unset
466/// fields keep OpenRouter's defaults. Slugs and limits are sent as given.
467#[non_exhaustive]
468#[derive(Clone, Debug, Default, PartialEq, Serialize)]
469pub struct ProviderPreferences {
470    /// Provider slugs tried first, in order.
471    #[serde(skip_serializing_if = "Option::is_none")]
472    pub order: Option<Vec<String>>,
473    /// The only eligible provider slugs.
474    #[serde(skip_serializing_if = "Option::is_none")]
475    pub only: Option<Vec<String>>,
476    /// Provider slugs never used.
477    #[serde(skip_serializing_if = "Option::is_none")]
478    pub ignore: Option<Vec<String>>,
479    /// Whether providers outside `order` may serve the request.
480    #[serde(skip_serializing_if = "Option::is_none")]
481    pub allow_fallbacks: Option<bool>,
482    /// Whether only providers taking every request parameter are eligible.
483    #[serde(skip_serializing_if = "Option::is_none")]
484    pub require_parameters: Option<bool>,
485    /// Whether providers that store request data are eligible.
486    #[serde(skip_serializing_if = "Option::is_none")]
487    pub data_collection: Option<DataCollection>,
488    /// Whether only zero-data-retention endpoints are eligible.
489    #[serde(skip_serializing_if = "Option::is_none")]
490    pub zdr: Option<bool>,
491    /// The order providers are tried in, in place of load balancing.
492    #[serde(skip_serializing_if = "Option::is_none")]
493    pub sort: Option<ProviderSort>,
494    /// The throughput below which a provider is tried later.
495    #[serde(skip_serializing_if = "Option::is_none")]
496    pub preferred_min_throughput: Option<ThroughputThreshold>,
497    /// The latency above which a provider is tried later.
498    #[serde(skip_serializing_if = "Option::is_none")]
499    pub preferred_max_latency: Option<LatencyThreshold>,
500    /// The price ceilings.
501    #[serde(skip_serializing_if = "Option::is_none")]
502    pub max_price: Option<MaxPrice>,
503    /// The quantizations a provider may serve the model at.
504    #[serde(skip_serializing_if = "Option::is_none")]
505    pub quantizations: Option<Vec<Quantization>>,
506}
507
508fn strings(items: impl IntoIterator<Item = impl Into<String>>) -> Option<Vec<String>> {
509    Some(items.into_iter().map(Into::into).collect())
510}
511
512impl ProviderPreferences {
513    /// No preference.
514    pub fn new() -> Self {
515        Self::default()
516    }
517
518    /// Try `providers` first, in order.
519    pub fn order(mut self, providers: impl IntoIterator<Item = impl Into<String>>) -> Self {
520        self.order = strings(providers);
521        self
522    }
523
524    /// Make only `providers` eligible.
525    pub fn only(mut self, providers: impl IntoIterator<Item = impl Into<String>>) -> Self {
526        self.only = strings(providers);
527        self
528    }
529
530    /// Never use `providers`.
531    pub fn ignore(mut self, providers: impl IntoIterator<Item = impl Into<String>>) -> Self {
532        self.ignore = strings(providers);
533        self
534    }
535
536    /// Whether providers outside `order` may serve the request.
537    pub fn allow_fallbacks(mut self, allow: bool) -> Self {
538        self.allow_fallbacks = Some(allow);
539        self
540    }
541
542    /// Whether only providers taking every request parameter are eligible.
543    pub fn require_parameters(mut self, require: bool) -> Self {
544        self.require_parameters = Some(require);
545        self
546    }
547
548    /// Set the data collection policy.
549    pub fn data_collection(mut self, policy: DataCollection) -> Self {
550        self.data_collection = Some(policy);
551        self
552    }
553
554    /// Whether only zero-data-retention endpoints are eligible.
555    pub fn zdr(mut self, enable: bool) -> Self {
556        self.zdr = Some(enable);
557        self
558    }
559
560    /// Try providers in the order `sort` gives.
561    pub fn sort(mut self, sort: impl Into<ProviderSort>) -> Self {
562        self.sort = Some(sort.into());
563        self
564    }
565
566    /// Try providers below `threshold` later.
567    pub fn preferred_min_throughput(mut self, threshold: ThroughputThreshold) -> Self {
568        self.preferred_min_throughput = Some(threshold);
569        self
570    }
571
572    /// Try providers above `threshold` later.
573    pub fn preferred_max_latency(mut self, threshold: LatencyThreshold) -> Self {
574        self.preferred_max_latency = Some(threshold);
575        self
576    }
577
578    /// Fail the request when no provider is under `price`.
579    pub fn max_price(mut self, price: MaxPrice) -> Self {
580        self.max_price = Some(price);
581        self
582    }
583
584    /// Make only providers serving one of `quantizations` eligible.
585    pub fn quantizations(mut self, quantizations: impl IntoIterator<Item = Quantization>) -> Self {
586        self.quantizations = Some(quantizations.into_iter().collect());
587        self
588    }
589
590    /// Try the cheapest provider first.
591    pub fn cheapest(self) -> Self {
592        self.sort(ProviderSortStrategy::Price)
593    }
594}
595
596/// OpenRouter's reply fields. Each is `None` when the reply lacks it.
597#[non_exhaustive]
598#[derive(Clone, Debug, Default, PartialEq)]
599pub struct OpenRouterExtras {
600    /// The upstream provider that served the request. Chat; Responses when
601    /// sent.
602    pub provider: Option<String>,
603    /// The upstream provider's own finish reason. Chat only.
604    pub native_finish_reason: Option<String>,
605    /// The service tier that served the request. Both routes.
606    pub service_tier: Option<String>,
607    /// The upstream backend's fingerprint. Chat only.
608    pub system_fingerprint: Option<String>,
609    /// The request's cost in credits. Both routes.
610    pub cost: Option<f64>,
611    /// The cost broken down by upstream charge. Both routes, keyed as each
612    /// route spells them.
613    pub cost_details: Option<Map<String, Value>>,
614    /// Whether the request ran on the caller's own provider key. Both
615    /// routes.
616    pub is_byok: Option<bool>,
617    /// Prompt token details: `usage.prompt_tokens_details` on Chat,
618    /// `usage.input_tokens_details` on Responses.
619    pub prompt_tokens_details: Option<Value>,
620    /// Server tool usage, such as web searches. Chat only.
621    pub server_tool_use_details: Option<Value>,
622    /// OpenRouter's routing metadata. Chat only.
623    pub openrouter_metadata: Option<Value>,
624    /// The message's annotations, such as URL citations. Chat only.
625    pub annotations: Option<Vec<Value>>,
626}
627
628impl ReplyExtras for OpenRouterExtras {
629    fn from_reply(api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
630        if api.as_str() == "openai.responses" {
631            return Ok(Self {
632                provider: reply_field(raw, "/provider")?,
633                service_tier: reply_field(raw, "/service_tier")?,
634                cost: reply_field(raw, "/usage/cost")?,
635                cost_details: reply_field(raw, "/usage/cost_details")?,
636                is_byok: reply_field(raw, "/usage/is_byok")?,
637                prompt_tokens_details: reply_field(raw, "/usage/input_tokens_details")?,
638                ..Self::default()
639            });
640        }
641        Ok(Self {
642            provider: reply_field(raw, "/provider")?,
643            native_finish_reason: reply_field(raw, "/choices/0/native_finish_reason")?,
644            service_tier: reply_field(raw, "/service_tier")?,
645            system_fingerprint: reply_field(raw, "/system_fingerprint")?,
646            cost: reply_field(raw, "/usage/cost")?,
647            cost_details: reply_field(raw, "/usage/cost_details")?,
648            is_byok: reply_field(raw, "/usage/is_byok")?,
649            prompt_tokens_details: reply_field(raw, "/usage/prompt_tokens_details")?,
650            server_tool_use_details: reply_field(raw, "/usage/server_tool_use_details")?,
651            openrouter_metadata: reply_field(raw, "/openrouter_metadata")?,
652            annotations: reply_field(raw, "/choices/0/message/annotations")?,
653        })
654    }
655}
656
657#[cfg(test)]
658mod tests;