mpc4j--psi同态加密中的参数(二)
=========书接上文=========
感悟,最近被大如传迷了眼,喜欢奇迹婉婉,嬿婉语录:我还年轻,慢慢学~
(猫咪前天看水獭超喜欢,上次去野生动物园看水獭表演,水獭真的很小只)
重要参数
在整个算法流程中有几个参数尤为重要
private Cmg21KwPirParams(CuckooHashBinType cuckooHashBinType, int binNum, int maxPartitionSizePerBin,
int itemEncodedSlotSize, int psLowDegree, int[] queryPowers,
long plainModulus, int polyModulusDegree, int[] coeffModulusBits,
int expectServerSize, int maxRetrievalSize) {
this.cuckooHashBinType = cuckooHashBinType;
this.binNum = binNum;
this.maxPartitionSizePerBin = maxPartitionSizePerBin;
this.itemEncodedSlotSize = itemEncodedSlotSize;
this.psLowDegree = psLowDegree;
this.queryPowers = queryPowers;
this.plainModulus = plainModulus;
this.polyModulusDegree = polyModulusDegree;
this.coeffModulusBits = coeffModulusBits;
this.expectServerSize = expectServerSize;
this.maxRetrievalSize = maxRetrievalSize;
this.itemPerCiphertext = polyModulusDegree / itemEncodedSlotSize;
this.ciphertextNum = binNum / itemPerCiphertext;
this.encryptionParams = Cmg21KwPirNativeUtils.genEncryptionParameters(polyModulusDegree, plainModulus, coeffModulusBits);
}
- 参数列表:
- cuckooHashBinType:指代一种哈希桶类型。
- binNum:哈希桶的数量。
- maxPartitionSizePerBin:每个哈希桶的最大分区大小。
- itemEncodedSlotSize:编码槽的大小。
- psLowDegree:多项式的低阶项。
- queryPowers:查询的幂次列表。
- plainModulus:明文模数。
- polyModulusDegree:多项式模数的度。
- coeffModulusBits:系数模数的位数。
- expectServerSize:预计的服务器大小。
- maxRetrievalSize:最大检索大小。
poly_modulus_degree is a positive power-of-two integer that determines how many integers modulo plain_modulus can be encoded into a single Microsoft SEAL plaintext。翻译一下就是一个明文中的整数模数。
举例分析
这里每个参数的关系是什么样的呢?
16M-4096为例
参数列表
- hash_bin_params: {
- cuckoo_hash_bin_type : NAIVE_3_HASH
- bin_num : 6552
- max_items_per_bin : 1304
}
- item_params: {
- felts_per_item : 5
}
- query_params: {
- ps_low_degree : 44
- query_powers : [1, 3, 11, 18, 45, 225]
}
- seal_params: {
- plain_modulus : 4079617
- poly_modulus_degree : 8192
- coeff_modulus_bits : [56, 56, 56, 50]
}
2024-05-09 04:07:58,798 [main] INFO edu.alibaba.mpc4j.s2pc.main.kwpir.KwPirMain - server: serverSetSize = 2097152, elementBitLength = 64, queryNumber = 2, parallel = true
itemPerCiphertext = polyModulusDegree / itemEncodedSlotSize;
1638=8192/5 这行代码将多项式模数的度数除以每个槽的大小,以确定每个密文包含多少个项。这里的 polyModulusDegree 是多项式模数的度数,itemEncodedSlotSize 指的就是一个item要拆成几个明文。这个值决定了每个密文中包含的项的数量。
ciphertextNum = binNum / itemPerCiphertext
4=6552/1638 可用的 bin 数量除以每个密文的项数,以确定需要多少个密文来存储所有的 bin。这个值决定了在加密数据库时需要生成多少个密文。
总结一下:这个概念就是比如一个密文有10项,如果有100个桶,则需要10个密文才能完成存储。
int partitionCount = CommonUtils.getUnitNum(binSize, params.getMaxPartitionSizePerBin()); 表示每个哈希桶的最大分区大小3
partitionCount=binSize % params.getMaxPartitionSizePerBin()可以解读为 3912/1304 3
itemPerCiphertext size: 1638
itemEncodedSlotSize size: 5
binSize: 3912
partitionCount size: 3
bigPartitionCount size: 3
labelPartitionCount size: 1
partitionSize: 1304
重点:binSize 和binNum可不一样,binNum是桶的数量,binSize是每个桶的binItem的数量
partitionSize和partitionCount 也要注意区分。
补充知识
在同态加密方案中,原始数据通常被表示为多项式的系数。加密操作会将原始数据转换为多项式,并使用多项式模数来对其进行加密。解密操作会使用密文多项式和密钥来恢复原始数据。
server端初始化
3.encodeDatabase
所以这个方法就是生成多项式的系数以及将这些系数编码成向量的过程。(不理解的看补充知识)
原方法有点长,我提取了三个方法。
private void encodeDatabase(Map<ByteBuffer, ByteBuffer> prfMap, List<List<HashBinEntry<ByteBuffer>>> hashBins) {
Zp64Poly zp64Poly = Zp64PolyFactory.createInstance(envType, params.getPlainModulus());
int itemPerCiphertext = params.getItemPerCiphertext();
int itemEncodedSlotSize = params.getItemEncodedSlotSize();
int partitionCount = CommonUtils.getUnitNum(binSize, params.getMaxPartitionSizePerBin());
int bigPartitionCount = binSize / params.getMaxPartitionSizePerBin();
int labelPartitionCount = CommonUtils.getUnitNum((labelByteLength + ivByteLength) * Byte.SIZE,
(PirUtils.getBitLength(params.getPlainModulus()) - 1) * itemEncodedSlotSize);
serverKeywordEncode = new ArrayList<>();
serverLabelEncode = new ArrayList<>();
// for each bucket, compute the coefficients of the polynomial f(x) = \prod_{y in bucket} (x - y)
// and coeffs of g(x), which has the property g(y) = label(y) for each y in bucket.
// ciphertext num is small, therefore we need to do parallel computation inside the loop
for (int i = 0; i < params.getCiphertextNum(); i++) {
int finalIndex = i;
for (int partition = 0; partition < partitionCount; partition++) {
// 计算分区大小和起始位置
int partitionSize, partitionStart;
partitionSize = partition < bigPartitionCount ?
params.getMaxPartitionSizePerBin() : binSize % params.getMaxPartitionSizePerBin();
partitionStart = params.getMaxPartitionSizePerBin() * partition;
System.out.println("finalIndex: " + finalIndex + "--partitionSize: " + partitionSize);
System.out.println("finalIndex: " + finalIndex + "--partitionStart: " + partitionStart);
// 计算关键词和标签的系数
long[][] fCoeffs = new long[itemPerCiphertext * itemEncodedSlotSize][];
long[][][] gCoeffs = new long[labelPartitionCount][itemPerCiphertext * itemEncodedSlotSize][];
// 计算关键词和标签的系数
computeCoefficients(fCoeffs, gCoeffs, hashBins, prfMap, finalIndex, partitionSize, partitionStart, itemPerCiphertext, itemEncodedSlotSize,labelPartitionCount, zp64Poly);
long[][] encodeElementVector = encodeElementVector(fCoeffs, partitionSize, itemEncodedSlotSize, itemPerCiphertext);
long[][][] encodeLabelVector = encodeLabelVector(gCoeffs, itemEncodedSlotSize, itemPerCiphertext, labelPartitionCount);
serverKeywordEncode.add(Cmg21KwPirNativeUtils.preprocessDatabase(
params.getEncryptionParams(), encodeElementVector, params.getPsLowDegree())
);
for (int j = 0; j < labelPartitionCount; j++) {
serverLabelEncode.add(Cmg21KwPirNativeUtils.preprocessDatabase(
params.getEncryptionParams(), encodeLabelVector[j], params.getPsLowDegree()
));
}
}
}
}
// 计算关键词和标签的系数
private void computeCoefficients(long[][] fCoeffs, long[][][] gCoeffs ,List<List<HashBinEntry<ByteBuffer>>> hashBins, Map<ByteBuffer, ByteBuffer> prfMap,
int finalIndex, int partitionSize, int partitionStart,
int itemPerCiphertext, int itemEncodedSlotSize, int labelPartitionCount, Zp64Poly zp64Poly){
IntStream itemIndexStream = IntStream.range(0, itemPerCiphertext);
itemIndexStream = parallel ? itemIndexStream.parallel() : itemIndexStream;
itemIndexStream.forEach(j -> {
long[][] currentBucketElement = new long[itemEncodedSlotSize][partitionSize];
long[][][] currentBucketLabels = new long[labelPartitionCount][itemEncodedSlotSize][partitionSize];
for (int l = 0; l < partitionSize; l++) {
HashBinEntry<ByteBuffer> entry = hashBins.get(
finalIndex * itemPerCiphertext + j).get(partitionStart + l
);
long[] temp = params.getHashBinEntryEncodedArray(entry, false, secureRandom);
for (int k = 0; k < itemEncodedSlotSize; k++) {
currentBucketElement[k][l] = temp[k];
}
}
for (int l = 0; l < itemEncodedSlotSize; l++) {
fCoeffs[j * itemEncodedSlotSize + l] = zp64Poly.rootInterpolate(
partitionSize, currentBucketElement[l], 0L
);
}
int nonEmptyBuckets = 0;
for (int l = 0; l < partitionSize; l++) {
HashBinEntry<ByteBuffer> entry = hashBins.get(
finalIndex * itemPerCiphertext + j).get(partitionStart + l
);
if (entry.getHashIndex() != HashBinEntry.DUMMY_ITEM_HASH_INDEX) {
for (int k = 0; k < itemEncodedSlotSize; k++) {
currentBucketElement[k][nonEmptyBuckets] = currentBucketElement[k][l];
}
byte[] oprf = entry.getItem().array();
// choose first 128 bits
byte[] keyBytes = new byte[CommonConstants.BLOCK_BYTE_LENGTH];
System.arraycopy(oprf, 0, keyBytes, 0, CommonConstants.BLOCK_BYTE_LENGTH);
byte[] iv = new byte[ivByteLength];
secureRandom.nextBytes(iv);
byte[] extendedIv = BytesUtils.paddingByteArray(iv, CommonConstants.BLOCK_BYTE_LENGTH);
byte[] plaintextLabel = prfMap.get(entry.getItem()).array();
byte[] extendedCipherLabel = streamCipher.ivEncrypt(keyBytes, extendedIv, plaintextLabel);
byte[] ciphertextLabel = new byte[ivByteLength + labelByteLength];
System.arraycopy(extendedCipherLabel, CommonConstants.BLOCK_BYTE_LENGTH - ivByteLength, ciphertextLabel, 0, ivByteLength + labelByteLength);
long[][] temp = params.encodeLabel(ciphertextLabel, labelPartitionCount);
for (int k = 0; k < labelPartitionCount; k++) {
for (int h = 0; h < itemEncodedSlotSize; h++) {
currentBucketLabels[k][h][nonEmptyBuckets] = temp[k][h];
}
}
nonEmptyBuckets++;
}
}
for (int l = 0; l < itemEncodedSlotSize; l++) {
for (int k = 0; k < labelPartitionCount; k++) {
long[] xArray = new long[nonEmptyBuckets];
long[] yArray = new long[nonEmptyBuckets];
for (int index = 0; index < nonEmptyBuckets; index++) {
xArray[index] = currentBucketElement[l][index];
yArray[index] = currentBucketLabels[k][l][index];
}
if (nonEmptyBuckets > 0) {
gCoeffs[k][j * itemEncodedSlotSize + l] = zp64Poly.interpolate(
nonEmptyBuckets, xArray, yArray
);
} else {
gCoeffs[k][j * itemEncodedSlotSize + l] = new long[0];
}
}
}
});
}
// 将关键词系数编码为向量
private long[][] encodeElementVector(long[][] fCoeffs, int partitionSize, int itemEncodedSlotSize, int itemPerCiphertext) {
long[][] encodeElementVector = new long[partitionSize + 1][params.getPolyModulusDegree()];
for (int j = 0; j < partitionSize + 1; j++) {
// encode the jth coefficients of all polynomials into a vector
for (int l = 0; l < itemEncodedSlotSize * itemPerCiphertext; l++) {
encodeElementVector[j][l] = fCoeffs[l][j];
}
for (int l = itemEncodedSlotSize * itemPerCiphertext; l < params.getPolyModulusDegree(); l++) {
encodeElementVector[j][l] = 0;
}
}
return encodeElementVector;
}
private long[][][] encodeLabelVector(long[][][] gCoeffs, int itemEncodedSlotSize, int itemPerCiphertext, int labelPartitionCount) {
int labelSize = IntStream.range(1, gCoeffs[0].length)
.map(j -> gCoeffs[0][j].length).filter(j -> j >= 1)
.max()
.orElse(1);
long[][][] encodeLabelVector = new long[labelPartitionCount][labelSize][params.getPolyModulusDegree()];
for (int j = 0; j < labelSize; j++) {
for (int k = 0; k < labelPartitionCount; k++) {
for (int l = 0; l < itemEncodedSlotSize * itemPerCiphertext; l++) {
if (gCoeffs[k][l].length == 0) {
encodeLabelVector[k][j][l] = Math.abs(secureRandom.nextLong()) % params.getPlainModulus();
} else {
encodeLabelVector[k][j][l] = (j < gCoeffs[k][l].length) ? gCoeffs[k][l][j] : 0;
}
}
for (int l = itemEncodedSlotSize * itemPerCiphertext; l < params.getPolyModulusDegree(); l++) {
encodeLabelVector[k][j][l] = 0;
}
}
}
return encodeLabelVector;
}
encodeDatabase 方法
- 此方法是生成多项式系数并编码为向量的入口点。它初始化了一些变量,如
zp64Poly、itemPerCiphertext、itemEncodedSlotSize、partitionCount、bigPartitionCount和labelPartitionCount。然后,它迭代处理每个密文,并计算每个分区的多项式系数,最后编码为向量并存储起来。 - 密文有4个,每个密文有3个分区,所以serverKeywordEncode返回是12个list,每个list包含1305个byte[]
computeCoefficients 方法
此方法计算多项式的系数,包括关键词和标签的系数。它迭代处理每个密文中的每个分区,计算出每个分区的多项式系数。这些系数用于构建多项式 f(x) 和 g(x)。
先遍历每个密文中的item,之前就说了每个item可以分解为多个明文,所以对于每个Item,currentBucketElement实际上是一个itemEncodedSlotSize*partitionSize(5*1304)的小矩阵,因此需要把哈希桶中取到的元素编译成数组,桶中取元素是这样的,先取是哪个桶,再取桶中的哪个分区的哪个值,把这个entry编译成5个明文的形式(实际上只取了一个元素), 这里分析一下差值多项式的组成:多项式的每行是这个临时桶的每行构造的差值多项式。
理论超级重要
我一直不清楚问什么在currentBucketElement中为什么要分成多个slot,结合APSI中的描述来看一下这个过程,如图所示。

第二个技巧是将每个项目分解为多个部分,然后将它们分别编码到连续的批处理槽中。换句话说,如果 plain_modulus 是一个 B-位质数,那么我们只将一个项目的 B - 1 位写入一个批处理槽,下一个 B - 1 位写入下一个槽。这种方法的一个缺点是,一个批处理的明文/密文现在只能容纳一部分 poly_modulus_degree 项目。例如,如果 plain_modulus 是一个 21-位质数,那么 4 个槽可以编码一个长度为 80 的项目,查询密文 Q (及其幂) 可以加密多达 poly_modulus_degree / 4 的接收者的项目。接收者现在的查询物品比以前少得多。解决方案是将杜鹃哈希表的大小与 poly_modulus_degree 解耦,简单地使用两个或更多的密文来加密 {X_i} (及其幂)。
现在来看看getHashBinEntryEncodedArray这个方法,
- plainModulus:22
- bitLength: 105
- shiftBits: 21
- shiftMask: 2097151
public long[] getHashBinEntryEncodedArray(HashBinEntry<ByteBuffer> hashBinEntry, boolean isReceiver,
SecureRandom secureRandom) {
long[] encodedArray = new long[itemEncodedSlotSize];
int bitLength = (BigInteger.valueOf(plainModulus).bitLength() - 1) * itemEncodedSlotSize;
assert bitLength >= 80;
int shiftBits = BigInteger.valueOf(plainModulus).bitLength() - 1;
BigInteger shiftMask = BigInteger.ONE.shiftLeft(shiftBits).subtract(BigInteger.ONE);
System.out.println("shiftMask: "+shiftMask);
if (hashBinEntry.getHashIndex() != -1) {
assert (hashBinEntry.getHashIndex() < 3) : "hash index should be [0, 1, 2]";
BigInteger input = BigIntegerUtils.byteArrayToNonNegBigInteger(hashBinEntry.getItem().array());
input = input.mod(BigInteger.ONE.shiftLeft(CommonConstants.BLOCK_BIT_LENGTH));
for (int i = 0; i < itemEncodedSlotSize; i++) {
encodedArray[i] = input.and(shiftMask).longValueExact();
input = input.shiftRight(shiftBits);
}
} else {
IntStream.range(0, itemEncodedSlotSize).forEach(i -> {
long random = Math.abs(secureRandom.nextLong()) % plainModulus / 4;
encodedArray[i] = random << 1 | (isReceiver ? 1L : 0L);
});
}
for (int i = 0; i < itemEncodedSlotSize; i++) {
assert encodedArray[i] < plainModulus;
}
return encodedArray;
}

现在宏观理解一下这个函数的过程
// 计算关键词和标签的系数
private void computeCoefficients(long[][] fCoeffs, long[][][] gCoeffs ,List<List<HashBinEntry<ByteBuffer>>> hashBins, Map<ByteBuffer, ByteBuffer> prfMap,
int finalIndex, int partitionSize, int partitionStart,
int itemPerCiphertext, int itemEncodedSlotSize, int labelPartitionCount, Zp64Poly zp64Poly){
IntStream itemIndexStream = IntStream.range(0, itemPerCiphertext);
itemIndexStream = parallel ? itemIndexStream.parallel() : itemIndexStream;
itemIndexStream.forEach(j -> {
long[][] currentBucketElement = new long[itemEncodedSlotSize][partitionSize];
long[][][] currentBucketLabels = new long[labelPartitionCount][itemEncodedSlotSize][partitionSize];
for (int l = 0; l < partitionSize; l++) {
HashBinEntry<ByteBuffer> entry = hashBins.get(
finalIndex * itemPerCiphertext + j).get(partitionStart + l
);
System.out.println("finalIndex: "+finalIndex+"--itemIndex: "+j+"--entry["+(finalIndex * itemPerCiphertext + j)+"]"+"--partition: "+(partitionStart + l));
long[] temp = params.getHashBinEntryEncodedArray(entry, false, secureRandom);
for (int k = 0; k < itemEncodedSlotSize; k++) {
currentBucketElement[k][l] = temp[k];
System.out.println("finalIndex: "+finalIndex+"--itemIndex: "+j+"--currentBucketElement["+k+"]["+l+"]"+"temp["+k+"]");
}
}
for (int l = 0; l < itemEncodedSlotSize; l++) {
fCoeffs[j * itemEncodedSlotSize + l] = zp64Poly.rootInterpolate(
partitionSize, currentBucketElement[l], 0L
);
System.out.println("finalIndex: "+finalIndex+"--itemIndex: "+j+"--fCoeffs["+(j * itemEncodedSlotSize + l)+"]"+"--currentBucketElement["+l+"]");
}
int nonEmptyBuckets = 0;
for (int l = 0; l < partitionSize; l++) {
HashBinEntry<ByteBuffer> entry = hashBins.get(
finalIndex * itemPerCiphertext + j).get(partitionStart + l
);
if (entry.getHashIndex() != HashBinEntry.DUMMY_ITEM_HASH_INDEX) {
for (int k = 0; k < itemEncodedSlotSize; k++) {
currentBucketElement[k][nonEmptyBuckets] = currentBucketElement[k][l];
}
byte[] oprf = entry.getItem().array();
// choose first 128 bits
byte[] keyBytes = new byte[CommonConstants.BLOCK_BYTE_LENGTH];
System.arraycopy(oprf, 0, keyBytes, 0, CommonConstants.BLOCK_BYTE_LENGTH);
byte[] iv = new byte[ivByteLength];
secureRandom.nextBytes(iv);
byte[] extendedIv = BytesUtils.paddingByteArray(iv, CommonConstants.BLOCK_BYTE_LENGTH);
byte[] plaintextLabel = prfMap.get(entry.getItem()).array();
byte[] extendedCipherLabel = streamCipher.ivEncrypt(keyBytes, extendedIv, plaintextLabel);
byte[] ciphertextLabel = new byte[ivByteLength + labelByteLength];
System.arraycopy(extendedCipherLabel, CommonConstants.BLOCK_BYTE_LENGTH - ivByteLength, ciphertextLabel, 0, ivByteLength + labelByteLength);
long[][] temp = params.encodeLabel(ciphertextLabel, labelPartitionCount);
for (int k = 0; k < labelPartitionCount; k++) {
for (int h = 0; h < itemEncodedSlotSize; h++) {
currentBucketLabels[k][h][nonEmptyBuckets] = temp[k][h];
}
}
nonEmptyBuckets++;
}
}
for (int l = 0; l < itemEncodedSlotSize; l++) {
for (int k = 0; k < labelPartitionCount; k++) {
long[] xArray = new long[nonEmptyBuckets];
long[] yArray = new long[nonEmptyBuckets];
for (int index = 0; index < nonEmptyBuckets; index++) {
xArray[index] = currentBucketElement[l][index];
yArray[index] = currentBucketLabels[k][l][index];
}
if (nonEmptyBuckets > 0) {
gCoeffs[k][j * itemEncodedSlotSize + l] = zp64Poly.interpolate(
nonEmptyBuckets, xArray, yArray
);
} else {
gCoeffs[k][j * itemEncodedSlotSize + l] = new long[0];
}
}
}
});
}
finalIndex: 0--itemIndex: 440--entry[440]--partition: 13
finalIndex: 0--itemIndex: 440--currentBucketElement[0][13]temp[0]
finalIndex: 0--itemIndex: 440--currentBucketElement[1][13]temp[1]
finalIndex: 0--itemIndex: 440--currentBucketElement[2][13]temp[2]
finalIndex: 0--itemIndex: 440--currentBucketElement[3][13]temp[3]
finalIndex: 0--itemIndex: 440--currentBucketElement[4][13]temp[4]finalIndex: 0--itemIndex: 440--fCoeffs[2200]--currentBucketElement[0]
finalIndex: 0--itemIndex: 440--fCoeffs[2201]--currentBucketElement[1]
finalIndex: 0--itemIndex: 440--fCoeffs[2202]--currentBucketElement[2]
finalIndex: 0--itemIndex: 440--fCoeffs[2203]--currentBucketElement[3]
finalIndex: 0--itemIndex: 440--fCoeffs[2204]--currentBucketElement[4]
finalIndex是密文的索引,根据这个打印举例一下是第0个密文中第440个项,从hashBins中取第440的桶中第13个数据temp。将temp用 params.getHashBinEntryEncodedArray分解成5份,然后每个临时桶的第13列就是对应的temp[]中的元素。
之前有说过,总体是4个密文。每个密文有1638项,这里是每一项取1304个桶中的数据,并且把每个数据分成5份存储到临时桶中,所以每个临时桶有5行、1304列,填充完这个临时桶的所有行和列。进行多项式差值运算,得到差值多项式系数,对于每一项构建5个多项式。
fCoeffs[2200]=fCoeffs[5*440]
zp64Poly.rootInterpolate构建多项式的时候,一个多项式有1304个差值点,多项式的X数组是 currentBucketElement[l],也就是临时桶的每行。

encodeElementVector 方法
此方法将关键词的系数编码为向量。它创建一个二维数组,将多项式的每个系数组合成一个向量,并在适当的位置进行填充。这些向量被用于构建关键词的加密数据。
finalIndex: 0--partition: 0--fCoeffs[8190][1305]
encodeElementVector[622][8189]--fCoeffs[8189][622]
初始化encodeElementVector[1305][8192] fCoeffs填充到第8189列,剩下两列填0 encodeElementVector[64][8190]: 0 encodeElementVector[64][8191]: 0
encodeLabelVector 方法
- 此方法将标签的系数编码为向量。它创建一个三维数组,用于存储标签的每个系数,然后将其编码为向量。这些向量被用于构建标签的加密数据。
更多推荐
所有评论(0)