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;