501 lines
19 KiB
Python
501 lines
19 KiB
Python
# built-in dependencies
|
|
import time
|
|
import json
|
|
from typing import Optional, Union, List
|
|
import multiprocessing
|
|
from contextlib import closing
|
|
import traceback
|
|
import copy
|
|
|
|
# 3rd party dependencies
|
|
from tqdm import tqdm
|
|
|
|
# project dependencies
|
|
from lightphe.models.Homomorphic import Homomorphic
|
|
from lightphe.models.Ciphertext import Ciphertext
|
|
from lightphe.models.Algorithm import Algorithm
|
|
from lightphe.models.Tensor import Fraction, EncryptedTensor
|
|
from lightphe.commons import phe_utils
|
|
from lightphe.commons.key_validator import validate_keys
|
|
from lightphe.commons.logger import Logger
|
|
|
|
# cryptosystems
|
|
from lightphe.cryptosystems.RSA import RSA
|
|
from lightphe.cryptosystems.ElGamal import ElGamal
|
|
from lightphe.cryptosystems.Paillier import Paillier
|
|
from lightphe.cryptosystems.DamgardJurik import DamgardJurik
|
|
from lightphe.cryptosystems.OkamotoUchiyama import OkamotoUchiyama
|
|
from lightphe.cryptosystems.Benaloh import Benaloh
|
|
from lightphe.cryptosystems.NaccacheStern import NaccacheStern
|
|
from lightphe.cryptosystems.GoldwasserMicali import GoldwasserMicali
|
|
from lightphe.cryptosystems.EllipticCurveElGamal import EllipticCurveElGamal
|
|
from lightphe.cryptosystems.SanderYoungYung import SanderYoungYung
|
|
from lightphe.cryptosystems.BonehGohNissim import BonehGohNissim
|
|
|
|
|
|
# pylint: disable=eval-used, simplifiable-if-expression, too-few-public-methods
|
|
|
|
logger = Logger(module="lightphe/__init__.py")
|
|
|
|
VERSION = "0.0.24"
|
|
|
|
|
|
class LightPHE:
|
|
__version__ = VERSION
|
|
|
|
def __init__(
|
|
self,
|
|
algorithm_name: str,
|
|
keys: Optional[dict] = None,
|
|
key_file: Optional[str] = None,
|
|
key_size: Optional[int] = None,
|
|
precision: int = 5,
|
|
form: Optional[str] = None,
|
|
curve: Optional[str] = None,
|
|
plaintext_limit: Optional[int] = None,
|
|
max_tries: int = 10000,
|
|
):
|
|
"""
|
|
Build LightPHE class
|
|
Args:
|
|
algorithm_name (str): RSA | ElGamal | Exponential-ElGamal | EllipticCurve-ElGamal
|
|
| Paillier | Damgard-Jurik | Okamoto-Uchiyama | Benaloh | Naccache-Stern
|
|
| Goldwasser-Micali | Sander-Young-Yung | Boneh-Goh-Nissim
|
|
keys (dict): optional private-public key pair
|
|
key_file (str): if keys were exported already, you can load them into cryptosystem
|
|
key_size (int): key size in bits
|
|
precision (int): precision for homomorphic operations on tensors
|
|
form (str): specifies the form of the elliptic curve.
|
|
Options: 'weierstrass' (default), 'edwards', 'koblitz'.
|
|
This parameter is only used if `algorithm_name` is 'EllipticCurve-ElGamal'.
|
|
curve (str): specifies the elliptic curve to use.
|
|
Options:
|
|
- e.g. ed25519, ed448 for edwards form
|
|
- e.g. secp256k1 for weierstrass form
|
|
- e.g. k-409 for koblitz form
|
|
List of all available curves:
|
|
https://github.com/serengil/LightECC?tab=readme-ov-file#supported-curves
|
|
This parameter is only used if `algorithm_name` is 'EllipticCurve-ElGamal'.
|
|
plaintext_limit (int, optional): Upper bound for plaintext values.
|
|
This parameter is only used if `algorithm_name` is 'Benaloh'.
|
|
max_tries (int): maximum attempts to generate keys. Default is 10000.
|
|
RSA, Benaloh, Naccache-Stern and Goldwasser-Micali algorithms
|
|
need multiple attempts to generate valid keys. Will be discarded
|
|
for other algorithms.
|
|
"""
|
|
self.algorithm_name = algorithm_name
|
|
self.precision = precision
|
|
self.form = form
|
|
self.curve = curve
|
|
|
|
if key_file is not None:
|
|
keys = self.restore_keys(target_file=key_file)
|
|
|
|
if keys is not None:
|
|
validate_keys(algorithm_name=algorithm_name, keys=keys)
|
|
|
|
self.cs: Homomorphic = self.__build_cryptosystem(
|
|
algorithm_name=algorithm_name,
|
|
keys=keys,
|
|
key_size=key_size,
|
|
form=form,
|
|
curve=curve,
|
|
plaintext_limit=plaintext_limit,
|
|
max_tries=max_tries,
|
|
)
|
|
|
|
def __build_cryptosystem(
|
|
self,
|
|
algorithm_name: str = "Paillier",
|
|
keys: Optional[dict] = None,
|
|
key_size: Optional[int] = None,
|
|
form: Optional[str] = None,
|
|
curve: Optional[str] = None,
|
|
plaintext_limit: Optional[int] = None,
|
|
max_tries: int = 10000,
|
|
) -> Union[
|
|
RSA,
|
|
ElGamal,
|
|
Paillier,
|
|
DamgardJurik,
|
|
OkamotoUchiyama,
|
|
Benaloh,
|
|
NaccacheStern,
|
|
EllipticCurveElGamal,
|
|
SanderYoungYung,
|
|
BonehGohNissim,
|
|
]:
|
|
"""
|
|
Build a cryptosystem among partially homomorphic algorithms
|
|
Args:
|
|
algorithm_name (str): RSA | ElGamal | Exponential-ElGamal | EllipticCurve-ElGamal
|
|
| Paillier | Damgard-Jurik | Okamoto-Uchiyama | Benaloh | Naccache-Stern
|
|
| Goldwasser-Micali | Edwards-ElGamal | Sander-Young-Yung | Boneh-Goh-Nissim
|
|
Default is Paillier.
|
|
keys (dict): optional private-public key pair
|
|
key_file (str): if keys are exported, you can load them into cryptosystem
|
|
key_size (int): key size in bits
|
|
form (str): specifies the form of the elliptic curve.
|
|
Options: 'weierstrass' (default), 'edwards'.
|
|
This parameter is only used if `algorithm_name` is 'EllipticCurve-ElGamal'.
|
|
curve (str): specifies the elliptic curve to use.
|
|
Options:
|
|
- ed25519, ed448 for edwards form
|
|
- secp256k1 for weierstrass form
|
|
This parameter is only used if `algorithm_name` is 'EllipticCurve-ElGamal'.
|
|
plaintext_limit (int, optional): Upper bound for plaintext values.
|
|
This parameter is only used if `algorithm_name` is 'Benaloh'.
|
|
max_tries (int): maximum attempts to generate keys. Default is 10000.
|
|
RSA, Benaloh, Naccache-Stern and Goldwasser-Micali algorithms
|
|
need multiple attempts to generate valid keys. Will be discarded
|
|
for other algorithms.
|
|
Returns
|
|
cryptosystem
|
|
"""
|
|
# build cryptosystem
|
|
if algorithm_name == Algorithm.RSA:
|
|
cs = RSA(keys=keys, key_size=key_size, max_tries=max_tries)
|
|
elif algorithm_name == Algorithm.ElGamal:
|
|
cs = ElGamal(keys=keys, key_size=key_size)
|
|
elif algorithm_name == Algorithm.ExponentialElGamal:
|
|
cs = ElGamal(keys=keys, key_size=key_size, exponential=True)
|
|
elif algorithm_name == Algorithm.EllipticCurveElGamal:
|
|
cs = EllipticCurveElGamal(
|
|
keys=keys, key_size=key_size, form=form, curve=curve
|
|
)
|
|
elif algorithm_name == Algorithm.Paillier:
|
|
cs = Paillier(keys=keys, key_size=key_size)
|
|
elif algorithm_name == Algorithm.DamgardJurik:
|
|
cs = DamgardJurik(keys=keys, key_size=key_size)
|
|
elif algorithm_name == Algorithm.OkamotoUchiyama:
|
|
cs = OkamotoUchiyama(keys=keys, key_size=key_size)
|
|
elif algorithm_name == Algorithm.Benaloh:
|
|
cs = Benaloh(
|
|
keys=keys,
|
|
key_size=key_size,
|
|
plaintext_limit=plaintext_limit,
|
|
max_tries=max_tries,
|
|
)
|
|
elif algorithm_name == Algorithm.NaccacheStern:
|
|
cs = NaccacheStern(keys=keys, key_size=key_size, max_tries=max_tries)
|
|
elif algorithm_name == Algorithm.GoldwasserMicali:
|
|
cs = GoldwasserMicali(keys=keys, key_size=key_size, max_tries=max_tries)
|
|
elif algorithm_name == Algorithm.SanderYoungYung:
|
|
cs = SanderYoungYung(
|
|
keys=keys, key_size=key_size, plaintext_limit=plaintext_limit
|
|
)
|
|
elif algorithm_name == Algorithm.BonehGohNissim:
|
|
cs = BonehGohNissim(keys=keys, key_size=key_size, max_tries=max_tries)
|
|
else:
|
|
raise ValueError(f"unimplemented algorithm - {algorithm_name}")
|
|
return cs
|
|
|
|
def encrypt(
|
|
self, plaintext: Union[int, float, list], silent: bool = False
|
|
) -> Union[Ciphertext, EncryptedTensor]:
|
|
"""
|
|
Encrypt a plaintext with a built cryptosystem
|
|
Args:
|
|
plaintext (int, float or tensor): message
|
|
silent (bool): set this to True if you do not want to see progress bar
|
|
Returns
|
|
ciphertext (from lightphe.models.Ciphertext import Ciphertext): encrypted message
|
|
"""
|
|
if self.cs.keys.get("public_key") is None:
|
|
raise ValueError("You must have public key to perform encryption")
|
|
|
|
if isinstance(plaintext, list):
|
|
# then encrypt tensors
|
|
return self.__encrypt_tensors(tensor=plaintext, silent=silent)
|
|
|
|
ciphertext = self.cs.encrypt(
|
|
plaintext=phe_utils.normalize_input(
|
|
value=plaintext, modulo=self.cs.plaintext_modulo
|
|
)
|
|
)
|
|
|
|
public_keys = self.cs.keys.copy()
|
|
if public_keys.get("private_key") is not None:
|
|
del public_keys["private_key"]
|
|
|
|
return Ciphertext(
|
|
algorithm_name=self.algorithm_name,
|
|
keys=public_keys,
|
|
value=ciphertext,
|
|
form=self.form,
|
|
curve=self.curve,
|
|
)
|
|
|
|
def decrypt(
|
|
self, ciphertext: Union[Ciphertext, EncryptedTensor]
|
|
) -> Union[int, List[int], List[float]]:
|
|
"""
|
|
Decrypt a ciphertext with a buit cryptosystem
|
|
Args:
|
|
ciphertext (from lightphe.models.Ciphertext import Ciphertext): encrypted message
|
|
Returns:
|
|
plaintext (int): restored message
|
|
"""
|
|
if self.cs.keys.get("private_key") is None:
|
|
raise ValueError("You must have private key to perform decryption")
|
|
|
|
if self.cs.keys.get("public_key") is None:
|
|
raise ValueError("You must have public key to perform decryption")
|
|
|
|
if isinstance(ciphertext, EncryptedTensor):
|
|
# then this is encrypted tensor
|
|
return self.__decrypt_tensors(encrypted_tensor=ciphertext)
|
|
|
|
return self.cs.decrypt(ciphertext=ciphertext.value)
|
|
|
|
def __encrypt_tensors(self, tensor: list, silent: bool = False) -> EncryptedTensor:
|
|
"""
|
|
Encrypt a given tensor
|
|
Args:
|
|
tensor (list of int or float)
|
|
silent (bool): set this to True if you do not want to see progress bar
|
|
Returns
|
|
encrypted tensor (list of encrypted tensor object)
|
|
"""
|
|
encrypted_tensor: List[Fraction] = []
|
|
|
|
encrypted_zero = self.cs.encrypt(plaintext=0)
|
|
divisor_encrypted = self.cs.encrypt(plaintext=10**self.precision)
|
|
|
|
num_workers = min(len(tensor), 2 * multiprocessing.cpu_count())
|
|
logger.debug(f"encrypting tensors in {num_workers} parallel")
|
|
|
|
with closing(multiprocessing.Pool(num_workers)) as pool:
|
|
funclist = []
|
|
|
|
for m in tensor:
|
|
f = pool.apply_async(
|
|
encrypt_float,
|
|
(
|
|
m,
|
|
divisor_encrypted,
|
|
self.cs,
|
|
self.precision,
|
|
encrypted_zero,
|
|
),
|
|
)
|
|
funclist.append(f)
|
|
|
|
tic = time.time()
|
|
encrypted_tensor = []
|
|
for f in tqdm(
|
|
funclist,
|
|
desc="Encrypting tensors",
|
|
disable=silent,
|
|
):
|
|
result = f.get(timeout=10)
|
|
encrypted_tensor.append(result)
|
|
|
|
toc = time.time()
|
|
logger.debug(f"encryption took {toc - tic} seconds")
|
|
|
|
public_cs = copy.deepcopy(self.cs)
|
|
if public_cs.keys.get("private_key") is not None:
|
|
del public_cs.keys["private_key"]
|
|
|
|
return EncryptedTensor(
|
|
fractions=encrypted_tensor,
|
|
cs=public_cs,
|
|
precision=self.precision,
|
|
)
|
|
|
|
def __decrypt_tensors(
|
|
self, encrypted_tensor: EncryptedTensor
|
|
) -> Union[List[int], List[float]]:
|
|
"""
|
|
Decrypt a given encrypted tensor
|
|
Args:
|
|
encrypted_tensor (list of encrypted tensor)
|
|
Returns:
|
|
List of plain tensors
|
|
"""
|
|
plain_tensor = []
|
|
for c in encrypted_tensor.fractions:
|
|
if isinstance(c, Fraction) is False:
|
|
raise ValueError("Ciphertext items must be type of Fraction")
|
|
|
|
sign = c.sign
|
|
abs_dividend = self.cs.decrypt(ciphertext=c.abs_dividend)
|
|
# dividend = self.cs.decrypt(ciphertext=c.dividend)
|
|
|
|
# TODO: do I really need encrypted divisor? cannot I store current_precision
|
|
divisor = self.cs.decrypt(ciphertext=c.divisor)
|
|
|
|
m = sign * abs_dividend / divisor
|
|
|
|
plain_tensor.append(m)
|
|
return plain_tensor
|
|
|
|
def regenerate_ciphertext(self, ciphertext: Ciphertext) -> Ciphertext:
|
|
"""
|
|
Generate a different ciphertext belonging to same plaintext
|
|
Args:
|
|
ciphertext (from lightphe.models.Ciphertext import Ciphertext): encrypted message
|
|
Returns:
|
|
ciphertext (from lightphe.models.Ciphertext import Ciphertext): encrypted message
|
|
"""
|
|
if self.cs.keys.get("public_key") is None:
|
|
raise ValueError("You must have public key to perform decryption")
|
|
|
|
ciphertext_new = self.cs.reencrypt(ciphertext=ciphertext.value)
|
|
return Ciphertext(
|
|
algorithm_name=self.algorithm_name, keys=self.cs.keys, value=ciphertext_new
|
|
)
|
|
|
|
def export_keys(self, target_file: str, public: bool = False) -> None:
|
|
"""
|
|
Export keys to a file
|
|
Args:
|
|
target_file (str): target file name
|
|
public (bool): set this to True if you will publish this
|
|
to publicly.
|
|
"""
|
|
keys = self.cs.keys
|
|
private_key = None
|
|
if public is True and keys.get("private_key") is not None:
|
|
private_key = keys["private_key"]
|
|
del keys["private_key"]
|
|
|
|
if public is False:
|
|
logger.warn(
|
|
"You did not set public arg to True. So, exported key has private key information."
|
|
"Do not share this to anyone"
|
|
)
|
|
|
|
with open(target_file, "w", encoding="UTF-8") as file:
|
|
file.write(json.dumps(keys))
|
|
|
|
# restore private key if you dropped
|
|
if private_key is not None:
|
|
self.cs.keys["private_key"] = private_key
|
|
|
|
def restore_keys(self, target_file: str) -> dict:
|
|
"""
|
|
Restore keys from a file
|
|
Args:
|
|
target_file (str): target file name
|
|
Returns:
|
|
keys (dict): private public key pair
|
|
"""
|
|
with open(target_file, "r", encoding="UTF-8") as file:
|
|
dict_str = file.read()
|
|
|
|
keys = eval(dict_str)
|
|
if not isinstance(keys, dict):
|
|
raise ValueError(
|
|
f"The content of the file {target_file} does not represent a valid dictionary."
|
|
)
|
|
|
|
if "private_key" in keys.keys():
|
|
logger.info(f"private-public key pair is restored from {target_file}")
|
|
elif "public_key" in keys.keys():
|
|
logger.info(f"public key is restored from {target_file}")
|
|
else:
|
|
raise ValueError(f"File {target_file} must have public_key key")
|
|
return keys
|
|
|
|
def create_ciphertext_obj(self, ciphertext: Union[int, tuple, list]) -> Ciphertext:
|
|
"""
|
|
Ciphertext objects have keys in addition ciphertext itself to perform
|
|
homomorphic operations.
|
|
Args:
|
|
ciphertext (int or tuple or list): ciphertext content
|
|
Returns:
|
|
Ciphertext
|
|
"""
|
|
return Ciphertext(
|
|
algorithm_name=self.algorithm_name, keys=self.cs.keys, value=ciphertext
|
|
)
|
|
|
|
|
|
def encrypt_float(
|
|
m: Union[int, float],
|
|
divisor_encrypted: int,
|
|
cs: Homomorphic,
|
|
precision: int,
|
|
encrypted_zero: int,
|
|
) -> Fraction:
|
|
"""
|
|
Encrypt a float value
|
|
Args:
|
|
m (int or float): message to encrypt
|
|
divisor_encrypted (int): pre-calculated encrypted divisor
|
|
cs (Homomorphic): cryptosystem itself
|
|
precision (int): define how many digits after dot
|
|
encrypted_zero (int): pre-calculated encrypted value of 0
|
|
Returns:
|
|
result (Fraction): encrypted float value
|
|
"""
|
|
try:
|
|
if m == 0:
|
|
# this is very common in VGG-Face embeddings
|
|
c = Fraction(
|
|
dividend=encrypted_zero,
|
|
divisor=divisor_encrypted,
|
|
abs_dividend=encrypted_zero,
|
|
sign=1,
|
|
)
|
|
elif isinstance(m, int):
|
|
dividend_encrypted = cs.encrypt(
|
|
plaintext=(m % cs.plaintext_modulo) * pow(10, precision)
|
|
)
|
|
abs_dividend_encrypted = (
|
|
dividend_encrypted
|
|
if m > 0
|
|
else cs.encrypt(
|
|
plaintext=(abs(m) % cs.plaintext_modulo) * pow(10, precision)
|
|
)
|
|
)
|
|
# divisor_encrypted = self.cs.encrypt(plaintext=pow(10, self.precision))
|
|
c = Fraction(
|
|
dividend=dividend_encrypted,
|
|
divisor=divisor_encrypted,
|
|
abs_dividend=abs_dividend_encrypted,
|
|
sign=1 if m >= 0 else -1,
|
|
)
|
|
elif isinstance(m, float):
|
|
# got `int too large to convert float` while m mod plaintext modulo
|
|
# when security level is set to 128
|
|
dividend, _ = phe_utils.fractionize(
|
|
value=(m % cs.plaintext_modulo if m > cs.plaintext_modulo else m),
|
|
modulo=cs.plaintext_modulo,
|
|
precision=precision,
|
|
)
|
|
abs_dividend = (
|
|
dividend
|
|
if m > 0
|
|
else phe_utils.fractionize(
|
|
value=(
|
|
abs(m) % cs.plaintext_modulo
|
|
if abs(m) > cs.plaintext_modulo
|
|
else abs(m)
|
|
),
|
|
modulo=cs.plaintext_modulo,
|
|
precision=precision,
|
|
)[0]
|
|
)
|
|
dividend_encrypted = cs.encrypt(plaintext=dividend)
|
|
abs_dividend_encrypted = (
|
|
dividend_encrypted if m > 0 else cs.encrypt(plaintext=abs_dividend)
|
|
)
|
|
# divisor_encrypted = self.cs.encrypt(plaintext=_divisor)
|
|
c = Fraction(
|
|
dividend=dividend_encrypted,
|
|
divisor=divisor_encrypted,
|
|
abs_dividend=abs_dividend_encrypted,
|
|
sign=1 if m >= 0 else -1,
|
|
)
|
|
else:
|
|
raise ValueError(f"unimplemented type - {type(m)}")
|
|
|
|
return c
|
|
except Exception as err:
|
|
logger.error(f"Exception while running encrypt_float: {str(err)}")
|
|
logger.error(traceback.format_exc())
|
|
raise err
|