tenferro_ops/ext_op.rs
1//! Out-of-tree extension-operation mechanism.
2//!
3//! This module implements the [`ExtensionOp`] trait and its process-local
4//! registry. Together they let external crates contribute fused primitives
5//! that participate in the [`crate::std_tensor_op::StdTensorOp`] graph through
6//! the single carrier variant
7//! `StdTensorOp::Extension(Arc<dyn ExtensionOp>)`.
8//!
9//! See `docs/spec/extension-op.md` for the normative contract. Key points:
10//!
11//! - Identity / hashing / equality are expressed on the trait so the
12//! type-erased `Arc<dyn ExtensionOp>` carrier can satisfy
13//! `Clone + Hash + Eq + Send + Sync + 'static` (computegraph's
14//! `GraphOperation` requirements).
15//! - AD rules are owned by the semantic AD registry in `tenferro-ad`; this
16//! extension-operation contract does not import an AD engine.
17//! - Extension ops themselves do not require process-global registration.
18//! Frontends carry them directly as `Arc<dyn ExtensionOp>`.
19
20use std::any::Any;
21use std::error::Error as StdError;
22use std::fmt::Debug;
23use std::hash::{Hash, Hasher};
24use std::sync::Arc;
25
26use computegraph::graph::GraphBuilder;
27use computegraph::types::ValueRef;
28use tenferro_tensor::{DType, ErrorKind, ValidationKind};
29
30use crate::std_tensor_op::StdTensorOp;
31use crate::sym_dim::SymDim;
32use crate::ExtensionShapeContext;
33
34#[doc(hidden)]
35pub use crate::shape_constraint::ExtensionShapeConstraint;
36
37/// Canonical result of one extension metadata inference callback.
38///
39/// # Examples
40///
41/// ```rust
42/// use tenferro_ops::ext_op::ExtensionShapeInference;
43///
44/// let inferred = ExtensionShapeInference {
45/// output_metas: Vec::new(),
46/// constraints: Vec::new(),
47/// };
48/// assert!(inferred.output_metas.is_empty());
49/// assert!(inferred.constraints.is_empty());
50/// ```
51#[doc(hidden)]
52#[derive(Clone, Debug, PartialEq, Eq)]
53pub struct ExtensionShapeInference {
54 /// Output dtype and symbolic-shape metadata in output-slot order.
55 pub output_metas: Vec<(DType, Vec<SymDim>)>,
56 /// Shape requirements recorded by the callback.
57 pub constraints: Vec<ExtensionShapeConstraint>,
58}
59
60/// Invoke an extension metadata callback after validating its declared arity.
61///
62/// # Examples
63///
64/// ```rust
65/// use std::any::Any;
66/// use std::sync::Arc;
67/// use tenferro_ops::ext_op::{invoke_extension_shape_inference, ExtensionOp};
68/// use tenferro_ops::{ExtensionShapeContext, SymDim};
69/// use tenferro_tensor::DType;
70///
71/// #[derive(Clone, Debug)]
72/// struct Identity;
73///
74/// impl ExtensionOp for Identity {
75/// fn family_id(&self) -> &'static str { "example.identity.v1" }
76/// fn payload_hash(&self, _hasher: &mut dyn std::hash::Hasher) {}
77/// fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
78/// other.as_any().downcast_ref::<Self>().is_some()
79/// }
80/// fn clone_arc(&self) -> Arc<dyn ExtensionOp> { Arc::new(self.clone()) }
81/// fn as_any(&self) -> &dyn Any { self }
82/// fn input_count(&self) -> usize { 1 }
83/// fn output_count(&self) -> usize { 1 }
84/// fn infer_output_meta(
85/// &self,
86/// ctx: &mut ExtensionShapeContext<'_>,
87/// ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
88/// Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
89/// }
90/// }
91///
92/// let shape = [SymDim::from(3usize)];
93/// let inferred = invoke_extension_shape_inference(
94/// &Identity,
95/// &[DType::F64],
96/// &[&shape],
97/// ).unwrap();
98/// assert_eq!(inferred.output_metas, vec![(DType::F64, shape.to_vec())]);
99/// ```
100#[doc(hidden)]
101pub fn invoke_extension_shape_inference(
102 op: &dyn ExtensionOp,
103 input_dtypes: &[DType],
104 input_shapes: &[&[SymDim]],
105) -> tenferro_tensor::Result<ExtensionShapeInference> {
106 let expected_inputs = op.input_count();
107 if input_dtypes.len() != expected_inputs || input_shapes.len() != expected_inputs {
108 return Err(tenferro_tensor::Error::invalid_argument(
109 "extension",
110 "input metadata",
111 format!(
112 "family_id={:?}: infer_output_meta expects {expected_inputs} input metadata entries, got {} dtypes and {} shapes",
113 op.family_id(),
114 input_dtypes.len(),
115 input_shapes.len()
116 ),
117 ));
118 }
119
120 let mut ctx =
121 ExtensionShapeContext::new_for_inference(op.family_id(), input_dtypes, input_shapes);
122 let output_metas = op.infer_output_meta(&mut ctx)?;
123 if output_metas.len() != op.output_count() {
124 return Err(tenferro_tensor::Error::invalid_argument(
125 "extension",
126 "output metadata",
127 format!(
128 "family_id={:?}: infer_output_meta produced {} output metadata entries; op declared {} outputs",
129 op.family_id(),
130 output_metas.len(),
131 op.output_count()
132 ),
133 ));
134 }
135
136 Ok(ExtensionShapeInference {
137 output_metas,
138 constraints: ctx.into_constraints(),
139 })
140}
141
142/// Error returned when an extension cannot expand itself into standard ops.
143///
144/// # Examples
145///
146/// ```
147/// use tenferro_ops::ext_op::ExtensionLoweringError;
148///
149/// let err = ExtensionLoweringError::new("example extension cannot lower");
150/// assert!(err.to_string().contains("cannot lower"));
151/// ```
152#[derive(Debug, thiserror::Error)]
153pub enum ExtensionLoweringError {
154 /// A lowering failure that has no typed source.
155 #[error("{message}")]
156 Message {
157 /// Human-readable lowering detail.
158 message: String,
159 /// Coarse classification supplied by the extension owner.
160 kind: ErrorKind,
161 },
162 /// A lowering failure retaining the domain source that caused it.
163 #[error("{source}")]
164 Source {
165 /// Coarse classification supplied by the extension owner.
166 kind: ErrorKind,
167 /// Original typed lowering source.
168 #[source]
169 source: Box<dyn StdError + Send + Sync + 'static>,
170 },
171}
172
173impl ExtensionLoweringError {
174 /// Create a lowering error with a human-readable diagnostic.
175 ///
176 /// # Examples
177 ///
178 /// ```
179 /// use tenferro_ops::ext_op::ExtensionLoweringError;
180 ///
181 /// let err = ExtensionLoweringError::new("shape must be static");
182 /// assert_eq!(err.to_string(), "shape must be static");
183 /// ```
184 pub fn new(message: impl Into<String>) -> Self {
185 Self::new_with_kind(
186 ErrorKind::Validation(ValidationKind::InvalidArgument),
187 message,
188 )
189 }
190
191 /// Create a lowering error with an explicit coarse classification.
192 ///
193 /// # Examples
194 ///
195 /// ```
196 /// use tenferro_ops::ext_op::ExtensionLoweringError;
197 /// use tenferro_tensor::ErrorKind;
198 ///
199 /// let err = ExtensionLoweringError::new_with_kind(
200 /// ErrorKind::Unsupported,
201 /// "extension is not supported by this lowering target",
202 /// );
203 /// assert_eq!(err.kind(), ErrorKind::Unsupported);
204 /// ```
205 pub fn new_with_kind(kind: ErrorKind, message: impl Into<String>) -> Self {
206 Self::Message {
207 message: message.into(),
208 kind,
209 }
210 }
211
212 /// Create a lowering error while retaining a typed source.
213 ///
214 /// # Examples
215 ///
216 /// ```
217 /// use std::error::Error as _;
218 /// use tenferro_ops::ext_op::ExtensionLoweringError;
219 ///
220 /// let source = std::io::Error::new(std::io::ErrorKind::Other, "shape unavailable");
221 /// let err = ExtensionLoweringError::from_source(source);
222 /// assert!(err.source().is_some());
223 /// ```
224 pub fn from_source<E>(source: E) -> Self
225 where
226 E: StdError + Send + Sync + 'static,
227 {
228 Self::from_source_with_kind(
229 ErrorKind::Validation(ValidationKind::InvalidArgument),
230 source,
231 )
232 }
233
234 /// Create a lowering error with a typed source and explicit classification.
235 ///
236 /// # Examples
237 ///
238 /// ```
239 /// use tenferro_ops::ext_op::ExtensionLoweringError;
240 /// use tenferro_tensor::ErrorKind;
241 ///
242 /// let err = ExtensionLoweringError::from_source_with_kind(
243 /// ErrorKind::BackendFailure,
244 /// std::io::Error::other("backend rejected lowering"),
245 /// );
246 /// assert_eq!(err.kind(), ErrorKind::BackendFailure);
247 /// assert!(std::error::Error::source(&err).is_some());
248 /// ```
249 pub fn from_source_with_kind<E>(kind: ErrorKind, source: E) -> Self
250 where
251 E: StdError + Send + Sync + 'static,
252 {
253 Self::Source {
254 kind,
255 source: Box::new(source),
256 }
257 }
258
259 /// Return the stable classification carried by this lowering failure.
260 ///
261 /// # Examples
262 ///
263 /// ```
264 /// use tenferro_ops::ext_op::ExtensionLoweringError;
265 /// use tenferro_tensor::ErrorKind;
266 ///
267 /// let error = ExtensionLoweringError::new_with_kind(
268 /// ErrorKind::Unsupported,
269 /// "target has no lowering",
270 /// );
271 /// assert_eq!(error.kind(), ErrorKind::Unsupported);
272 /// ```
273 pub fn kind(&self) -> ErrorKind {
274 match self {
275 Self::Message { kind, .. } | Self::Source { kind, .. } => *kind,
276 }
277 }
278}
279
280/// Result returned by [`ExtensionOp::lower_to_standard_ops`].
281pub type ExtensionLoweringResult =
282 std::result::Result<ExtensionStandardLowering, ExtensionLoweringError>;
283
284/// Typed result of trying to lower an extension into standard tensor ops.
285///
286/// Callers branch on this enum instead of encoding unsupported capability as a
287/// successful empty sentinel.
288///
289/// # Examples
290///
291/// ```
292/// use tenferro_ops::ext_op::ExtensionStandardLowering;
293///
294/// let outcome = ExtensionStandardLowering::Unsupported;
295/// assert!(matches!(outcome, ExtensionStandardLowering::Unsupported));
296/// ```
297#[derive(Clone, Debug, PartialEq, Eq)]
298pub enum ExtensionStandardLowering {
299 /// The extension emitted standard tensor graph outputs.
300 Lowered(Vec<ValueRef<StdTensorOp>>),
301 /// The extension has no standard-op lowering for the supplied metadata.
302 Unsupported,
303}
304
305/// The contract every out-of-tree extension primitive must satisfy.
306///
307/// Implementations appear in the core graph as
308/// `StdTensorOp::Extension(Arc<dyn ExtensionOp>)`. Every method is part of the
309/// `ExtensionOp` spec (`docs/spec/extension-op.md`); the short form:
310///
311/// - identity via [`family_id`][Self::family_id] + [`payload_hash`][Self::payload_hash]
312/// + [`payload_eq`][Self::payload_eq];
313/// - fixed arity via [`input_count`][Self::input_count] / [`output_count`][Self::output_count];
314/// - shape / dtype inference via [`infer_output_meta`][Self::infer_output_meta];
315/// - optional fixed-shape standard-op expansion via
316/// [`lower_to_standard_ops`][Self::lower_to_standard_ops] for peer lowerers
317/// such as XLA that cannot execute extension runtimes;
318/// - AD via separately registered role-specific extension rules.
319///
320/// # Downcast convention
321///
322/// Implementations MUST also implement [`Any`] so that
323/// [`ExtensionOp::payload_eq`] can downcast a trait-object reference to
324/// the concrete type. The helper [`ExtensionOp::as_any`] returns
325/// `&dyn Any` for this purpose. Implementations usually define it as
326/// `fn as_any(&self) -> &dyn Any { self }`.
327///
328/// # Examples
329///
330/// ```
331/// # use std::any::Any;
332/// use std::sync::Arc;
333/// use tenferro_ops::ext_op::ExtensionOp;
334/// use tenferro_ops::{ExtensionShapeContext, SymDim};
335/// use tenferro_tensor::DType;
336///
337/// #[derive(Clone, Debug)]
338/// struct IdentityExt;
339///
340/// impl ExtensionOp for IdentityExt {
341/// fn family_id(&self) -> &'static str { "example.identity.v1" }
342/// fn payload_hash(&self, _hasher: &mut dyn std::hash::Hasher) {}
343/// fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
344/// other.as_any().downcast_ref::<IdentityExt>().is_some()
345/// }
346/// fn clone_arc(&self) -> Arc<dyn ExtensionOp> { Arc::new(self.clone()) }
347/// fn as_any(&self) -> &dyn Any { self }
348/// fn input_count(&self) -> usize { 1 }
349/// fn output_count(&self) -> usize { 1 }
350/// fn infer_output_meta(
351/// &self,
352/// ctx: &mut ExtensionShapeContext<'_>,
353/// ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
354/// Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
355/// }
356/// }
357///
358/// let op: Arc<dyn ExtensionOp> = Arc::new(IdentityExt);
359/// assert_eq!(op.input_count(), 1);
360/// ```
361pub trait ExtensionOp: Debug + Send + Sync + 'static {
362 // ----- Identity, hashing, equality (spec Section 5) -----
363
364 /// Stable, process-independent family identifier.
365 ///
366 /// MUST be unique per extension *family* (payload schema), not per
367 /// *instance*, and MUST follow the reserved format
368 /// `"<crate-name>.<op-name>.v<major>"`.
369 fn family_id(&self) -> &'static str;
370
371 /// Hash the payload (everything except `family_id`).
372 ///
373 /// Implementations MUST be pure and deterministic across calls on the same
374 /// value. Hashes MUST NOT include transient state such as allocation
375 /// addresses or atomically updated counters.
376 fn payload_hash(&self, hasher: &mut dyn Hasher);
377
378 /// Structural equality against another extension value.
379 ///
380 /// The carrier's `PartialEq` impl first compares `family_id`s. When the
381 /// family IDs match, it calls `payload_eq`. Implementations MUST return
382 /// `true` iff the payloads are semantically equal AND
383 /// `other.family_id() == self.family_id()`.
384 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool;
385
386 /// Deep-clone the payload behind an `Arc`.
387 ///
388 /// The carrier's `Clone` impl uses `Arc::clone` on the fast path; this
389 /// method exists for rare cases that need a second independent `Arc`.
390 fn clone_arc(&self) -> Arc<dyn ExtensionOp>;
391
392 /// Upcast this extension to `&dyn Any` for downcasting in `payload_eq`.
393 ///
394 /// Implementations SHOULD return `self` verbatim. The method is
395 /// object-safe (no `Self: Sized` bound) so it can be called on an
396 /// `&dyn ExtensionOp`; that's what makes
397 /// `other.as_any().downcast_ref::<ConcreteType>()` work from
398 /// [`Self::payload_eq`] implementations.
399 fn as_any(&self) -> &dyn Any;
400
401 // ----- Arity (spec Section 6) -----
402
403 /// Number of primal inputs. MUST be constant for any given
404 /// `Arc<dyn ExtensionOp>` value.
405 fn input_count(&self) -> usize;
406
407 /// Number of outputs. MUST match the length of the vector returned by a
408 /// successful [`Self::infer_output_meta`] call.
409 fn output_count(&self) -> usize;
410
411 /// Declare observable semantic effects for this extension payload.
412 ///
413 /// The compatibility default is deliberately `Undeclared`, not pure.
414 /// Semantic-program construction rejects an undeclared payload so an
415 /// extension cannot silently acquire purity during migration.
416 fn semantic_effects(&self) -> ExtensionEffectDeclaration<'_> {
417 ExtensionEffectDeclaration::Undeclared
418 }
419
420 /// Declare semantic output aliasing for this extension payload.
421 ///
422 /// The compatibility default is deliberately `Undeclared`, not fresh.
423 /// Execution-only users may continue to carry an older payload, while
424 /// semantic-program construction requires an explicit declaration.
425 fn semantic_aliases(&self) -> ExtensionAliasDeclaration<'_> {
426 ExtensionAliasDeclaration::Undeclared
427 }
428
429 /// Canonical identity of an externally defined scalar this operation carries.
430 ///
431 /// A semantic program's identity must be reproducible across processes, while an
432 /// externally defined scalar's tag is a process-local `TypeId`. An operation
433 /// whose inputs or outputs are externally defined declares the stable name here,
434 /// and the runtime stamps it onto that operation's external value metadata. The
435 /// name belongs to the contribution that owns the scalar, and two scalars must
436 /// not share one.
437 ///
438 /// # Examples
439 ///
440 /// ```rust
441 /// use std::any::Any;
442 /// use std::hash::Hasher;
443 /// use tenferro_ops::ext_op::ExtensionOp;
444 /// use tenferro_ops::{ExtensionShapeContext, SymDim};
445 /// use tenferro_tensor::DType;
446 ///
447 /// #[derive(Clone, Debug)]
448 /// struct Declared;
449 ///
450 /// impl ExtensionOp for Declared {
451 /// fn family_id(&self) -> &'static str {
452 /// "example.declared.v1"
453 /// }
454 /// fn payload_hash(&self, _hasher: &mut dyn Hasher) {}
455 /// fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
456 /// other.as_any().downcast_ref::<Self>().is_some()
457 /// }
458 /// fn clone_arc(&self) -> std::sync::Arc<dyn ExtensionOp> {
459 /// std::sync::Arc::new(self.clone())
460 /// }
461 /// fn as_any(&self) -> &dyn Any {
462 /// self
463 /// }
464 /// fn input_count(&self) -> usize {
465 /// 1
466 /// }
467 /// fn output_count(&self) -> usize {
468 /// 1
469 /// }
470 /// fn scalar_identity(&self) -> Option<&'static str> {
471 /// Some("example.scalar.v1")
472 /// }
473 /// fn infer_output_meta(
474 /// &self,
475 /// _context: &mut ExtensionShapeContext<'_>,
476 /// ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
477 /// Ok(vec![(DType::F64, Vec::new())])
478 /// }
479 /// }
480 ///
481 /// #[derive(Clone, Debug)]
482 /// struct Undeclared;
483 ///
484 /// impl ExtensionOp for Undeclared {
485 /// fn family_id(&self) -> &'static str {
486 /// "example.undeclared.v1"
487 /// }
488 /// fn payload_hash(&self, _hasher: &mut dyn Hasher) {}
489 /// fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
490 /// other.as_any().downcast_ref::<Self>().is_some()
491 /// }
492 /// fn clone_arc(&self) -> std::sync::Arc<dyn ExtensionOp> {
493 /// std::sync::Arc::new(self.clone())
494 /// }
495 /// fn as_any(&self) -> &dyn Any {
496 /// self
497 /// }
498 /// fn input_count(&self) -> usize {
499 /// 1
500 /// }
501 /// fn output_count(&self) -> usize {
502 /// 1
503 /// }
504 /// fn infer_output_meta(
505 /// &self,
506 /// _context: &mut ExtensionShapeContext<'_>,
507 /// ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
508 /// Ok(vec![(DType::F64, Vec::new())])
509 /// }
510 /// }
511 ///
512 /// // A contribution declares the identity of its own scalar, and an operation
513 /// // that carries no external scalar keeps the default.
514 /// assert_eq!(Declared.scalar_identity(), Some("example.scalar.v1"));
515 /// assert_eq!(Undeclared.scalar_identity(), None);
516 /// ```
517 fn scalar_identity(&self) -> Option<&'static str> {
518 None
519 }
520
521 // ----- Shape and dtype inference (spec Section 7) -----
522
523 /// Infer output dtypes and shapes for each output slot.
524 ///
525 /// The canonical inference driver validates arity before invoking this
526 /// callback. Implementations MUST validate rank, dtype, axis, and other
527 /// input-derived metadata through `ctx` before using it. Invalid public
528 /// input must return a typed error rather than an empty sentinel or panic.
529 ///
530 /// On success, the returned vector MUST have length `self.output_count()`,
531 /// one `(dtype, shape)` entry per output slot. Shapes use [`SymDim`] so
532 /// extension ops compose with graph-global symbolic metadata.
533 ///
534 /// # Errors
535 ///
536 /// Returns [`tenferro_tensor::Error::Validation`] for invalid rank, axis,
537 /// or dtype metadata, or [`tenferro_tensor::Error::RuntimeState`] when the
538 /// output contract cannot be inferred from unavailable metadata.
539 fn infer_output_meta(
540 &self,
541 ctx: &mut ExtensionShapeContext<'_>,
542 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>>;
543
544 /// Try to expand this extension into standard tensor graph operations.
545 ///
546 /// Return [`ExtensionStandardLowering::Lowered`] after adding only standard
547 /// [`StdTensorOp`] operations to `builder`.
548 /// [`ExtensionStandardLowering::Unsupported`] means a peer lowerer may try a
549 /// configured fallback; an [`ExtensionLoweringError`] remains a real
550 /// lowering failure and must not be converted into a capability miss.
551 ///
552 /// # Errors
553 ///
554 /// Returns [`ExtensionLoweringError`] when the payload or input metadata
555 /// cannot be lowered safely.
556 fn lower_to_standard_ops(
557 &self,
558 _builder: &mut GraphBuilder<StdTensorOp>,
559 _inputs: &[ValueRef<StdTensorOp>],
560 _input_dtypes: &[DType],
561 _input_shapes: &[&[SymDim]],
562 ) -> ExtensionLoweringResult {
563 Ok(ExtensionStandardLowering::Unsupported)
564 }
565
566 /// Optionally return an equivalent op that produces only live outputs.
567 ///
568 /// `live_outputs` is aligned with this op's current output slots. Return
569 /// `None` when the family does not support output pruning. Return
570 /// `Some(op)` only when the new op's outputs are exactly the live output
571 /// slots, in ascending slot order, and `op.output_count()` equals the
572 /// number of `true` entries in `live_outputs`.
573 fn prune_outputs(&self, _live_outputs: &[bool]) -> Option<Arc<dyn ExtensionOp>> {
574 None
575 }
576
577 // AD rules are registered separately in `tenferro-ad`.
578}
579
580/// Access mode for one extension-declared semantic resource.
581#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
582pub enum ExtensionEffectAccess {
583 /// Read-only access.
584 Read,
585 /// Mutating access.
586 Write,
587}
588
589/// Backend-neutral resource access declared by an extension payload.
590#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
591pub struct ExtensionEffect {
592 /// Stable versioned resource family.
593 pub family: &'static str,
594 /// Family-local resource identity.
595 pub key: u64,
596 /// Read or write access.
597 pub access: ExtensionEffectAccess,
598}
599
600/// Explicit effect declaration returned by an extension payload.
601#[derive(Clone, Copy, Debug, PartialEq, Eq)]
602pub enum ExtensionEffectDeclaration<'a> {
603 /// The payload has not been migrated to the semantic contract.
604 Undeclared,
605 /// Complete ordered effect list; an empty slice explicitly means pure.
606 Declared(&'a [ExtensionEffect]),
607}
608
609/// One extension-declared output alias.
610#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
611pub enum ExtensionAlias {
612 /// The output is semantically fresh.
613 Fresh {
614 /// Operation-local output index.
615 output: usize,
616 },
617 /// The output is a view of an input.
618 ViewOf {
619 /// Operation-local output index.
620 output: usize,
621 /// Operation-local input index.
622 input: usize,
623 },
624 /// The output must alias an input.
625 MustAlias {
626 /// Operation-local output index.
627 output: usize,
628 /// Operation-local input index.
629 input: usize,
630 },
631 /// The output aliases an external typed resource.
632 ExternalAlias {
633 /// Operation-local output index.
634 output: usize,
635 /// Stable versioned resource family.
636 family: &'static str,
637 /// Family-local resource identity.
638 key: u64,
639 },
640}
641
642/// Explicit alias declaration returned by an extension payload.
643#[derive(Clone, Copy, Debug, PartialEq, Eq)]
644pub enum ExtensionAliasDeclaration<'a> {
645 /// The payload has not been migrated to the semantic contract.
646 Undeclared,
647 /// Every output is semantically fresh.
648 AllFresh,
649 /// Complete ordered alias list.
650 Declared(&'a [ExtensionAlias]),
651}
652
653/// Thin adapter that lets a generic `H: Hasher` satisfy the object-safe
654/// `&mut dyn Hasher` signature required by [`ExtensionOp::payload_hash`].
655///
656/// Only `write` and `finish` are load-bearing from the generic hasher; the
657/// various `write_u8` / `write_u16` default implementations in `Hasher`
658/// delegate to `write`. The adapter preserves that behaviour.
659pub(crate) struct DynHasherProxy<'a, H: Hasher + ?Sized> {
660 inner: &'a mut H,
661}
662
663impl<'a, H: Hasher + ?Sized> DynHasherProxy<'a, H> {
664 pub(crate) fn new(inner: &'a mut H) -> Self {
665 Self { inner }
666 }
667}
668
669impl<H: Hasher + ?Sized> Hasher for DynHasherProxy<'_, H> {
670 fn finish(&self) -> u64 {
671 self.inner.finish()
672 }
673
674 fn write(&mut self, bytes: &[u8]) {
675 self.inner.write(bytes);
676 }
677}
678
679/// Hash an `Arc<dyn ExtensionOp>` payload using the extension's
680/// [`ExtensionOp::family_id`] plus [`ExtensionOp::payload_hash`]. Shared
681/// between the `StdTensorOp::Extension` carrier's `Hash` impl and callers
682/// that need to fingerprint an `ExtensionOp` independently.
683pub(crate) fn hash_extension<H: Hasher>(op: &(dyn ExtensionOp + '_), state: &mut H) {
684 op.family_id().as_bytes().hash(state);
685 op.payload_hash(&mut DynHasherProxy::new(state));
686}
687
688/// Structural equality used by the `StdTensorOp::Extension` carrier.
689///
690/// Short-circuits on `family_id` inequality so two extensions with
691/// accidentally similar payloads but different families cannot be unified
692/// by the op interner.
693pub(crate) fn ext_op_eq(a: &dyn ExtensionOp, b: &dyn ExtensionOp) -> bool {
694 a.family_id() == b.family_id() && a.payload_eq(b)
695}