Skip to main content

prebindgen_registry/
domain.rs

1//! Valid subsets of scalar representations used by custom conversions.
2
3use std::ops::{Bound, RangeBounds};
4
5use proc_macro2::TokenStream;
6use quote::{quote, ToTokens};
7
8/// Scalar representations that can carry a declarative domain.
9pub trait DomainScalar: Copy + 'static {
10    #[doc(hidden)]
11    fn domain_value(self) -> ScalarValue;
12    #[doc(hidden)]
13    fn domain_type() -> syn::Type;
14}
15
16macro_rules! impl_ints {
17    ($(($t:ty, $v:ident)),* $(,)?) => {$(
18        impl DomainScalar for $t {
19            fn domain_value(self) -> ScalarValue { ScalarValue::$v(self) }
20            fn domain_type() -> syn::Type { syn::parse_quote!($t) }
21        }
22    )*};
23}
24impl_ints!(
25    (i8, I8),
26    (i16, I16),
27    (i32, I32),
28    (i64, I64),
29    (i128, I128),
30    (u8, U8),
31    (u16, U16),
32    (u32, U32),
33    (u64, U64),
34    (u128, U128),
35);
36
37impl DomainScalar for f32 {
38    fn domain_value(self) -> ScalarValue {
39        ScalarValue::F32(self.to_bits())
40    }
41    fn domain_type() -> syn::Type {
42        syn::parse_quote!(f32)
43    }
44}
45impl DomainScalar for f64 {
46    fn domain_value(self) -> ScalarValue {
47        ScalarValue::F64(self.to_bits())
48    }
49    fn domain_type() -> syn::Type {
50        syn::parse_quote!(f64)
51    }
52}
53
54/// Type-erased scalar. Floats retain their raw IEEE representation.
55#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
56pub enum ScalarValue {
57    I8(i8),
58    I16(i16),
59    I32(i32),
60    I64(i64),
61    I128(i128),
62    U8(u8),
63    U16(u16),
64    U32(u32),
65    U64(u64),
66    U128(u128),
67    F32(u32),
68    F64(u64),
69}
70
71impl ScalarValue {
72    pub fn rust_expr(self) -> syn::Expr {
73        match self {
74            Self::I8(v) => syn::parse_quote!(#v),
75            Self::I16(v) => syn::parse_quote!(#v),
76            Self::I32(v) => syn::parse_quote!(#v),
77            Self::I64(v) => syn::parse_quote!(#v),
78            Self::I128(v) => syn::parse_quote!(#v),
79            Self::U8(v) => syn::parse_quote!(#v),
80            Self::U16(v) => syn::parse_quote!(#v),
81            Self::U32(v) => syn::parse_quote!(#v),
82            Self::U64(v) => syn::parse_quote!(#v),
83            Self::U128(v) => syn::parse_quote!(#v),
84            Self::F32(v) => syn::parse_quote!(f32::from_bits(#v)),
85            Self::F64(v) => syn::parse_quote!(f64::from_bits(#v)),
86        }
87    }
88
89    /// A literal expression suitable for C header generation. Arbitrary NaN
90    /// payloads and infinities have no portable C constant spelling.
91    pub fn portable_expr(self) -> Option<syn::Expr> {
92        match self {
93            Self::F32(bits) => {
94                let value = f32::from_bits(bits);
95                value
96                    .is_finite()
97                    .then(|| syn::parse_str(&format!("{value:?}f32")).expect("finite f32 literal"))
98            }
99            Self::F64(bits) => {
100                let value = f64::from_bits(bits);
101                value
102                    .is_finite()
103                    .then(|| syn::parse_str(&format!("{value:?}f64")).expect("finite f64 literal"))
104            }
105            _ => Some(self.rust_expr()),
106        }
107    }
108
109    pub fn ty(self) -> syn::Type {
110        match self {
111            Self::I8(_) => syn::parse_quote!(i8),
112            Self::I16(_) => syn::parse_quote!(i16),
113            Self::I32(_) => syn::parse_quote!(i32),
114            Self::I64(_) => syn::parse_quote!(i64),
115            Self::I128(_) => syn::parse_quote!(i128),
116            Self::U8(_) => syn::parse_quote!(u8),
117            Self::U16(_) => syn::parse_quote!(u16),
118            Self::U32(_) => syn::parse_quote!(u32),
119            Self::U64(_) => syn::parse_quote!(u64),
120            Self::U128(_) => syn::parse_quote!(u128),
121            Self::F32(_) => syn::parse_quote!(f32),
122            Self::F64(_) => syn::parse_quote!(f64),
123        }
124    }
125
126    fn raw_eq(self, value: &TokenStream) -> TokenStream {
127        match self {
128            Self::F32(bits) => quote!((#value).to_bits() == #bits),
129            Self::F64(bits) => quote!((#value).to_bits() == #bits),
130            _ => {
131                let v = self.rust_expr();
132                quote!((#value) == #v)
133            }
134        }
135    }
136
137    fn cmp_expr(self, value: &TokenStream, op: &str) -> TokenStream {
138        let v = self.rust_expr();
139        match op {
140            ">" => quote!((#value) > #v),
141            ">=" => quote!((#value) >= #v),
142            "<" => quote!((#value) < #v),
143            "<=" => quote!((#value) <= #v),
144            _ => unreachable!(),
145        }
146    }
147
148    fn is_integer_min(self) -> bool {
149        matches!(
150            self,
151            Self::I8(i8::MIN)
152                | Self::I16(i16::MIN)
153                | Self::I32(i32::MIN)
154                | Self::I64(i64::MIN)
155                | Self::I128(i128::MIN)
156                | Self::U8(u8::MIN)
157                | Self::U16(u16::MIN)
158                | Self::U32(u32::MIN)
159                | Self::U64(u64::MIN)
160                | Self::U128(u128::MIN)
161        )
162    }
163
164    fn is_integer_max(self) -> bool {
165        matches!(
166            self,
167            Self::I8(i8::MAX)
168                | Self::I16(i16::MAX)
169                | Self::I32(i32::MAX)
170                | Self::I64(i64::MAX)
171                | Self::I128(i128::MAX)
172                | Self::U8(u8::MAX)
173                | Self::U16(u16::MAX)
174                | Self::U32(u32::MAX)
175                | Self::U64(u64::MAX)
176                | Self::U128(u128::MAX)
177        )
178    }
179}
180
181#[derive(Clone)]
182enum Base {
183    Range {
184        start: Bound<ScalarValue>,
185        end: Bound<ScalarValue>,
186    },
187    Values(Vec<ScalarValue>),
188}
189
190/// Legal values of a custom conversion's scalar representation.
191#[derive(Clone)]
192pub struct RepresentationDomain {
193    ty: syn::Type,
194    base: Base,
195    excluded: Vec<ScalarValue>,
196}
197
198impl RepresentationDomain {
199    pub fn range<T: DomainScalar, R: RangeBounds<T>>(range: R) -> Self {
200        let cvt = |b: Bound<&T>| match b {
201            Bound::Included(v) => Bound::Included((*v).domain_value()),
202            Bound::Excluded(v) => Bound::Excluded((*v).domain_value()),
203            Bound::Unbounded => Bound::Unbounded,
204        };
205        let start = cvt(range.start_bound());
206        let end = cvt(range.end_bound());
207        assert!(
208            !bound_is_nan(&start) && !bound_is_nan(&end),
209            "representation-domain range bounds cannot be NaN"
210        );
211        assert!(
212            range_is_nonempty(&start, &end),
213            "representation-domain range cannot be empty"
214        );
215        Self {
216            ty: T::domain_type(),
217            base: Base::Range { start, end },
218            excluded: vec![],
219        }
220    }
221
222    pub fn values<T: DomainScalar>(values: impl IntoIterator<Item = T>) -> Self {
223        let values = dedup(values.into_iter().map(T::domain_value).collect());
224        assert!(
225            !values.is_empty(),
226            "representation-domain valid set cannot be empty"
227        );
228        Self {
229            ty: T::domain_type(),
230            base: Base::Values(values),
231            excluded: vec![],
232        }
233    }
234
235    pub fn exclude<T: DomainScalar>(&mut self, values: impl IntoIterator<Item = T>) {
236        assert_eq!(
237            prebindgen_flat::flat::TypeKey::from_type(&self.ty),
238            prebindgen_flat::flat::TypeKey::from_type(&T::domain_type()),
239            "representation-domain exclusions must use the base domain's scalar type"
240        );
241        self.excluded
242            .extend(values.into_iter().map(T::domain_value));
243        self.excluded = dedup(std::mem::take(&mut self.excluded));
244    }
245
246    pub fn ty(&self) -> &syn::Type {
247        &self.ty
248    }
249
250    /// An expression testing whether `value` lies inside the legal domain —
251    /// what a generated bounds check is built from.
252    pub fn contains_expr(&self, value: TokenStream) -> TokenStream {
253        let base = match &self.base {
254            Base::Range { start, end } => {
255                let lo = bound_expr(start, &value, true);
256                let hi = bound_expr(end, &value, false);
257                let not_nan = if matches!(
258                    self.ty.to_token_stream().to_string().as_str(),
259                    "f32" | "f64"
260                ) {
261                    quote!(!(#value).is_nan())
262                } else {
263                    quote!(true)
264                };
265                quote!(#not_nan && #lo && #hi)
266            }
267            Base::Values(values) => {
268                let checks = values.iter().map(|v| v.raw_eq(&value));
269                quote!(false #(|| #checks)*)
270            }
271        };
272        let excluded = self.excluded.iter().map(|v| v.raw_eq(&value));
273        quote!((#base) && !(false #(|| #excluded)*))
274    }
275
276    /// Derive a bounded number of stable values outside the legal domain.
277    ///
278    /// Adapter-facing, with [`ScalarValue::portable_expr`]: a back-end that
279    /// gives a sum type a niche-based ABI needs the values to reserve.
280    pub fn niche_values(&self, limit: usize) -> Vec<ScalarValue> {
281        let mut out = Vec::new();
282        let extra = match &self.base {
283            Base::Values(v) => v.len(),
284            _ => 0,
285        };
286        let budget = limit.saturating_add(extra).max(1);
287        match self.ty.to_token_stream().to_string().as_str() {
288            "i8" => ints!(out, self, i8, I8, budget),
289            "i16" => ints!(out, self, i16, I16, budget),
290            "i32" => ints!(out, self, i32, I32, budget),
291            "i64" => ints!(out, self, i64, I64, budget),
292            "i128" => ints!(out, self, i128, I128, budget),
293            "u8" => ints!(out, self, u8, U8, budget),
294            "u16" => ints!(out, self, u16, U16, budget),
295            "u32" => ints!(out, self, u32, U32, budget),
296            "u64" => ints!(out, self, u64, U64, budget),
297            "u128" => ints!(out, self, u128, U128, budget),
298            "f32" => float32_candidates(&mut out, budget),
299            "f64" => float64_candidates(&mut out, budget),
300            _ => unreachable!(),
301        }
302        out.extend(self.excluded.iter().copied());
303        let mut selected = Vec::new();
304        for value in out {
305            if !self.contains(value) && !selected.contains(&value) {
306                selected.push(value);
307                if selected.len() == limit {
308                    break;
309                }
310            }
311        }
312        selected
313    }
314
315    fn contains(&self, value: ScalarValue) -> bool {
316        let base = match &self.base {
317            Base::Range { start, end } => {
318                !is_nan(value) && lower_ok(value, start) && upper_ok(value, end)
319            }
320            Base::Values(values) => values.contains(&value),
321        };
322        base && !self.excluded.contains(&value)
323    }
324}
325
326macro_rules! ints {
327    ($out:expr, $d:expr, $ty:ty, $variant:ident, $budget:expr) => {{
328        let mut hi = <$ty>::MAX;
329        let mut lo = <$ty>::MIN;
330        for _ in 0..$budget {
331            let h = ScalarValue::$variant(hi);
332            if !$d.contains(h) {
333                $out.push(h);
334            }
335            hi = hi.saturating_sub(1);
336            let l = ScalarValue::$variant(lo);
337            if !$d.contains(l) {
338                $out.push(l);
339            }
340            lo = lo.saturating_add(1);
341        }
342    }};
343}
344use ints;
345
346fn float32_candidates(out: &mut Vec<ScalarValue>, n: usize) {
347    for i in 0..n {
348        out.push(ScalarValue::F32(f32::MAX.to_bits() - i as u32));
349        out.push(ScalarValue::F32((-f32::MAX).to_bits() - i as u32));
350    }
351    out.push(ScalarValue::F32(f32::INFINITY.to_bits()));
352    out.push(ScalarValue::F32(f32::NEG_INFINITY.to_bits()));
353    for i in 0..n {
354        out.push(ScalarValue::F32(0x7fc0_0000 + i as u32));
355    }
356}
357fn float64_candidates(out: &mut Vec<ScalarValue>, n: usize) {
358    for i in 0..n {
359        out.push(ScalarValue::F64(f64::MAX.to_bits() - i as u64));
360        out.push(ScalarValue::F64((-f64::MAX).to_bits() - i as u64));
361    }
362    out.push(ScalarValue::F64(f64::INFINITY.to_bits()));
363    out.push(ScalarValue::F64(f64::NEG_INFINITY.to_bits()));
364    for i in 0..n {
365        out.push(ScalarValue::F64(0x7ff8_0000_0000_0000 + i as u64));
366    }
367}
368
369fn bound_expr(bound: &Bound<ScalarValue>, value: &TokenStream, lower: bool) -> TokenStream {
370    match bound {
371        Bound::Unbounded => quote!(true),
372        Bound::Included(v) if lower && v.is_integer_min() => quote!(true),
373        Bound::Included(v) if !lower && v.is_integer_max() => quote!(true),
374        Bound::Included(v) if lower => v.cmp_expr(value, ">="),
375        Bound::Excluded(v) if lower => v.cmp_expr(value, ">"),
376        Bound::Included(v) => v.cmp_expr(value, "<="),
377        Bound::Excluded(v) => v.cmp_expr(value, "<"),
378    }
379}
380fn lower_ok(v: ScalarValue, b: &Bound<ScalarValue>) -> bool {
381    match b {
382        Bound::Unbounded => true,
383        Bound::Included(x) => cmp(v, *x).is_some_and(|v| v >= 0),
384        Bound::Excluded(x) => cmp(v, *x).is_some_and(|v| v > 0),
385    }
386}
387fn upper_ok(v: ScalarValue, b: &Bound<ScalarValue>) -> bool {
388    match b {
389        Bound::Unbounded => true,
390        Bound::Included(x) => cmp(v, *x).is_some_and(|v| v <= 0),
391        Bound::Excluded(x) => cmp(v, *x).is_some_and(|v| v < 0),
392    }
393}
394fn cmp(a: ScalarValue, b: ScalarValue) -> Option<i8> {
395    macro_rules! c {
396        ($a:expr, $b:expr) => {
397            Some(if $a < $b {
398                -1
399            } else if $a > $b {
400                1
401            } else {
402                0
403            })
404        };
405    }
406    match (a, b) {
407        (ScalarValue::I8(a), ScalarValue::I8(b)) => c!(a, b),
408        (ScalarValue::I16(a), ScalarValue::I16(b)) => c!(a, b),
409        (ScalarValue::I32(a), ScalarValue::I32(b)) => c!(a, b),
410        (ScalarValue::I64(a), ScalarValue::I64(b)) => c!(a, b),
411        (ScalarValue::I128(a), ScalarValue::I128(b)) => c!(a, b),
412        (ScalarValue::U8(a), ScalarValue::U8(b)) => c!(a, b),
413        (ScalarValue::U16(a), ScalarValue::U16(b)) => c!(a, b),
414        (ScalarValue::U32(a), ScalarValue::U32(b)) => c!(a, b),
415        (ScalarValue::U64(a), ScalarValue::U64(b)) => c!(a, b),
416        (ScalarValue::U128(a), ScalarValue::U128(b)) => c!(a, b),
417        (ScalarValue::F32(a), ScalarValue::F32(b)) => {
418            f32::from_bits(a).partial_cmp(&f32::from_bits(b)).map(ord)
419        }
420        (ScalarValue::F64(a), ScalarValue::F64(b)) => {
421            f64::from_bits(a).partial_cmp(&f64::from_bits(b)).map(ord)
422        }
423        _ => None,
424    }
425}
426fn ord(v: std::cmp::Ordering) -> i8 {
427    match v {
428        std::cmp::Ordering::Less => -1,
429        std::cmp::Ordering::Equal => 0,
430        std::cmp::Ordering::Greater => 1,
431    }
432}
433fn is_nan(v: ScalarValue) -> bool {
434    match v {
435        ScalarValue::F32(v) => f32::from_bits(v).is_nan(),
436        ScalarValue::F64(v) => f64::from_bits(v).is_nan(),
437        _ => false,
438    }
439}
440fn bound_is_nan(v: &Bound<ScalarValue>) -> bool {
441    match v {
442        Bound::Included(v) | Bound::Excluded(v) => is_nan(*v),
443        Bound::Unbounded => false,
444    }
445}
446fn range_is_nonempty(start: &Bound<ScalarValue>, end: &Bound<ScalarValue>) -> bool {
447    let (start_value, start_included) = match start {
448        Bound::Unbounded => return true,
449        Bound::Included(value) => (*value, true),
450        Bound::Excluded(value) => (*value, false),
451    };
452    let (end_value, end_included) = match end {
453        Bound::Unbounded => return true,
454        Bound::Included(value) => (*value, true),
455        Bound::Excluded(value) => (*value, false),
456    };
457    cmp(start_value, end_value)
458        .is_some_and(|ordering| ordering < 0 || (ordering == 0 && start_included && end_included))
459}
460fn dedup(values: Vec<ScalarValue>) -> Vec<ScalarValue> {
461    let mut out = Vec::new();
462    for v in values {
463        if !out.contains(&v) {
464            out.push(v);
465        }
466    }
467    out
468}
469
470#[cfg(test)]
471mod tests {
472    use super::*;
473
474    #[test]
475    fn integer_range_derives_extreme_niches() {
476        let domain = RepresentationDomain::range(0u64..=1_000_000);
477        assert_eq!(
478            domain.niche_values(3),
479            vec![
480                ScalarValue::U64(u64::MAX),
481                ScalarValue::U64(u64::MAX - 1),
482                ScalarValue::U64(u64::MAX - 2),
483            ]
484        );
485    }
486
487    #[test]
488    fn integer_extreme_bounds_do_not_emit_useless_comparisons() {
489        let domain = RepresentationDomain::range(0u64..=1_000_000);
490        let expr = domain.contains_expr(quote!(value)).to_string();
491        assert!(!expr.contains(">= 0u64"), "{expr}");
492        assert!(expr.contains("<= 1000000u64"), "{expr}");
493
494        let full = RepresentationDomain::range(i32::MIN..=i32::MAX);
495        let expr = full.contains_expr(quote!(value)).to_string();
496        assert!(!expr.contains(">="), "{expr}");
497        assert!(!expr.contains("<="), "{expr}");
498    }
499
500    #[test]
501    fn valid_set_and_exclusions_use_raw_float_bits() {
502        let domain = RepresentationDomain::values([0.0f64, -0.0f64, 1.0]);
503        assert!(domain.contains(ScalarValue::F64(0.0f64.to_bits())));
504        assert!(domain.contains(ScalarValue::F64((-0.0f64).to_bits())));
505
506        let mut range = RepresentationDomain::range(-1.0f64..=1.0f64);
507        range.exclude([0.5f64]);
508        assert!(!range.contains(ScalarValue::F64(0.5f64.to_bits())));
509        assert!(!range.contains(ScalarValue::F64(f64::NAN.to_bits())));
510    }
511
512    #[test]
513    #[should_panic(expected = "cannot be NaN")]
514    fn nan_range_bound_is_rejected() {
515        let _ = RepresentationDomain::range(f64::NAN..=1.0);
516    }
517
518    #[test]
519    #[should_panic(expected = "range cannot be empty")]
520    fn empty_range_is_rejected() {
521        let _ = RepresentationDomain::range(2u8..2u8);
522    }
523}