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}