Skip to main content

rig_core/completion/
provider_options.rs

1//! Typed per-provider request options and reply extras. A provider's
2//! [`ProviderExtension`] marker names its key, its serialize-only `Options`
3//! and its `Extras` view of a reply. [`ProviderOptions`] holds entries for
4//! several providers at once, and each wire reads only the entry named by
5//! its own provider.
6//!
7//! An entry is an object of sections: [`SHARED`] (`"*"`) for fields every
8//! route of the provider spells the same, and one section per route, named
9//! by its API (`"openai.chat"`, `"openai.responses"`). A wire merges `"*"`,
10//! then the section of the route it encodes, at the top level of the body,
11//! above the mapped generation options and below `additional_params`. A
12//! section for another route is skipped.
13//!
14//! Each `Options` type names its provider ([`ExtensionOptions::Ext`]), so
15//! [`ProviderOptions::set`] and `CompletionRequest::provider_option` take
16//! the options by value, with no provider type and no `?`. Options that do
17//! not serialize fail the request's encode.
18//!
19//! ```
20//! use rig_core::completion::{CompletionRequest, ProviderOptions};
21//! use rig_core::providers::openrouter::extension::{OpenRouterExt, OpenRouterOptions};
22//!
23//! let request =
24//!     CompletionRequest::new("hi").provider_option(OpenRouterOptions::new().session_id("s-1"));
25//! assert!(request.provider_options.contains::<OpenRouterExt>());
26//! ```
27
28use std::collections::BTreeMap;
29use std::fmt;
30use std::panic::{RefUnwindSafe, UnwindSafe};
31use std::sync::Arc;
32
33use serde::{Deserialize, Deserializer, Serialize, Serializer};
34use serde_json::{Map, Value};
35
36use crate::completion::{CompletionRequest, ReplayTarget};
37use crate::message::Api;
38
39/// The section every route of a provider reads.
40pub const SHARED: &str = "*";
41
42/// A provider's typed request options and reply extras.
43pub trait ProviderExtension {
44    /// The provider's key: the name its wires report as
45    /// [`ReplayTarget::provider`] and stamp on a reply's `Origin`.
46    ///
47    /// Stable API from 0.44: requests store provider options under this
48    /// key, and the model catalog files the provider's models under it, so
49    /// rig does not rename a provider's key in a minor release.
50    const PROVIDER: &'static str;
51    /// The request options: an object of sections, [`SHARED`] and route API
52    /// names, each an object of body keys.
53    type Options: ExtensionOptions;
54    /// The typed view of a reply's `raw` document.
55    type Extras: ReplyExtras;
56}
57
58/// A provider's request options. Serialize-only: the body keys they write
59/// are their only output. They are plain data, `Send + Sync` on every
60/// target, since a request carrying them is a component of an ECS world,
61/// and unwind safe, so a request holding them stays unwind safe.
62pub trait ExtensionOptions:
63    Serialize + Clone + fmt::Debug + Send + Sync + UnwindSafe + RefUnwindSafe + 'static
64{
65    /// The provider these options are for: the entry
66    /// [`ProviderOptions::set`] stores them under. A type that serves as the
67    /// `Options` of several providers names one of them here, and is stored
68    /// for the others with [`ProviderOptions::with`].
69    type Ext: ProviderExtension<Options = Self>;
70
71    /// The fields `target` cannot send for `request`, each a top-level body
72    /// key with the reason. Each one set is reported through the request's
73    /// [`OnUnsupported`](crate::completion::OnUnsupported) policy under the
74    /// name `"<provider>.<section>.<field>"`. It must not call
75    /// [`options::param`](crate::completion::options::param), which reads
76    /// it. By default every field is sent.
77    fn unsupported(
78        &self,
79        target: &dyn ReplayTarget,
80        request: &CompletionRequest,
81    ) -> Vec<(&'static str, String)> {
82        let _ = (target, request);
83        Vec::new()
84    }
85}
86
87/// A typed view of a reply's provider document.
88pub trait ReplyExtras: Sized {
89    /// Read the view from `raw`, the reply's document on `api`.
90    ///
91    /// # Errors
92    ///
93    /// When `raw` does not hold the view's shape.
94    fn from_reply(api: &Api, raw: &Value) -> Result<Self, serde_json::Error>;
95}
96
97/// The value at `pointer` in a reply document, read as `T`: `None` when it
98/// is absent, `null` or an empty list, which a unary body may state where
99/// the document a stream rebuilds omits it. The one reader every `Extras`
100/// type's `from_reply` goes through.
101///
102/// # Errors
103///
104/// When the value is present but is not a `T`.
105pub(crate) fn reply_field<T: serde::de::DeserializeOwned>(
106    raw: &Value,
107    pointer: &str,
108) -> Result<Option<T>, serde_json::Error> {
109    match raw.pointer(pointer) {
110        None | Some(Value::Null) => Ok(None),
111        Some(Value::Array(items)) if items.is_empty() => Ok(None),
112        Some(value) => T::deserialize(value).map(Some),
113    }
114}
115
116/// Why a provider's options cannot be stored.
117#[non_exhaustive]
118#[derive(Debug, thiserror::Error)]
119pub enum OptionsError {
120    /// The options failed to serialize.
121    #[error("{provider} options do not serialize: {source}")]
122    Serialize {
123        /// The provider key.
124        provider: &'static str,
125        /// The serializer's error.
126        #[source]
127        source: serde_json::Error,
128    },
129    /// The options are not an object whose every value is an object.
130    #[error("{provider} options must serialize to an object of sections, each an object")]
131    NotSections {
132        /// The provider key.
133        provider: &'static str,
134    },
135}
136
137/// The refusal check of the typed options an entry was made from.
138trait Refusals: fmt::Debug + Send + Sync + UnwindSafe + RefUnwindSafe {
139    fn refusals(
140        &self,
141        target: &dyn ReplayTarget,
142        request: &CompletionRequest,
143    ) -> Vec<(&'static str, String)>;
144}
145
146impl<T: ExtensionOptions> Refusals for T {
147    fn refusals(
148        &self,
149        target: &dyn ReplayTarget,
150        request: &CompletionRequest,
151    ) -> Vec<(&'static str, String)> {
152        self.unsupported(target, request)
153    }
154}
155
156/// One provider's entry: its sections and the typed options they came
157/// from, or the error the typed options failed to serialize with.
158#[derive(Clone)]
159enum Entry {
160    Sections {
161        sections: Map<String, Value>,
162        typed: Option<Arc<dyn Refusals>>,
163    },
164    Failed(Failure),
165}
166
167/// The error of options that did not serialize. It is never mutated once
168/// made, so a request holding it stays unwind safe although
169/// `serde_json::Error` is not.
170#[derive(Clone)]
171struct Failure(Arc<OptionsError>);
172
173impl UnwindSafe for Failure {}
174impl RefUnwindSafe for Failure {}
175
176impl Entry {
177    /// The sections, when the options serialized.
178    fn sections(&self) -> Option<&Map<String, Value>> {
179        match self {
180            Self::Sections { sections, .. } => Some(sections),
181            Self::Failed(_) => None,
182        }
183    }
184}
185
186/// `P`'s `options` as sections, empty when they write no field.
187fn sections_of<P: ProviderExtension>(
188    options: &P::Options,
189) -> Result<Map<String, Value>, OptionsError> {
190    let provider = P::PROVIDER;
191    let value = serde_json::to_value(options)
192        .map_err(|source| OptionsError::Serialize { provider, source })?;
193    let Value::Object(sections) = value else {
194        return Err(OptionsError::NotSections { provider });
195    };
196    non_empty_sections(sections).ok_or(OptionsError::NotSections { provider })
197}
198
199/// Typed options for several providers, one entry per provider key. Each
200/// wire reads only the entry named by its own provider. Equality and the
201/// serialized form are the entries' sections. A deserialized entry has no
202/// typed options behind it, so its fields are sent as written: inserting the
203/// typed options again restores their refusal check.
204///
205/// An entry [`Self::set`] stores from options that do not serialize holds
206/// the [`OptionsError`] instead of sections. It fails the encode of every
207/// request that carries it, whichever provider the request goes to, and it
208/// fails serializing this value, so the error is never dropped. Two failed
209/// entries are equal when their errors read the same.
210#[derive(Clone, Default)]
211pub struct ProviderOptions(BTreeMap<String, Entry>);
212
213impl ProviderOptions {
214    /// No entry.
215    pub fn new() -> Self {
216        Self::default()
217    }
218
219    /// Store `options` as `P`'s entry, replacing any entry it had. Options
220    /// that write no field leave `P` with no entry.
221    ///
222    /// # Errors
223    ///
224    /// When `options` does not serialize to an object of object sections.
225    pub fn insert<P: ProviderExtension>(
226        &mut self,
227        options: &P::Options,
228    ) -> Result<&mut Self, OptionsError> {
229        let sections = sections_of::<P>(options)?;
230        self.put(P::PROVIDER, sections, || Arc::new(options.clone()));
231        Ok(self)
232    }
233
234    /// `self` with `options` as the entry of their provider,
235    /// [`ExtensionOptions::Ext`], replacing any entry it had. Options that
236    /// write no field leave that provider with no entry.
237    ///
238    /// It cannot fail: options that do not serialize to an object of object
239    /// sections are stored as a failed entry, which fails the encode of a
240    /// request that carries it with the [`OptionsError`] as the source.
241    /// [`Self::with`] reports the same error at once.
242    ///
243    /// The entry is always stored under `O::Ext`'s key, the built-in
244    /// provider the options type belongs to. A third-party provider whose
245    /// extension reuses a built-in options type (say `OpenAiOptions` for an
246    /// OpenAI-compatible gateway) must store them with
247    /// [`Self::with::<P>`](Self::with) instead, or its wire never reads them.
248    ///
249    /// ```
250    /// use rig_core::completion::ProviderOptions;
251    /// use rig_core::providers::openrouter::extension::{OpenRouterExt, OpenRouterOptions};
252    ///
253    /// let options = ProviderOptions::new().set(OpenRouterOptions::new().session_id("s-1"));
254    /// assert!(options.contains::<OpenRouterExt>());
255    /// ```
256    pub fn set<O: ExtensionOptions>(mut self, options: O) -> Self {
257        let provider = <O::Ext as ProviderExtension>::PROVIDER;
258        match sections_of::<O::Ext>(&options) {
259            Ok(sections) => self.put(provider, sections, || Arc::new(options)),
260            Err(error) => {
261                self.0
262                    .insert(provider.to_owned(), Entry::Failed(Failure(Arc::new(error))));
263            }
264        }
265        self
266    }
267
268    /// Store `sections` as `provider`'s entry, with the typed options
269    /// `typed` makes, or remove the entry when `sections` is empty.
270    fn put(
271        &mut self,
272        provider: &str,
273        sections: Map<String, Value>,
274        typed: impl FnOnce() -> Arc<dyn Refusals>,
275    ) {
276        if sections.is_empty() {
277            self.0.remove(provider);
278        } else {
279            self.0.insert(
280                provider.to_owned(),
281                Entry::Sections {
282                    sections,
283                    typed: Some(typed()),
284                },
285            );
286        }
287    }
288
289    /// `self` with `options` as `P`'s entry. See [`Self::insert`].
290    ///
291    /// # Errors
292    ///
293    /// As [`Self::insert`].
294    pub fn with<P: ProviderExtension>(
295        mut self,
296        options: &P::Options,
297    ) -> Result<Self, OptionsError> {
298        self.insert::<P>(options)?;
299        Ok(self)
300    }
301
302    /// `P`'s sections as the wire receives them, when it has an entry.
303    /// `None` for a failed entry ([`Self::set`]).
304    pub fn get<P: ProviderExtension>(&self) -> Option<&Map<String, Value>> {
305        self.0.get(P::PROVIDER).and_then(Entry::sections)
306    }
307
308    /// Remove `P`'s entry.
309    pub fn remove<P: ProviderExtension>(&mut self) {
310        self.0.remove(P::PROVIDER);
311    }
312
313    /// Whether `P` has an entry.
314    pub fn contains<P: ProviderExtension>(&self) -> bool {
315        self.0.contains_key(P::PROVIDER)
316    }
317
318    /// Whether no provider has an entry.
319    pub fn is_empty(&self) -> bool {
320        self.0.is_empty()
321    }
322
323    /// `self` with each entry of `over` in place of its own for the same
324    /// provider. An agent's options overlaid with a run's give the run's
325    /// entry for every provider the run names.
326    pub fn overlay(mut self, over: &ProviderOptions) -> ProviderOptions {
327        for (provider, entry) in &over.0 {
328            self.0.insert(provider.clone(), entry.clone());
329        }
330        self
331    }
332
333    /// The sections of `provider`'s entry.
334    pub(crate) fn sections(&self, provider: &str) -> Option<&Map<String, Value>> {
335        self.0.get(provider).and_then(Entry::sections)
336    }
337
338    /// The error of the first entry whose options did not serialize.
339    pub(crate) fn failure(&self) -> Option<&Arc<OptionsError>> {
340        self.0.values().find_map(|entry| match entry {
341            Entry::Failed(Failure(error)) => Some(error),
342            Entry::Sections { .. } => None,
343        })
344    }
345
346    /// The refusals of the typed options behind `provider`'s entry.
347    pub(crate) fn refusals(
348        &self,
349        provider: &str,
350        target: &dyn ReplayTarget,
351        request: &CompletionRequest,
352    ) -> Vec<(&'static str, String)> {
353        match self.0.get(provider) {
354            Some(Entry::Sections {
355                typed: Some(typed), ..
356            }) => typed.refusals(target, request),
357            _ => Vec::new(),
358        }
359    }
360
361    /// Remove `field` from the `sections` of `provider`'s entry.
362    pub(crate) fn remove_field(&mut self, provider: &str, sections: &[&str], field: &str) {
363        if let Some(Entry::Sections {
364            sections: entry, ..
365        }) = self.0.get_mut(provider)
366        {
367            for section in sections {
368                if let Some(Value::Object(fields)) = entry.get_mut(*section) {
369                    fields.shift_remove(field);
370                }
371            }
372        }
373    }
374}
375
376/// `sections` without its empty sections, or `None` when one is not an
377/// object.
378fn non_empty_sections(sections: Map<String, Value>) -> Option<Map<String, Value>> {
379    let mut kept = Map::new();
380    for (name, section) in sections {
381        match section {
382            Value::Object(fields) if fields.is_empty() => {}
383            Value::Object(fields) => {
384                kept.insert(name, Value::Object(fields));
385            }
386            _ => return None,
387        }
388    }
389    Some(kept)
390}
391
392impl fmt::Debug for Entry {
393    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
394        match self {
395            Self::Sections { sections, .. } => sections.fmt(f),
396            Self::Failed(Failure(error)) => {
397                f.debug_tuple("Failed").field(&error.to_string()).finish()
398            }
399        }
400    }
401}
402
403impl fmt::Debug for ProviderOptions {
404    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
405        f.debug_map().entries(&self.0).finish()
406    }
407}
408
409impl PartialEq for Entry {
410    fn eq(&self, other: &Self) -> bool {
411        match (self, other) {
412            (Self::Sections { sections: a, .. }, Self::Sections { sections: b, .. }) => a == b,
413            (Self::Failed(Failure(a)), Self::Failed(Failure(b))) => a.to_string() == b.to_string(),
414            _ => false,
415        }
416    }
417}
418
419impl PartialEq for ProviderOptions {
420    fn eq(&self, other: &Self) -> bool {
421        self.0 == other.0
422    }
423}
424
425impl Serialize for ProviderOptions {
426    fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
427        if let Some(error) = self.failure() {
428            return Err(serde::ser::Error::custom(error));
429        }
430        serializer.collect_map(
431            self.0
432                .iter()
433                .filter_map(|(provider, entry)| Some((provider, entry.sections()?))),
434        )
435    }
436}
437
438impl<'de> Deserialize<'de> for ProviderOptions {
439    fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
440        let entries = BTreeMap::<String, Map<String, Value>>::deserialize(deserializer)?;
441        let mut options = BTreeMap::new();
442        for (provider, sections) in entries {
443            let sections = non_empty_sections(sections).ok_or_else(|| {
444                serde::de::Error::custom(format!(
445                    "{provider} options must be an object of sections, each an object"
446                ))
447            })?;
448            if !sections.is_empty() {
449                options.insert(
450                    provider,
451                    Entry::Sections {
452                        sections,
453                        typed: None,
454                    },
455                );
456            }
457        }
458        Ok(Self(options))
459    }
460}
461
462#[cfg(test)]
463mod tests;