1use std::ops::{Bound, RangeBounds};
4
5use proc_macro2::TokenStream;
6use quote::{quote, ToTokens};
7
8pub 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#[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 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#[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 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 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}