1use tenferro_ops::dim_expr::DimExpr;
2use tenferro_ops::shape_extent::ShapeExtent;
3use tenferro_tensor::DType;
4
5use super::EffectResourceError;
6
7#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
9#[non_exhaustive]
10pub enum SemanticProvenanceKind {
11 Builder,
13 Imported,
15 Derived,
17}
18
19#[derive(Clone)]
20pub(crate) struct SemanticProvenance {
21 kind: SemanticProvenanceKind,
22 label: Option<std::sync::Arc<str>>,
23}
24
25impl SemanticProvenance {
26 pub(crate) fn builder(label: Option<&str>) -> Self {
27 Self {
28 kind: SemanticProvenanceKind::Builder,
29 label: label.map(std::sync::Arc::from),
30 }
31 }
32
33 pub(crate) fn view(&self) -> SemanticProvenanceView<'_> {
34 SemanticProvenanceView {
35 kind: self.kind,
36 label: self.label.as_deref(),
37 }
38 }
39}
40
41#[derive(Clone, Copy)]
43pub struct SemanticProvenanceView<'a> {
44 kind: SemanticProvenanceKind,
45 label: Option<&'a str>,
46}
47
48impl<'a> SemanticProvenanceView<'a> {
49 pub const fn kind(self) -> SemanticProvenanceKind {
51 self.kind
52 }
53
54 pub const fn label(self) -> Option<&'a str> {
56 self.label
57 }
58}
59
60impl std::fmt::Debug for SemanticProvenanceView<'_> {
61 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
62 formatter
63 .debug_struct("SemanticProvenanceView")
64 .field("kind", &self.kind)
65 .field("has_label", &self.label.is_some())
66 .finish()
67 }
68}
69
70#[derive(Clone, Debug, PartialEq, Eq, Hash)]
72pub struct ProgramValueMetadata {
73 dtype: DType,
74 shape: Box<[ShapeExtent<DimExpr>]>,
75 scalar_identity: Option<&'static str>,
76}
77
78impl ProgramValueMetadata {
79 pub fn new(dtype: DType, shape: impl IntoIterator<Item = DimExpr>) -> Self {
81 Self {
82 dtype,
83 shape: shape.into_iter().map(ShapeExtent::Exact).collect(),
84 scalar_identity: None,
85 }
86 }
87
88 pub fn from_extents(
90 dtype: DType,
91 shape: impl IntoIterator<Item = ShapeExtent<DimExpr>>,
92 ) -> Self {
93 Self {
94 dtype,
95 shape: shape.into_iter().collect(),
96 scalar_identity: None,
97 }
98 }
99
100 #[must_use]
121 pub fn with_scalar_identity(mut self, identity: &'static str) -> Self {
122 self.scalar_identity = Some(identity);
123 self
124 }
125
126 #[must_use]
128 pub const fn scalar_identity(&self) -> Option<&'static str> {
129 self.scalar_identity
130 }
131
132 pub const fn dtype(&self) -> DType {
134 self.dtype
135 }
136
137 pub fn shape(&self) -> &[ShapeExtent<DimExpr>] {
139 &self.shape
140 }
141
142 pub(crate) fn logical_retained_bytes(&self) -> Option<usize> {
143 checked_sum([
144 self.shape
145 .len()
146 .checked_mul(std::mem::size_of::<ShapeExtent<DimExpr>>())?,
147 checked_sum_options(self.shape.iter().map(shape_extent_logical_retained_bytes))?,
148 ])
149 }
150}
151
152#[derive(Clone, Debug, PartialEq, Eq, Hash)]
154pub struct ProgramInputSpec {
155 metadata: ProgramValueMetadata,
156}
157
158impl ProgramInputSpec {
159 pub fn new(dtype: DType, shape: impl IntoIterator<Item = DimExpr>) -> Self {
161 Self {
162 metadata: ProgramValueMetadata::new(dtype, shape),
163 }
164 }
165
166 #[must_use]
183 pub fn with_scalar_identity(mut self, identity: &'static str) -> Self {
184 self.metadata = self.metadata.with_scalar_identity(identity);
185 self
186 }
187
188 pub fn from_metadata(metadata: ProgramValueMetadata) -> Self {
190 Self { metadata }
191 }
192
193 pub const fn metadata(&self) -> &ProgramValueMetadata {
195 &self.metadata
196 }
197}
198
199#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
201#[non_exhaustive]
202pub enum ProgramShapeRelation {
203 Equal,
205 LessEqual,
207 GreaterEqual,
209}
210
211#[derive(Clone, Debug)]
213pub struct ShapeGuard {
214 relation: ProgramShapeRelation,
215 lhs: DimExpr,
216 rhs: DimExpr,
217 source_family: Option<&'static str>,
218}
219
220impl ShapeGuard {
221 pub fn new(relation: ProgramShapeRelation, lhs: DimExpr, rhs: DimExpr) -> Self {
223 Self {
224 relation,
225 lhs,
226 rhs,
227 source_family: None,
228 }
229 }
230
231 pub(crate) fn with_source_family(mut self, family: &'static str) -> Self {
232 self.source_family = Some(family);
233 self
234 }
235
236 pub const fn relation(&self) -> ProgramShapeRelation {
238 self.relation
239 }
240
241 pub const fn lhs(&self) -> &DimExpr {
243 &self.lhs
244 }
245
246 pub const fn rhs(&self) -> &DimExpr {
248 &self.rhs
249 }
250
251 pub const fn source_family(&self) -> Option<&'static str> {
256 self.source_family
257 }
258
259 pub(crate) fn logical_retained_bytes(&self) -> Option<usize> {
260 checked_sum([
261 dim_expr_logical_retained_bytes(&self.lhs)?,
262 dim_expr_logical_retained_bytes(&self.rhs)?,
263 ])
264 }
265}
266
267impl PartialEq for ShapeGuard {
268 fn eq(&self, other: &Self) -> bool {
269 self.relation == other.relation && self.lhs == other.lhs && self.rhs == other.rhs
270 }
271}
272
273impl Eq for ShapeGuard {}
274
275impl std::hash::Hash for ShapeGuard {
276 fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
277 self.relation.hash(state);
278 self.lhs.hash(state);
279 self.rhs.hash(state);
280 }
281}
282
283fn shape_extent_logical_retained_bytes(extent: &ShapeExtent<DimExpr>) -> Option<usize> {
284 match extent {
285 ShapeExtent::Exact(expression) | ShapeExtent::UpperBound(expression) => {
286 dim_expr_logical_retained_bytes(expression)
287 }
288 ShapeExtent::Unknown => Some(0),
289 }
290}
291
292fn dim_expr_logical_retained_bytes(expression: &DimExpr) -> Option<usize> {
293 match expression {
294 DimExpr::Const(_) | DimExpr::InputDim { .. } => Some(0),
295 DimExpr::Add(left, right)
296 | DimExpr::Sub(left, right)
297 | DimExpr::Mul(left, right)
298 | DimExpr::FloorDiv(left, right)
299 | DimExpr::Min(left, right)
300 | DimExpr::Max(left, right) => checked_sum([
301 2usize.checked_mul(std::mem::size_of::<DimExpr>())?,
302 dim_expr_logical_retained_bytes(left)?,
303 dim_expr_logical_retained_bytes(right)?,
304 ]),
305 }
306}
307
308fn checked_sum(values: impl IntoIterator<Item = usize>) -> Option<usize> {
309 values
310 .into_iter()
311 .try_fold(0usize, |sum, value| sum.checked_add(value))
312}
313
314fn checked_sum_options(values: impl IntoIterator<Item = Option<usize>>) -> Option<usize> {
315 values
316 .into_iter()
317 .try_fold(0usize, |sum, value| sum.checked_add(value?))
318}
319
320#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
322pub struct EffectResource {
323 family: &'static str,
324 key: u64,
325}
326
327impl EffectResource {
328 pub fn new(family: &'static str, key: u64) -> Result<Self, EffectResourceError> {
335 let version = family.rsplit_once(".v").map(|(_, version)| version);
336 if family.is_empty()
337 || !version.is_some_and(|version| {
338 !version.is_empty() && version.bytes().all(|byte| byte.is_ascii_digit())
339 })
340 {
341 return Err(EffectResourceError::InvalidFamily);
342 }
343 Ok(Self { family, key })
344 }
345
346 pub const fn family(self) -> &'static str {
348 self.family
349 }
350
351 pub const fn key(self) -> u64 {
353 self.key
354 }
355}
356
357#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
359pub enum EffectAccess {
360 Read,
362 Write,
364}
365
366#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
368pub struct Effect {
369 resource: EffectResource,
370 access: EffectAccess,
371}
372
373impl Effect {
374 pub const fn new(resource: EffectResource, access: EffectAccess) -> Self {
376 Self { resource, access }
377 }
378
379 pub const fn resource(self) -> EffectResource {
381 self.resource
382 }
383
384 pub const fn access(self) -> EffectAccess {
386 self.access
387 }
388}
389
390#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
392pub enum AliasKind {
393 Fresh,
395 ViewOf,
397 MustAlias,
399 ExternalAlias,
401}
402
403#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
405pub struct Alias {
406 kind: AliasKind,
407 output: usize,
408 input: Option<usize>,
409 resource: Option<EffectResource>,
410}
411
412impl Alias {
413 pub const fn fresh(output: usize) -> Self {
415 Self {
416 kind: AliasKind::Fresh,
417 output,
418 input: None,
419 resource: None,
420 }
421 }
422
423 pub const fn view_of(output: usize, input: usize) -> Self {
425 Self {
426 kind: AliasKind::ViewOf,
427 output,
428 input: Some(input),
429 resource: None,
430 }
431 }
432
433 pub const fn must_alias(output: usize, input: usize) -> Self {
435 Self {
436 kind: AliasKind::MustAlias,
437 output,
438 input: Some(input),
439 resource: None,
440 }
441 }
442
443 pub const fn external(output: usize, resource: EffectResource) -> Self {
445 Self {
446 kind: AliasKind::ExternalAlias,
447 output,
448 input: None,
449 resource: Some(resource),
450 }
451 }
452
453 pub const fn kind(self) -> AliasKind {
455 self.kind
456 }
457
458 pub const fn output(self) -> usize {
460 self.output
461 }
462
463 pub const fn input(self) -> Option<usize> {
465 self.input
466 }
467
468 pub const fn resource(self) -> Option<EffectResource> {
470 self.resource
471 }
472}
473
474#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
476#[non_exhaustive]
477pub enum SemanticPlacementKind {
478 Any,
480 SameAsInput,
482}
483
484#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
486pub struct SemanticPlacementConstraint {
487 kind: SemanticPlacementKind,
488 input: Option<usize>,
489}
490
491impl SemanticPlacementConstraint {
492 pub const fn any() -> Self {
494 Self {
495 kind: SemanticPlacementKind::Any,
496 input: None,
497 }
498 }
499
500 pub const fn same_as_input(input: usize) -> Self {
502 Self {
503 kind: SemanticPlacementKind::SameAsInput,
504 input: Some(input),
505 }
506 }
507
508 pub const fn kind(self) -> SemanticPlacementKind {
510 self.kind
511 }
512
513 pub const fn input(self) -> Option<usize> {
515 self.input
516 }
517}