"""
Ring-LWE-Based Single-Input Inner-Product Functional Encryption
| Based on
| "Efficient Lattice-Based Inner-Product Functional Encryption"
| By Bermudo Mera, Karmakar, Marc, Soleimanian
| ePrint: 2021/046
|
| This module follows the ring-LWE inner-product FE line and is additionally
| informed by the GoFE reference implementation in `innerprod/simple/ringlwe.go`.
* type: public-key encryption
* setting: Integer based
* note: research prototype with SIMD-style matrix encryption
"""
from __future__ import annotations
import json
import logging
import math
import os
import gmpy2 as gp
from pyfe4ai.schemes.ipfe import IPFEAbsCrypto
from pyfe4ai.schemes.ipfe import IPFEAbsKeyGenerator
from pyfe4ai.schemes.ipfe import ParameterCacheMixin
from pyfe4ai.utils.crypto_constants import CryptoCONST
from pyfe4ai.utils.ring_lwe_utils import (
center_matrix,
decode_vector,
mat_vec_mul,
matrix_check_bound,
next_ntt_prime,
poly_add,
poly_mul,
poly_mul_negacyclic,
poly_neg,
transpose,
)
from pyfe4ai.utils.sampling_utils import discrete_gaussian_matrix
from pyfe4ai.utils.sampling_utils import rand_uniform_vector
from pyfe4ai.utils.exceptions import FEKeyError, FEValidationError
logger = logging.getLogger(__name__)
# Backward-compatible aliases (private names used by internal callers)
_poly_add = poly_add
_poly_neg = poly_neg
_poly_mul_negacyclic = poly_mul_negacyclic
_poly_mul = poly_mul
_matrix_check_bound = matrix_check_bound
_transpose = transpose
_mat_vec_mul = mat_vec_mul
_center_matrix = center_matrix
_decode_vector = decode_vector
[docs]
class SIFERingLWEKeyGenerator(IPFEAbsKeyGenerator, ParameterCacheMixin):
"""Key generator for Ring-LWE-based single-input inner-product FE."""
_scheme_type = CryptoCONST.TYPE_SIFE_RING_LWE
[docs]
def __init__(self, config: dict, **kwargs) -> None:
"""Perform the __init__ operation.
Args:
config: Scheme configuration dict.
"""
super().__init__(config, **kwargs)
self.eta = config.get("eta", CryptoCONST.SIFE_DEFAULT_ETA)
self.ring_n = config.get("ring_n", None)
self.bound_x = gp.mpz(config.get("bound_x", 4))
self.bound_y = gp.mpz(config.get("bound_y", 4))
self._load_parameters()
def _apply_parameters(self, param: dict) -> None:
self.p = gp.mpz(param["p"])
self.q = gp.mpz(param["q"])
self.ring_n = int(param["ring_n"])
self.sigma1 = float(param["sigma1"])
self.sigma2 = float(param["sigma2"])
self.sigma3 = float(param["sigma3"])
self.A = [gp.mpz(v) for v in param["A"]]
def _param_verification(self, param: dict) -> bool:
return (
param.get("sec_param") == self.sec_param
and param.get("eta") == self.eta
and param.get("bound_x") == gp.digits(self.bound_x)
and param.get("bound_y") == gp.digits(self.bound_y)
and (self.ring_n is None or param.get("ring_n") == self.ring_n)
)
def _generate_and_save(self, param_file: str) -> None:
l = self.eta
kappa = float(self.sec_param)
sigma = 1.0
sigma1 = math.sqrt(float(4 * l)) * sigma * float(self.bound_x)
kappa_sqrt = math.sqrt(kappa)
p = gp.mpz(self.bound_x * self.bound_y * gp.mpz(2 * l))
if self.ring_n is None:
for pow_exp in range(5, 11):
ring_n = 1 << pow_exp
sigma2 = math.sqrt(float(2 * (l + 2) * ring_n * ring_n)) * sigma
sigma2 *= sigma1 * kappa_sqrt
sigma3 = sigma2 * math.sqrt(2.0)
q_float = sigma1 * sigma2 * kappa * float(2 * ring_n)
q_float += kappa_sqrt * sigma3
q_float *= float(self.bound_y) * float(2 * l)
q = gp.mpz(int(q_float)) * p
if q > 0:
self.ring_n = ring_n
break
sigma2 = math.sqrt(float(2 * (l + 2) * self.ring_n * self.ring_n)) * sigma
sigma2 *= sigma1 * kappa_sqrt
sigma3 = sigma2 * math.sqrt(2.0)
q_float = sigma1 * sigma2 * kappa * float(2 * self.ring_n)
q_float += kappa_sqrt * sigma3
q_float *= float(self.bound_y) * float(2 * l)
q = gp.mpz(int(q_float)) * p
self.p = p
self.q = next_ntt_prime(q + 1, self.ring_n)
self.sigma1 = sigma1
self.sigma2 = sigma2
self.sigma3 = sigma3
self.A = rand_uniform_vector(self.ring_n, self.q)
with open(param_file, "w", encoding="utf-8") as f:
json.dump(
{
"sec_param": self.sec_param,
"eta": self.eta,
"ring_n": self.ring_n,
"bound_x": gp.digits(self.bound_x),
"bound_y": gp.digits(self.bound_y),
"p": gp.digits(self.p),
"q": gp.digits(self.q),
"sigma1": self.sigma1,
"sigma2": self.sigma2,
"sigma3": self.sigma3,
"A": [gp.digits(v) for v in self.A],
},
f,
)
[docs]
def setup(self) -> None:
sk = discrete_gaussian_matrix(self.eta, self.ring_n, self.sigma1)
noise = discrete_gaussian_matrix(self.eta, self.ring_n, self.sigma1)
pk = []
for i in range(self.eta):
pk_i = _poly_mul(self.A, sk[i], self.q)
pk.append(_poly_add(pk_i, noise[i], self.q))
self.msk = {"sk": sk}
self.mpk = {"pk": pk}
logger.info("SIFE Ring-LWE setup successfully")
[docs]
def get_public_parameters(self) -> dict:
return {
"A": [gp.digits(v) for v in self.A],
"p": gp.digits(self.p),
"q": gp.digits(self.q),
"eta": self.eta,
"ring_n": self.ring_n,
"bound_x": gp.digits(self.bound_x),
"bound_y": gp.digits(self.bound_y),
"sigma1": self.sigma1,
"sigma2": self.sigma2,
"sigma3": self.sigma3,
"sec_param": self.sec_param,
}
[docs]
def get_private_keys(self, nid: str = "nid_default", **kwargs) -> dict | None:
"""Perform the get_private_keys operation.
Args:
nid: Node identifier.
"""
return {"pk": [[gp.digits(v) for v in row] for row in self.mpk["pk"]]}
[docs]
def get_decryption_keys(self, sid: str, **kwargs) -> dict | None:
"""Perform the get_decryption_keys operation.
Args:
sid: Session / decryption-key identifier.
"""
credentials = kwargs.get("credentials", None)
if not credentials:
raise FEKeyError("need credentials for SIFE Ring-LWE decryption key generation")
fusion_weight = credentials.get("fusion_weight")
if not isinstance(fusion_weight, list):
raise FEValidationError("invalid fusion weights provided, need a list")
if len(fusion_weight) != self.eta:
raise FEValidationError("invalid fusion weights provided, length mismatch")
if any(abs(int(v)) > self.bound_y for v in fusion_weight):
raise FEValidationError("fusion weight exceeds configured bound_y")
sk_t = _transpose(self.msk["sk"])
y_vec = [gp.mpz(v) for v in fusion_weight]
sk_y = _mat_vec_mul(sk_t, y_vec, self.q)
return {"sk_y": [gp.digits(v) for v in sk_y]}
[docs]
class SIFERingLWE(IPFEAbsCrypto):
"""Crypto operations for Ring-LWE-based single-input inner-product FE."""
[docs]
def __init__(self, config: dict, **kwargs) -> None:
"""Perform the __init__ operation.
Args:
config: Scheme configuration dict.
"""
super().__init__(config, **kwargs)
self.pp = self.keys["pp"]
self.A = [gp.mpz(v) for v in self.pp["A"]]
if self._has_private_keys():
self.sk = self.keys["sk"]
[docs]
def encrypt(self, matrix_pt: list[list[int]]) -> dict:
if not self._has_public_parameters():
raise FEKeyError("no public parameters provided for encryption")
if not self._has_private_keys():
raise FEKeyError("no encryption key provided")
if not isinstance(matrix_pt, list) or not matrix_pt or not isinstance(matrix_pt[0], list):
raise FEValidationError("plaintext matrix must be a list of rows")
if len(matrix_pt) != self.pp["eta"]:
raise FEValidationError("invalid plaintext row count")
if len(matrix_pt[0]) > self.pp["ring_n"]:
raise FEValidationError("plaintext column count exceeds ring dimension")
if not _matrix_check_bound(matrix_pt, gp.mpz(self.pp["bound_x"])):
raise FEValidationError("plaintext exceeds configured bound_x")
q = gp.mpz(self.pp["q"])
pk = [[gp.mpz(v) for v in row] for row in self.sk["pk"]]
ring_n = self.pp["ring_n"]
r = discrete_gaussian_matrix(1, ring_n, self.pp["sigma2"])[0]
noise = discrete_gaussian_matrix(self.pp["eta"], ring_n, self.pp["sigma3"])
ct0 = []
for i in range(self.pp["eta"]):
ct0_i = _poly_mul(pk[i], r, q)
ct0.append(_poly_add(ct0_i, noise[i], q))
centered = _center_matrix(matrix_pt, gp.mpz(self.pp["p"]), q, ring_n)
ct0 = [_poly_add(a, b, q) for a, b in zip(ct0, centered)]
ct1 = _poly_mul(self.A, r, q)
e = discrete_gaussian_matrix(1, ring_n, self.pp["sigma2"])[0]
ct1 = _poly_add(ct1, e, q)
return {
"ct0": [[gp.digits(v) for v in row] for row in ct0],
"ct1": [gp.digits(v) for v in ct1],
"k": len(matrix_pt[0]),
}
[docs]
def decrypt(self, dct_ct: dict, dk: dict, fusion_weight: list):
"""Decrypt ciphertexts and recover the inner product.
Args:
dct_ct: Ciphertext dict (or dict of per-client ciphertexts).
dk: Functional decryption key.
fusion_weight: Fusion weight vector (or dict of per-client weight vectors).
"""
if not dk:
raise FEKeyError("no decryption key provided")
if len(fusion_weight) != self.pp["eta"]:
raise FEValidationError("invalid fusion weight length")
if any(abs(int(v)) > gp.mpz(self.pp["bound_y"]) for v in fusion_weight):
raise FEValidationError("fusion weight exceeds configured bound_y")
q = gp.mpz(self.pp["q"])
p = gp.mpz(self.pp["p"])
ct0 = [[gp.mpz(v) for v in row] for row in dct_ct["ct0"]]
ct1 = [gp.mpz(v) for v in dct_ct["ct1"]]
sk_y = [gp.mpz(v) for v in dk["sk_y"]]
y_vec = [gp.mpz(v) for v in fusion_weight]
ct0_t = _transpose(ct0)
lhs = _mat_vec_mul(ct0_t, y_vec, q)
rhs = _poly_mul(ct1, sk_y, q)
d = _poly_add(lhs, _poly_neg(rhs, q), q)
return _decode_vector(d, p, q, int(dct_ct["k"]))
[docs]
def encrypt_lst_ndarray(self, lst_ndarray: list, **kwargs) -> list | None:
"""Perform the encrypt_lst_ndarray operation.
Args:
lst_ndarray: List of numpy arrays to encrypt element-wise.
"""
raise NotImplementedError()
[docs]
def compute_lst_ndarray_ct(self, dict_ndarray_ct: dict, **kwargs) -> list | None:
"""Compute inner products on encrypted ndarray ciphertexts.
Args:
dict_ndarray_ct: Encrypted ndarray ciphertexts (list or dict).
"""
raise NotImplementedError()
[docs]
def decrypt_lst_ndarray_ct(self, dict_ndarray_ct: dict, **kwargs) -> list | None:
"""Decrypt encrypted ndarray ciphertexts element-wise.
Args:
dict_ndarray_ct: Encrypted ndarray ciphertexts (list or dict).
"""
raise NotImplementedError()