[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