From 30e1a1bde9fda18c00b0ec7daaa830a71a34218a Mon Sep 17 00:00:00 2001 From: hlsxx Date: Fri, 28 Aug 2026 18:50:52 +0200 Subject: [PATCH 1/2] fix: const eval cast of single-variant enum --- crates/hir-ty/src/consteval/tests.rs | 13 ++++++++ crates/hir-ty/src/mir/eval.rs | 9 +++++- crates/hir-ty/src/mir/lower.rs | 48 ++++++++++++++++++++++++---- 3 files changed, 62 insertions(+), 8 deletions(-) diff --git a/crates/hir-ty/src/consteval/tests.rs b/crates/hir-ty/src/consteval/tests.rs index db8f94e40354..c2fb4fa53880 100644 --- a/crates/hir-ty/src/consteval/tests.rs +++ b/crates/hir-ty/src/consteval/tests.rs @@ -2542,6 +2542,19 @@ fn unsupported_str_binary_op_in_eval_rvalue() { ); } +#[test] +fn enum_variant_as_usize_in_const() { + check_number( + r#" +enum Foo { + Bar = 4, +} +const GOAL: usize = Foo::Bar as usize; + "#, + 4, + ); +} + #[test] fn const_loop() { check_fail( diff --git a/crates/hir-ty/src/mir/eval.rs b/crates/hir-ty/src/mir/eval.rs index 104552baa01e..7419aef5e6b1 100644 --- a/crates/hir-ty/src/mir/eval.rs +++ b/crates/hir-ty/src/mir/eval.rs @@ -1468,7 +1468,14 @@ impl<'a, 'db> Evaluator<'a, 'db> { Rvalue::Discriminant(p) => { let ty = self.place_ty(p, locals)?; let bytes = self.eval_place(p, locals)?.get(self)?; - let result = self.compute_discriminant(ty, bytes)?; + let result = if let Some(f) = locals.body.owner.as_variant() + && let Some((AdtId::EnumId(e), _)) = ty.as_adt() + && f.lookup(self.db).parent == e + { + i128::from_le_bytes(pad16(bytes, IsSigned::Yes)) + } else { + self.compute_discriminant(ty, bytes)? + }; Owned(result.to_le_bytes().to_vec()) } Rvalue::Repeat(it, len) => { diff --git a/crates/hir-ty/src/mir/lower.rs b/crates/hir-ty/src/mir/lower.rs index 3b82ed6f798e..1f9bc392115a 100644 --- a/crates/hir-ty/src/mir/lower.rs +++ b/crates/hir-ty/src/mir/lower.rs @@ -966,23 +966,57 @@ impl<'a, 'db> MirLowerCtx<'a, 'db> { self.lower_expr_to_place(id, place, current) } Expr::Cast { expr, type_ref: _ } => { - let Some((it, current)) = self.lower_expr_to_some_operand(*expr, current)? else { - return Ok(None); - }; // Since we don't have THIR, this is the "zipped" version of [rustc's HIR lowering](https://github.com/rust-lang/rust/blob/e71f9529121ca8f687e4b725e3c9adc3f1ebab4d/compiler/rustc_mir_build/src/thir/cx/expr.rs#L165-L178) // and [THIR lowering as RValue](https://github.com/rust-lang/rust/blob/a4601859ae3875732797873612d424976d9e3dd0/compiler/rustc_mir_build/src/build/expr/as_rvalue.rs#L193-L313) - let rvalue = if self.infer.coercion_casts.contains(expr) { - Rvalue::Use(it) + let (rvalue, current) = if self.infer.coercion_casts.contains(expr) { + let Some((it, current)) = self.lower_expr_to_some_operand(*expr, current)? + else { + return Ok(None); + }; + (Rvalue::Use(it), current) } else { let source_ty = self.infer.expr_ty(*expr); let target_ty = self.infer.expr_ty(expr_id); + let (it, source_ty, current) = if let TyKind::Adt(adt, _) = source_ty.kind() + && adt.is_enum() + { + let Some((enum_place, current)) = + self.lower_expr_as_place(current, *expr, true)? + else { + return Ok(None); + }; + let discr_ty = Ty::new_int(self.interner(), rustc_type_ir::IntTy::I128); + let discr_place: Place<'db> = + self.temp(discr_ty, current, expr_id.into())?.into(); + + self.push_assignment( + current, + discr_place, + Rvalue::Discriminant(enum_place.store()), + expr_id.into(), + ); + ( + Operand { + kind: OperandKind::Copy(discr_place.store()), + span: Some(expr_id.into()), + }, + discr_ty, + current, + ) + } else { + let Some((it, current)) = + self.lower_expr_to_some_operand(*expr, current)? + else { + return Ok(None); + }; + (it, source_ty, current) + }; let cast_kind = if source_ty.as_reference().is_some() { CastKind::PointerCoercion(PointerCast::ArrayToPointer) } else { cast_kind(self.db, source_ty, target_ty)? }; - - Rvalue::Cast(cast_kind, it, target_ty.store()) + (Rvalue::Cast(cast_kind, it, target_ty.store()), current) }; self.push_assignment(current, place, rvalue, expr_id.into()); Ok(Some(current)) From 185846c0b61785283e354bef6d8f3c430bc9d76c Mon Sep 17 00:00:00 2001 From: hlsxx Date: Mon, 14 Sep 2026 20:10:29 +0200 Subject: [PATCH 2/2] fix: avoid recursion when casting enum variant discriminants When an enum variant discriminant initializer references another variant via a cast, e.g. `B = Foo::A as isize + 1` the previous implementation would lower `Foo::A` as a place and then extract its discriminant via `Rvalue::Discriminant`. This caused infinite recursion during consteval because evaluating the discriminant of one variant requires evaluating another. # Conflicts: # crates/hir-ty/src/mir/lower.rs --- crates/hir-ty/src/mir/eval.rs | 9 +--- crates/hir-ty/src/mir/lower.rs | 80 ++++++++++++++++++++-------------- 2 files changed, 49 insertions(+), 40 deletions(-) diff --git a/crates/hir-ty/src/mir/eval.rs b/crates/hir-ty/src/mir/eval.rs index 7419aef5e6b1..104552baa01e 100644 --- a/crates/hir-ty/src/mir/eval.rs +++ b/crates/hir-ty/src/mir/eval.rs @@ -1468,14 +1468,7 @@ impl<'a, 'db> Evaluator<'a, 'db> { Rvalue::Discriminant(p) => { let ty = self.place_ty(p, locals)?; let bytes = self.eval_place(p, locals)?.get(self)?; - let result = if let Some(f) = locals.body.owner.as_variant() - && let Some((AdtId::EnumId(e), _)) = ty.as_adt() - && f.lookup(self.db).parent == e - { - i128::from_le_bytes(pad16(bytes, IsSigned::Yes)) - } else { - self.compute_discriminant(ty, bytes)? - }; + let result = self.compute_discriminant(ty, bytes)?; Owned(result.to_le_bytes().to_vec()) } Rvalue::Repeat(it, len) => { diff --git a/crates/hir-ty/src/mir/lower.rs b/crates/hir-ty/src/mir/lower.rs index 1f9bc392115a..7358c14af764 100644 --- a/crates/hir-ty/src/mir/lower.rs +++ b/crates/hir-ty/src/mir/lower.rs @@ -977,40 +977,45 @@ impl<'a, 'db> MirLowerCtx<'a, 'db> { } else { let source_ty = self.infer.expr_ty(*expr); let target_ty = self.infer.expr_ty(expr_id); - let (it, source_ty, current) = if let TyKind::Adt(adt, _) = source_ty.kind() - && adt.is_enum() - { - let Some((enum_place, current)) = - self.lower_expr_as_place(current, *expr, true)? - else { - return Ok(None); - }; - let discr_ty = Ty::new_int(self.interner(), rustc_type_ir::IntTy::I128); - let discr_place: Place<'db> = - self.temp(discr_ty, current, expr_id.into())?.into(); + let (it, source_ty, current) = + if let Some(VariantId::EnumVariantId(variant_id)) = + self.infer.variant_resolution_for_expr(*expr) + { + self.lower_variant_discriminant(current, variant_id)? + } else if let TyKind::Adt(adt, _) = source_ty.kind() + && adt.is_enum() + { + let Some((enum_place, current)) = + self.lower_expr_as_place(current, *expr, true)? + else { + return Ok(None); + }; + let discr_ty = Ty::new_int(self.interner(), rustc_type_ir::IntTy::I128); + let discr_place: Place<'db> = + self.temp(discr_ty, current, expr_id.into())?.into(); - self.push_assignment( - current, - discr_place, - Rvalue::Discriminant(enum_place.store()), - expr_id.into(), - ); - ( - Operand { - kind: OperandKind::Copy(discr_place.store()), - span: Some(expr_id.into()), - }, - discr_ty, - current, - ) - } else { - let Some((it, current)) = - self.lower_expr_to_some_operand(*expr, current)? - else { - return Ok(None); + self.push_assignment( + current, + discr_place, + Rvalue::Discriminant(enum_place.store()), + expr_id.into(), + ); + ( + Operand { + kind: OperandKind::Copy(discr_place.store()), + span: Some(expr_id.into()), + }, + discr_ty, + current, + ) + } else { + let Some((it, current)) = + self.lower_expr_to_some_operand(*expr, current)? + else { + return Ok(None); + }; + (it, source_ty, current) }; - (it, source_ty, current) - }; let cast_kind = if source_ty.as_reference().is_some() { CastKind::PointerCoercion(PointerCast::ArrayToPointer) } else { @@ -1553,6 +1558,17 @@ impl<'a, 'db> MirLowerCtx<'a, 'db> { Ok(prev_block) } + fn lower_variant_discriminant( + &mut self, + current: BasicBlockId, + variant_id: EnumVariantId, + ) -> Result<'db, (Operand, Ty<'db>, BasicBlockId)> { + let discriminant = self.const_eval_discriminant(variant_id)?; + let discr_ty = Ty::new_int(self.interner(), rustc_type_ir::IntTy::I128); + let operand = Operand::from_bytes(Box::new(discriminant.to_le_bytes()), discr_ty); + Ok((operand, discr_ty, current)) + } + fn lower_call_and_args( &mut self, func: Operand,