Skip to main content

tenferro_ad/
context.rs

1//! Explicit ownership for automatic-differentiation rule sets.
2
3use std::sync::Arc;
4
5use tenferro_runtime::program::FrozenProgram;
6use tenferro_runtime::{CacheStats, Result, TracedTensor};
7
8// SemanticCompatDispatcher removed in Unification 7.
9// Extension AD is handled exclusively by SemanticExtensionRuleSet.
10use crate::semantic_extension::{SemanticExtensionRegistryError, SemanticExtensionRuleSet};
11use crate::semantic_transform::{
12    semantic_jvp, semantic_vjp, SemanticAdProgram, SemanticAdTransformError,
13};
14use crate::transform_cache::{
15    AdTransformCache, AdTransformCacheLimits, SemanticAdTransformCacheKey,
16};
17
18/// Stats for caches owned by an [`AdContext`].
19///
20/// `retained_bytes` fields are logical payload estimates, not process RSS.
21///
22/// # Examples
23///
24/// ```rust
25/// use tenferro_ad::AdContext;
26///
27/// let ad = AdContext::builder().build().unwrap();
28/// assert_eq!(ad.cache_stats().unwrap().ad_transforms.entries, 0);
29/// ```
30#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
31pub struct AdContextCacheStats {
32    /// AD transform graph memoization cache.
33    pub ad_transforms: CacheStats,
34}
35
36/// Explicit automatic-differentiation context.
37///
38/// `AdContext` owns the extension AD rules used by traced AD transforms.
39/// It also owns the AD transform cache shared by context-driven traced AD and
40/// eager runtimes created from this context.
41///
42/// # Examples
43///
44/// ```rust
45/// use tenferro_ad::AdContext;
46///
47/// let ad = AdContext::builder().build().unwrap();
48/// assert!(ad
49///     .semantic_extension_rules()
50///     .lookup_linearize("example.missing.v1")
51///     .is_none());
52/// ```
53#[derive(Clone, Debug)]
54pub struct AdContext {
55    semantic_extension_rules: SemanticExtensionRuleSet,
56    ad_transform_cache: Arc<AdTransformCache>,
57}
58
59impl AdContext {
60    /// Start building an explicit AD context.
61    ///
62    /// # Examples
63    ///
64    /// ```rust
65    /// use tenferro_ad::AdContext;
66    ///
67    /// let _builder = AdContext::builder();
68    /// ```
69    pub fn builder() -> AdContextBuilder {
70        AdContextBuilder::default()
71    }
72
73    pub(crate) fn with_rules_and_transform_cache(
74        semantic_extension_rules: SemanticExtensionRuleSet,
75        ad_transform_cache: Arc<AdTransformCache>,
76    ) -> Self {
77        Self {
78            semantic_extension_rules,
79            ad_transform_cache,
80        }
81    }
82
83    /// Return semantic-program extension AD rules owned by this context.
84    ///
85    /// # Examples
86    ///
87    /// ```rust
88    /// use tenferro_ad::AdContext;
89    ///
90    /// let ad = AdContext::builder().build().unwrap();
91    /// assert!(ad
92    ///     .semantic_extension_rules()
93    ///     .lookup_linearize("example.missing.v1")
94    ///     .is_none());
95    /// ```
96    pub fn semantic_extension_rules(&self) -> &SemanticExtensionRuleSet {
97        &self.semantic_extension_rules
98    }
99
100    /// Transform a frozen semantic program into its forward-mode derivative.
101    ///
102    /// `active_inputs` follows source-program input order. Active tangent
103    /// seeds are appended after all primal inputs.
104    ///
105    /// # Errors
106    ///
107    /// Returns [`SemanticAdTransformError::ActivityArity`] when
108    /// `active_inputs` has the wrong length,
109    /// [`SemanticAdTransformError::Extension`] when an extension rule rejects
110    /// the transform, or the corresponding `Query`, `Build`, `Finish`, or
111    /// `Cache` variant when program import, construction, finalization, or
112    /// cache access fails.
113    pub fn jvp_program(
114        &self,
115        input: &FrozenProgram,
116        active_inputs: &[bool],
117    ) -> std::result::Result<SemanticAdProgram, SemanticAdTransformError> {
118        let key = SemanticAdTransformCacheKey::jvp(input, active_inputs);
119        if let Some(cached) = self
120            .ad_transform_cache
121            .get_semantic(&key, input)
122            .map_err(SemanticAdTransformError::Cache)?
123        {
124            return cached
125                .as_ref()
126                .with_input_prefix_bindings_from(input)
127                .map_err(SemanticAdTransformError::from);
128        }
129        let transformed = semantic_jvp(input, active_inputs, &self.semantic_extension_rules)?;
130        self.ad_transform_cache
131            .put_semantic(key, input, Arc::new(transformed.clone()))
132            .map_err(SemanticAdTransformError::Cache)?;
133        Ok(transformed)
134    }
135
136    /// Transform a frozen semantic program into its reverse-mode derivative.
137    ///
138    /// `active_inputs` selects requested primal-input cotangents and
139    /// `active_outputs` selects primal outputs that receive appended seeds.
140    ///
141    /// # Errors
142    ///
143    /// Returns [`SemanticAdTransformError::ActivityArity`] when either activity
144    /// mask has the wrong length,
145    /// [`SemanticAdTransformError::Extension`] when an extension rule rejects
146    /// the transform, or the corresponding `Query`, `Build`, `Finish`, or
147    /// `Cache` variant when program import, construction, finalization, or
148    /// cache access fails.
149    pub fn vjp_program(
150        &self,
151        input: &FrozenProgram,
152        active_inputs: &[bool],
153        active_outputs: &[bool],
154    ) -> std::result::Result<SemanticAdProgram, SemanticAdTransformError> {
155        let key = SemanticAdTransformCacheKey::vjp(input, active_inputs, active_outputs);
156        if let Some(cached) = self
157            .ad_transform_cache
158            .get_semantic(&key, input)
159            .map_err(SemanticAdTransformError::Cache)?
160        {
161            return cached
162                .as_ref()
163                .with_input_prefix_bindings_from(input)
164                .map_err(SemanticAdTransformError::from);
165        }
166        let transformed = semantic_vjp(
167            input,
168            active_inputs,
169            active_outputs,
170            &self.semantic_extension_rules,
171        )?;
172        self.ad_transform_cache
173            .put_semantic(key, input, Arc::new(transformed.clone()))
174            .map_err(SemanticAdTransformError::Cache)?;
175        Ok(transformed)
176    }
177
178    pub(crate) fn ad_transform_cache(&self) -> Arc<AdTransformCache> {
179        Arc::clone(&self.ad_transform_cache)
180    }
181
182    /// Return AD transform cache retention limits.
183    ///
184    /// # Examples
185    ///
186    /// ```rust
187    /// use tenferro_ad::AdContext;
188    ///
189    /// let ad = AdContext::builder().build().unwrap();
190    /// assert!(ad.ad_transform_cache_limits().unwrap().max_entries().get() > 0);
191    /// ```
192    ///
193    /// # Errors
194    ///
195    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the cache lock is
196    /// poisoned or its state cannot be inspected.
197    pub fn ad_transform_cache_limits(&self) -> Result<AdTransformCacheLimits> {
198        self.ad_transform_cache.limits()
199    }
200
201    /// Replace AD transform cache retention limits.
202    ///
203    /// # Examples
204    ///
205    /// ```rust
206    /// use std::num::NonZeroUsize;
207    /// use tenferro_ad::{AdContext, AdTransformCacheLimits};
208    ///
209    /// let ad = AdContext::builder().build().unwrap();
210    /// let limits = AdTransformCacheLimits::new(NonZeroUsize::new(1).unwrap());
211    /// ad.set_ad_transform_cache_limits(limits).unwrap();
212    /// assert_eq!(ad.ad_transform_cache_limits().unwrap(), limits);
213    /// ```
214    ///
215    /// # Errors
216    ///
217    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the cache lock is
218    /// poisoned while updating the limits.
219    pub fn set_ad_transform_cache_limits(&self, limits: AdTransformCacheLimits) -> Result<()> {
220        self.ad_transform_cache.set_limits(limits)
221    }
222
223    /// Clear AD transform cache entries owned by this context.
224    ///
225    /// # Examples
226    ///
227    /// ```rust
228    /// use tenferro_ad::AdContext;
229    ///
230    /// let ad = AdContext::builder().build().unwrap();
231    /// ad.clear_ad_transform_caches().unwrap();
232    /// assert_eq!(ad.ad_transform_cache_stats().unwrap().entries, 0);
233    /// ```
234    ///
235    /// # Errors
236    ///
237    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the cache lock is
238    /// poisoned while clearing entries.
239    pub fn clear_ad_transform_caches(&self) -> Result<()> {
240        self.ad_transform_cache.clear()
241    }
242
243    /// Return AD transform cache-entry and retained-byte stats.
244    ///
245    /// # Examples
246    ///
247    /// ```rust
248    /// use tenferro_ad::AdContext;
249    ///
250    /// let ad = AdContext::builder().build().unwrap();
251    /// assert_eq!(ad.ad_transform_cache_stats().unwrap().entries, 0);
252    /// ```
253    ///
254    /// # Errors
255    ///
256    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the cache lock is
257    /// poisoned while collecting statistics.
258    pub fn ad_transform_cache_stats(&self) -> Result<CacheStats> {
259        self.ad_transform_cache.stats()
260    }
261
262    /// Clear every cache owned by this AD context.
263    ///
264    /// # Examples
265    ///
266    /// ```rust
267    /// use tenferro_ad::AdContext;
268    ///
269    /// let ad = AdContext::builder().build().unwrap();
270    /// ad.clear_caches().unwrap();
271    /// assert_eq!(ad.cache_stats().unwrap().ad_transforms.entries, 0);
272    /// ```
273    ///
274    /// # Errors
275    ///
276    /// Returns [`tenferro_runtime::Error::RuntimeState`] if either owned cache
277    /// cannot be locked because its state is poisoned.
278    pub fn clear_caches(&self) -> Result<()> {
279        self.clear_ad_transform_caches()
280    }
281
282    /// Return aggregate cache-entry and retained-byte stats for this AD context.
283    ///
284    /// # Examples
285    ///
286    /// ```rust
287    /// use tenferro_ad::AdContext;
288    ///
289    /// let ad = AdContext::builder().build().unwrap();
290    /// assert_eq!(ad.cache_stats().unwrap().ad_transforms.retained_bytes, 0);
291    /// ```
292    ///
293    /// # Errors
294    ///
295    /// Returns [`tenferro_runtime::Error::RuntimeState`] if an owned cache lock
296    /// is poisoned while collecting statistics.
297    pub fn cache_stats(&self) -> Result<AdContextCacheStats> {
298        Ok(AdContextCacheStats {
299            ad_transforms: self.ad_transform_cache_stats()?,
300        })
301    }
302
303    /// Gradient of a scalar traced output with respect to a traced input.
304    ///
305    /// For complex scalar outputs, tenferro returns the Hermitian-adjoint
306    /// cotangent. To compare seed-`1` scalar gradients with JAX's public
307    /// `grad` values, use the complex conjugate of this result. See
308    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
309    ///
310    /// # Examples
311    ///
312    /// ```rust
313    /// use tenferro_ad::AdContext;
314    /// use tenferro_runtime::TracedTensor;
315    ///
316    /// let ad = AdContext::builder().build().unwrap();
317    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
318    /// let loss = (&x * &x).unwrap();
319    /// let grad = ad.grad(&loss, &x).unwrap();
320    /// assert_eq!(grad.rank, 0);
321    /// ```
322    ///
323    /// # Errors
324    ///
325    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] when `output` is not
326    /// scalar, [`tenferro_runtime::Error::UnsupportedAdRule`] when a graph op
327    /// lacks a registered rule, or a typed [`tenferro_runtime::Error::Validation`]
328    /// / backend error when graph metadata or execution is invalid. An inactive
329    /// `wrt` returns [`tenferro_runtime::Error::Validation`] with
330    /// `argument: "wrt"`; use [`grad_optional`](Self::grad_optional) to observe
331    /// that state.
332    pub fn grad(&self, output: &TracedTensor, wrt: &TracedTensor) -> Result<TracedTensor> {
333        crate::traced::grad_with_rules_and_cache(
334            output,
335            wrt,
336            &self.semantic_extension_rules,
337            Some(self.ad_transform_cache.as_ref()),
338        )
339    }
340
341    /// Gradient that returns `None` when `wrt` is inactive.
342    ///
343    /// # Examples
344    ///
345    /// ```rust
346    /// use tenferro_ad::AdContext;
347    /// use tenferro_runtime::TracedTensor;
348    ///
349    /// let ad = AdContext::builder().build().unwrap();
350    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
351    /// let loss = (&x * &x).unwrap();
352    /// assert!(ad.grad_optional(&loss, &x).unwrap().is_some());
353    /// ```
354    ///
355    /// # Errors
356    ///
357    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
358    /// output, [`tenferro_runtime::Error::UnsupportedAdRule`] for an
359    /// unregistered AD rule, or a typed [`tenferro_runtime::Error::Validation`]
360    /// / backend error from graph construction and
361    /// execution.
362    pub fn grad_optional(
363        &self,
364        output: &TracedTensor,
365        wrt: &TracedTensor,
366    ) -> Result<Option<TracedTensor>> {
367        crate::traced::grad_optional_with_rules_and_cache(
368            output,
369            wrt,
370            &self.semantic_extension_rules,
371            Some(self.ad_transform_cache.as_ref()),
372        )
373    }
374
375    /// Forward-mode Jacobian-vector product.
376    ///
377    /// # Examples
378    ///
379    /// ```rust
380    /// use tenferro_ad::AdContext;
381    /// use tenferro_runtime::TracedTensor;
382    ///
383    /// let ad = AdContext::builder().build().unwrap();
384    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
385    /// let dx = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
386    /// let y = (&x * &x).unwrap();
387    /// let dy = ad.jvp(&y, &x, &dx).unwrap();
388    /// assert_eq!(dy.rank, 0);
389    /// ```
390    ///
391    /// # Errors
392    ///
393    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when the graph
394    /// has no JVP rule, [`tenferro_runtime::Error::Validation`] for
395    /// inconsistent tangent metadata, or a typed backend/runtime-state error
396    /// during evaluation. An inactive `wrt` returns
397    /// [`tenferro_runtime::Error::Validation`] with `argument: "wrt"`; use
398    /// [`jvp_optional`](Self::jvp_optional) to observe that state.
399    pub fn jvp(
400        &self,
401        output: &TracedTensor,
402        wrt: &TracedTensor,
403        tangent: &TracedTensor,
404    ) -> Result<TracedTensor> {
405        crate::traced::jvp_with_rules_and_cache(
406            output,
407            wrt,
408            tangent,
409            &self.semantic_extension_rules,
410            Some(self.ad_transform_cache.as_ref()),
411        )
412    }
413
414    /// Forward-mode Jacobian-vector product that returns `None` for inactive output.
415    ///
416    /// # Examples
417    ///
418    /// ```rust
419    /// use tenferro_ad::AdContext;
420    /// use tenferro_runtime::TracedTensor;
421    ///
422    /// let ad = AdContext::builder().build().unwrap();
423    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
424    /// let dx = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
425    /// let y = (&x * &x).unwrap();
426    /// assert!(ad.jvp_optional(&y, &x, &dx).unwrap().is_some());
427    /// ```
428    ///
429    /// # Errors
430    ///
431    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when the graph
432    /// has no JVP rule, [`tenferro_runtime::Error::Validation`] for
433    /// inconsistent tangent metadata, or a typed backend/runtime-state error
434    /// during evaluation.
435    pub fn jvp_optional(
436        &self,
437        output: &TracedTensor,
438        wrt: &TracedTensor,
439        tangent: &TracedTensor,
440    ) -> Result<Option<TracedTensor>> {
441        crate::traced::jvp_optional_with_rules_and_cache(
442            output,
443            wrt,
444            tangent,
445            &self.semantic_extension_rules,
446            Some(self.ad_transform_cache.as_ref()),
447        )
448    }
449
450    /// Forward-mode directional derivative for multiple distinct traced leaves.
451    ///
452    /// Reachable leaves are transformed together in one derivative graph.
453    /// Unreachable leaves contribute nothing; an empty or fully unreachable
454    /// request returns `None`. Duplicate `wrt` leaves are rejected before the
455    /// transform because one semantic seed slot cannot accept two tangents.
456    ///
457    /// # Examples
458    ///
459    /// ```rust
460    /// use tenferro_ad::AdContext;
461    /// use tenferro_runtime::TracedTensor;
462    ///
463    /// let ad = AdContext::builder().build().unwrap();
464    /// let x = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
465    /// let y = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
466    /// let dx = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
467    /// let dy = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
468    /// let output = (&x * &y).unwrap();
469    /// assert!(ad.jvp_many(&output, &[(&x, &dx), (&y, &dy)]).unwrap().is_some());
470    /// ```
471    ///
472    /// # Errors
473    ///
474    /// Returns [`tenferro_runtime::Error::Validation`] for duplicate leaves or
475    /// incompatible tangent metadata, [`tenferro_runtime::Error::UnsupportedAdRule`]
476    /// when a required rule is unavailable, or a typed runtime-state error when
477    /// derivative graph construction fails.
478    ///
479    /// # Deferred errors
480    ///
481    /// Symbolic shape constraints may fail during later compilation or execution.
482    pub fn jvp_many(
483        &self,
484        output: &TracedTensor,
485        wrt_tangents: &[(&TracedTensor, &TracedTensor)],
486    ) -> Result<Option<TracedTensor>> {
487        crate::traced::jvp_many_with_rules_and_cache(
488            output,
489            wrt_tangents,
490            &self.semantic_extension_rules,
491            Some(self.ad_transform_cache.as_ref()),
492        )
493    }
494
495    /// Reverse-mode vector-Jacobian product.
496    ///
497    /// Complex cotangents use tenferro's Hermitian real-inner-product
498    /// convention. Non-real complex cotangent seeds therefore need an explicit
499    /// seed-convention comparison when matching JAX. See
500    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
501    ///
502    /// # Examples
503    ///
504    /// ```rust
505    /// use tenferro_ad::AdContext;
506    /// use tenferro_runtime::TracedTensor;
507    ///
508    /// let ad = AdContext::builder().build().unwrap();
509    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
510    /// let dy = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
511    /// let y = (&x * &x).unwrap();
512    /// let dx = ad.vjp(&y, &x, &dy).unwrap();
513    /// assert_eq!(dx.rank, 0);
514    /// ```
515    ///
516    /// # Errors
517    ///
518    /// Returns [`tenferro_runtime::Error::Validation`] when the cotangent
519    /// metadata is incompatible, [`tenferro_runtime::Error::UnsupportedAdRule`]
520    /// when a VJP rule is unavailable, or a typed backend/runtime-state error
521    /// during execution. An inactive `wrt` returns
522    /// [`tenferro_runtime::Error::Validation`] with `argument: "wrt"`; use
523    /// [`vjp_optional`](Self::vjp_optional) to observe that state.
524    pub fn vjp(
525        &self,
526        output: &TracedTensor,
527        wrt: &TracedTensor,
528        cotangent: &TracedTensor,
529    ) -> Result<TracedTensor> {
530        crate::traced::vjp_with_rules_and_cache(
531            output,
532            wrt,
533            cotangent,
534            &self.semantic_extension_rules,
535            Some(self.ad_transform_cache.as_ref()),
536        )
537    }
538
539    /// Reverse-mode vector-Jacobian product that returns `None` for inactive input.
540    ///
541    /// # Examples
542    ///
543    /// ```rust
544    /// use tenferro_ad::AdContext;
545    /// use tenferro_runtime::TracedTensor;
546    ///
547    /// let ad = AdContext::builder().build().unwrap();
548    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
549    /// let dy = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
550    /// let y = (&x * &x).unwrap();
551    /// assert!(ad.vjp_optional(&y, &x, &dy).unwrap().is_some());
552    /// ```
553    ///
554    /// # Errors
555    ///
556    /// Returns [`tenferro_runtime::Error::Validation`] when the cotangent
557    /// metadata is incompatible, [`tenferro_runtime::Error::UnsupportedAdRule`]
558    /// when a VJP rule is unavailable, or a typed backend/runtime-state error
559    /// during execution.
560    pub fn vjp_optional(
561        &self,
562        output: &TracedTensor,
563        wrt: &TracedTensor,
564        cotangent: &TracedTensor,
565    ) -> Result<Option<TracedTensor>> {
566        crate::traced::vjp_optional_with_rules_and_cache(
567            output,
568            wrt,
569            cotangent,
570            &self.semantic_extension_rules,
571            Some(self.ad_transform_cache.as_ref()),
572        )
573    }
574
575    /// Reverse-mode products for multiple traced leaves in one derivative graph.
576    ///
577    /// Results align with `wrts`; unreachable leaves produce `None`. Duplicate
578    /// leaves are allowed and repeat the same traced derivative without
579    /// accumulating the cotangent twice. An empty request validates that the
580    /// cotangent has concrete data, then returns an empty vector.
581    ///
582    /// # Examples
583    ///
584    /// ```rust
585    /// use tenferro_ad::AdContext;
586    /// use tenferro_runtime::TracedTensor;
587    ///
588    /// let ad = AdContext::builder().build().unwrap();
589    /// let x = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
590    /// let y = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
591    /// let seed = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
592    /// let output = (&x * &y).unwrap();
593    /// let products = ad.vjp_many(&output, &[&x, &y], &seed).unwrap();
594    /// assert!(products.iter().all(Option::is_some));
595    /// ```
596    ///
597    /// # Errors
598    ///
599    /// Returns [`tenferro_runtime::Error::Validation`] for invalid cotangent
600    /// metadata, [`tenferro_runtime::Error::UnsupportedAdRule`] when a required
601    /// rule is unavailable, or a typed runtime-state error when derivative graph
602    /// construction fails.
603    ///
604    /// # Deferred errors
605    ///
606    /// Symbolic shape constraints may fail during later compilation or execution.
607    pub fn vjp_many(
608        &self,
609        output: &TracedTensor,
610        wrts: &[&TracedTensor],
611        cotangent: &TracedTensor,
612    ) -> Result<Vec<Option<TracedTensor>>> {
613        crate::traced::vjp_many_with_rules_and_cache(
614            output,
615            wrts,
616            cotangent,
617            &self.semantic_extension_rules,
618            Some(self.ad_transform_cache.as_ref()),
619        )
620    }
621}
622
623/// Builder for [`AdContext`].
624///
625/// # Examples
626///
627/// ```rust
628/// use tenferro_ad::AdContextBuilder;
629///
630/// let ad = AdContextBuilder::new().build().unwrap();
631/// assert!(ad
632///     .semantic_extension_rules()
633///     .lookup_linearize("example.missing.v1")
634///     .is_none());
635/// ```
636#[derive(Clone, Debug, Default)]
637pub struct AdContextBuilder {
638    semantic_extension_rules: SemanticExtensionRuleSet,
639}
640
641impl AdContextBuilder {
642    /// Create an empty builder.
643    ///
644    /// # Examples
645    ///
646    /// ```rust
647    /// use tenferro_ad::AdContextBuilder;
648    ///
649    /// let _builder = AdContextBuilder::new();
650    /// ```
651    pub fn new() -> Self {
652        Self::default()
653    }
654
655    /// Include an owned semantic-program extension AD rule set.
656    ///
657    /// # Errors
658    ///
659    /// Returns [`SemanticExtensionRegistryError::MalformedFamilyId`] when a
660    /// family identifier is invalid, or
661    /// [`SemanticExtensionRegistryError::DuplicateRule`] when the same family
662    /// and role were already supplied.
663    pub fn with_semantic_extension_rules(
664        mut self,
665        rules: SemanticExtensionRuleSet,
666    ) -> std::result::Result<Self, SemanticExtensionRegistryError> {
667        self.semantic_extension_rules.merge(rules)?;
668        Ok(self)
669    }
670
671    /// Build the context.
672    ///
673    /// Semantic extension rules have already been validated and merged by
674    /// [`Self::with_semantic_extension_rules`].
675    ///
676    /// # Examples
677    ///
678    /// ```rust
679    /// use tenferro_ad::AdContext;
680    ///
681    /// let ad = AdContext::builder().build().unwrap();
682    /// assert!(ad
683    ///     .semantic_extension_rules()
684    ///     .lookup_linearize("example.missing.v1")
685    ///     .is_none());
686    /// ```
687    ///
688    /// # Errors
689    ///
690    /// The error type is [`std::convert::Infallible`], so this finalization step
691    /// never returns `Err` after semantic rule registration. It retains a
692    /// `Result` so callers can compose it with the fallible registration step.
693    pub fn build(self) -> std::result::Result<AdContext, std::convert::Infallible> {
694        Ok(AdContext {
695            semantic_extension_rules: self.semantic_extension_rules,
696            ad_transform_cache: Arc::new(AdTransformCache::new()),
697        })
698    }
699}