[llvm] [ConstantTime][LLVM] Add llvm.ct.select intrinsic with generic SelectionDAG lowering (PR #166702)

Akshay K via llvm-commits llvm-commits at lists.llvm.org
Thu Sep 10 10:53:19 PDT 2026


================
@@ -4352,6 +4353,188 @@ bool SelectionDAGLegalize::ExpandNode(SDNode *Node) {
     }
     Results.push_back(Tmp1);
     break;
+  case ISD::CT_SELECT: {
+    // Constant-time select: F ^ ((T ^ F) & Mask), Mask = 0 - (cond & 1).
+    // Bitwise-only — no select/cmov SDNode is constructed, so the CT property
+    // holds against any combiner that targets those opcodes. FP types operate
+    // on the same-size integer; vectors build the mask as a scalar then splat
+    // (avoids illegal vNi1).
+    //
+    // The masked-diff is routed through a virtual register (CopyToReg /
+    // CopyFromReg) below as a forward-looking DAGCombine barrier. This is
+    // *not* required for correctness against any combiner in tree today —
+    // DAGCombiner has no rewrite that recognizes XOR/AND/XOR-with-sext-mask
+    // and reconstructs a SELECT. The chain edge is defense-in-depth against
+    // a hypothetical future fold of that form: the dependency partitions the
+    // bitwise sequence into a region the combiner can't see through. Cost is
+    // at most a coalesce-able MOV per call. ISD::ARITH_FENCE would serve the
+    // same purpose (ISel already selects it for any legal type); switching to
+    // it is left to a follow-up.
+    Tmp1 = Node->getOperand(0); // cond
+    Tmp2 = Node->getOperand(1); // T
+    Tmp3 = Node->getOperand(2); // F
+    EVT VT = Tmp2.getValueType();
+
+    // Memory-blend for FP scalars whose same-size integer isn't legal (f64 on
+    // i386 no-SSE, x86_fp80, fp128). The bitcast-to-int expansion below can't
+    // run since LegalizeDAG is the last legalization stage. Spill T/F, blend
+    // chunk-by-chunk at a legal int width in place over the T slot, reload it
+    // as the FP type. Fixed and scalable FP vectors fall through to the
+    // unified path below.
+    if (VT.isFloatingPoint() && !VT.isVector() &&
+        !TLI.isTypeLegal(VT.changeTypeToInteger())) {
+      const DataLayout &DL = DAG.getDataLayout();
+      Type *VTTy = VT.getTypeForEVT(*DAG.getContext());
+      unsigned StorageBytes = DL.getTypeStoreSize(VTTy);
+      assert(StorageBytes > 0 && "FP type with zero storage size");
+
+      // Blend at the widest legal scalar integer width. Chunks narrower than
+      // that (e.g. the 2-byte tail of x86_fp80) are zero-extended on load and
+      // truncated on store, so every value in the DAG has a legal type.
+      MVT BlendVT;
+      for (MVT MV : {MVT::i64, MVT::i32, MVT::i16, MVT::i8})
+        if (TLI.isTypeLegal(MV)) {
+          BlendVT = MV;
+          break;
+        }
+      assert(BlendVT.isValid() && "no legal scalar integer type");
+      unsigned BlendBytes = BlendVT.getSizeInBits() / 8;
+
+      MachineFunction &MF = DAG.getMachineFunction();
+      SDValue StackT = DAG.CreateStackTemporary(VT);
+      SDValue StackF = DAG.CreateStackTemporary(VT);
+      int FIT = cast<FrameIndexSDNode>(StackT.getNode())->getIndex();
+      int FIF = cast<FrameIndexSDNode>(StackF.getNode())->getIndex();
+      MachinePointerInfo PIT = MachinePointerInfo::getFixedStack(MF, FIT);
+      MachinePointerInfo PIF = MachinePointerInfo::getFixedStack(MF, FIF);
+
+      SDValue Chain = DAG.getEntryNode();
+      Chain = DAG.getStore(Chain, dl, Tmp2, StackT, PIT);
+      Chain = DAG.getStore(Chain, dl, Tmp3, StackF, PIF);
+
+      // Walk the storage in power-of-2 chunks, widest first, so every access
+      // stays naturally aligned within the max-aligned stack temporaries.
+      unsigned Offset = 0;
+      while (Offset < StorageBytes) {
+        unsigned ChunkBytes =
+            std::min(BlendBytes, llvm::bit_floor(StorageBytes - Offset));
+        MVT MemVT = MVT::getIntegerVT(ChunkBytes * 8);
+        TypeSize Off = TypeSize::getFixed(Offset);
+        SDValue TPtr = DAG.getMemBasePlusOffset(StackT, Off, dl);
+        SDValue FPtr = DAG.getMemBasePlusOffset(StackF, Off, dl);
+
+        SDValue Ti, Fi;
+        if (MemVT == BlendVT) {
+          Ti = DAG.getLoad(BlendVT, dl, Chain, TPtr, PIT.getWithOffset(Offset));
+          Chain = Ti.getValue(1);
+          Fi = DAG.getLoad(BlendVT, dl, Chain, FPtr, PIF.getWithOffset(Offset));
+        } else {
+          Ti = DAG.getExtLoad(ISD::ZEXTLOAD, dl, BlendVT, Chain, TPtr,
+                              PIT.getWithOffset(Offset), MemVT);
+          Chain = Ti.getValue(1);
+          Fi = DAG.getExtLoad(ISD::ZEXTLOAD, dl, BlendVT, Chain, FPtr,
+                              PIF.getWithOffset(Offset), MemVT);
+        }
+        Chain = Fi.getValue(1);
+
+        // Blend this chunk via CT_SELECT on the legal integer type (the
+        // recursive node is Expand'd by the scalar-int branch below), then
+        // store it back over the now-dead chunk of the T slot. The serial
+        // chain keeps each chunk's loads ahead of the store.
+        SDValue Ri = DAG.getCTSelect(dl, BlendVT, Tmp1, Ti, Fi);
+
+        if (MemVT == BlendVT)
+          Chain = DAG.getStore(Chain, dl, Ri, TPtr, PIT.getWithOffset(Offset));
+        else
+          Chain = DAG.getTruncStore(Chain, dl, Ri, TPtr,
+                                    PIT.getWithOffset(Offset), MemVT);
+        Offset += ChunkBytes;
+      }
+
+      Tmp1 = DAG.getLoad(VT, dl, Chain, StackT, PIT);
+      Results.push_back(Tmp1);
+      break;
+    }
+
+    SDValue WorkingT = Tmp2;
+    SDValue WorkingF = Tmp3;
+    EVT WorkingVT = VT;
+
+    bool IsFP = VT.isVector() ? VT.getVectorElementType().isFloatingPoint()
+                              : VT.isFloatingPoint();
+    if (IsFP) {
+      WorkingVT = VT.changeTypeToInteger();
+      // Scalars with an illegal integer twin took the memory path above; FP
+      // vector types are expected to have a legal integer counterpart.
+      assert(TLI.isTypeLegal(WorkingVT) &&
+             "no legal same-width integer type for CT_SELECT expansion");
+      WorkingT = DAG.getBitcast(WorkingVT, Tmp2);
+      WorkingF = DAG.getBitcast(WorkingVT, Tmp3);
+    }
+
+    // Compute the all-ones/all-zeros mask as a scalar, then splat for vectors.
+    // The element type of a legal vector is not always a legal scalar type
+    // (i32 on riscv64, i64 on riscv32), so build the scalar mask at a legal
+    // width: BUILD_VECTOR and SPLAT_VECTOR implicitly truncate a wider
+    // scalar, and if every legal scalar is narrower than the element, splat
+    // at that width and sign-extend the mask vector (0 and -1 survive both).
+    EVT MaskEltVT =
+        WorkingVT.isVector() ? WorkingVT.getVectorElementType() : WorkingVT;
+    EVT MaskSclVT = MaskEltVT;
+    if (!TLI.isTypeLegal(MaskSclVT)) {
+      assert(WorkingVT.isVector() && "scalar CT_SELECT type must be legal");
+      MaskSclVT = TLI.getLegalTypeToTransformTo(*DAG.getContext(), MaskSclVT);
+      assert(TLI.isTypeLegal(MaskSclVT) && "no legal scalar integer type");
+    }
+    SDValue ScalarCond = Tmp1;
+    if (ScalarCond.getValueType() != MaskSclVT)
+      ScalarCond = DAG.getAnyExtOrTrunc(ScalarCond, dl, MaskSclVT);
+    SDValue ScalarMask =
+        DAG.getNode(ISD::SUB, dl, MaskSclVT, DAG.getConstant(0, dl, MaskSclVT),
+                    DAG.getNode(ISD::AND, dl, MaskSclVT, ScalarCond,
+                                DAG.getConstant(1, dl, MaskSclVT)));
+    SDValue Mask;
+    if (!WorkingVT.isVector()) {
+      Mask = ScalarMask;
+    } else if (MaskSclVT.bitsGE(MaskEltVT)) {
+      Mask = DAG.getSplat(WorkingVT, dl, ScalarMask);
+    } else {
+      EVT NarrowVT =
+          WorkingVT.changeVectorElementType(*DAG.getContext(), MaskSclVT);
+      assert(TLI.isTypeLegal(NarrowVT) &&
+             "no legal vector type for the CT_SELECT mask splat");
+      Mask = DAG.getNode(ISD::SIGN_EXTEND, dl, WorkingVT,
+                         DAG.getSplat(NarrowVT, dl, ScalarMask));
+    }
+
+    // F ^ ((T ^ F) & Mask)
+    SDValue XorTF = DAG.getNode(ISD::XOR, dl, WorkingVT, WorkingT, WorkingF);
+    SDValue TM = DAG.getNode(ISD::AND, dl, WorkingVT, XorTF, Mask);
+
+    // Forward-looking DAGCombine barrier (see header): route the masked-diff
----------------
kumarak wrote:

Switched to `ISD::ARITH_FENCE`, so the legalizer no longer creates vregs or calls `getRegClassFor`. ISel lowers the fence to the existing tied meta pseudo, and InstrEmitter handles register-class selection normally.  Re-generated the test output that includes `#ARITH_FENCE` and a few instances of register shuffling. 

https://github.com/llvm/llvm-project/pull/166702


More information about the llvm-commits mailing list