# pip3 install pyfinite
from pyfinite import ffield
class SAES:
def __init__(self, key):
X_generator = 0b10011
Y_generator = 0b10001
Z_generator = 0b101
self.a = 0b1101
self.b = 0b1001
self.X_field = ffield.FField(4, gen=X_generator, useLUT=0)
self.Y_field = ffield.FField(4, gen=Y_generator, useLUT=0)
self.Z_field = ffield.FField(4, gen=Z_generator, useLUT=0)
self._init_S()
self._extend_key(key)
def _init_S(self):
self.S_box = {}
self.S_box_inv = {}
for i in range(16):
N = self.X_field.Inverse(i)
Ny = self.Y_field.Multiply(N, 1)
res = self.Y_field.Add(self.Y_field.Multiply(Ny, self.a), self.b)
self.S_box[i] = res
self.S_box_inv[res] = i
def S(self, x):
return self.S_box[x]
def S_inv(self, x):
return self.S_box_inv[x]
def _sub_nib(self, nibble_pair):
n1 = (0b11110000 & nibble_pair) >> 4
n2 = 0b1111 & nibble_pair
n1_sub = self.S(n1)
n2_sub = self.S(n2)
return n1_sub * 16 + n2_sub
def _rot_nib(self, nibble_pair):
n1 = (0b11110000 & nibble_pair) >> 4
n2 = 0b1111 & nibble_pair
return n2 * 16 + n1
def _bytes_to_matrix(self, byte_array):
matrix = [[(byte_array[0] & 0b11110000) >> 4, (byte_array[1] & 0b11110000) >> 4],
[byte_array[0] & 0b1111, byte_array[1] & 0b1111]]
return matrix
def _add_key(self, i, state):
key_part = self.extended_key[i * 2: i * 2 + 2]
key_matrix = self._bytes_to_matrix(key_part)
result = [[key_matrix[0][0] ^ state[0][0], key_matrix[0][1] ^ state[0][1]],
[key_matrix[1][0] ^ state[1][0], key_matrix[1][1] ^ state[1][1]]]
return result
def _nibble_substitution(self, state):
sub_matrix = [[self.S(state[0][0]), self.S(state[0][1])],
[self.S(state[1][0]), self.S(state[1][1])]]
return sub_matrix
def _nibble_substitution_inv(self, state):
sub_matrix = [[self.S_inv(state[0][0]), self.S_inv(state[0][1])],
[self.S_inv(state[1][0]), self.S_inv(state[1][1])]]
return sub_matrix
def _shift_row(self, state):
return [[state[0][0], state[0][1]],
[state[1][1], state[1][0]]]
def _mix_columns(self, state):
Ni1 = state[0][0]
Nj1 = state[1][0]
Ni2 = state[0][1]
Nj2 = state[1][1]
return [[self.X_field.Add(Ni1, self.X_field.Multiply(Nj1, 0b100)), self.X_field.Add(Ni2, self.X_field.Multiply(Nj2, 0b100))],
[self.X_field.Add(self.X_field.Multiply(Ni1, 0b100), Nj1), self.X_field.Add(self.X_field.Multiply(Ni2, 0b100), Nj2)]]
def _mix_columns_inv(self, state):
Ni1 = state[0][0]
Nj1 = state[1][0]
Ni2 = state[0][1]
Nj2 = state[1][1]
return [[self.Z_field.Add(self.X_field.Multiply(Ni1, 0b1001), self.X_field.Multiply(Nj1, 0b10)), self.Z_field.Add(self.X_field.Multiply(Ni2, 0b1001), self.X_field.Multiply(Nj2, 0b10))],
[self.Z_field.Add(self.X_field.Multiply(Ni1, 0b10), self.X_field.Multiply(Nj1, 0b1001)), self.Z_field.Add(self.X_field.Multiply(Ni2, 0b10), self.X_field.Multiply(Nj2, 0b1001))]]
def _extend_key(self, key):
if len(key) != 2:
raise Exception('Invalid key length')
W = [ord(key[0]), ord(key[1]),0,0,0,0]
RC = []
x_2 = [0,1,0,0]
x2 = self.X_field.ConvertListToElement(x_2)
for i in range(1,4):
x_i = [0,0,0,0]
x_i[-i-1] = 1
xi = self.X_field.ConvertListToElement(x_i)
RC.append(self.X_field.Multiply(xi, x2))
RCON = [0] + [rc * 16 for rc in RC]
for i in range(2, 6):
if i % 2 == 0:
k = self._sub_nib(self._rot_nib(W[i - 1]))
W[i] = RCON[i//2] ^ k ^ W[i - 2]
else:
W[i] = W[i - 1] ^ W[i - 2]
self.extended_key = W
def _print_state(self, state):
print([[bin(state[0][0]), bin(state[0][1])],
[bin(state[1][0]), bin(state[1][1])]])
def _encrypt_bytes(self, data_bytes):
state = self._bytes_to_matrix(data_bytes)
state = self._add_key(0, state)
state = self._nibble_substitution(state)
state = self._shift_row(state)
state = self._mix_columns(state)
state = self._add_key(1, state)
state = self._nibble_substitution(state)
state = self._shift_row(state)
state = self._add_key(2 , state)
return [state[0][0] * 16 + state[1][0], state[0][1] * 16 + state[1][1]]
def _decrypt_bytes(self, data_bytes):
state = self._bytes_to_matrix(data_bytes)
state = self._add_key(2, state)
state = self._shift_row(state)
state = self._nibble_substitution_inv(state)
state = self._add_key(1, state)
state = self._mix_columns_inv(state)
state = self._shift_row(state)
state = self._nibble_substitution_inv(state)
state = self._add_key(0, state)
return [state[0][0] * 16 + state[1][0], state[0][1] * 16 + state[1][1]]
def encrypt(self, string_data):
data_bytes = [ord(x) for x in string_data]
n = len(data_bytes)
if n % 2 == 1:
data_bytes.append(ord(' '))
encrypted_bytes = []
for i in range(0, n, 2):
data_bytes_slice = data_bytes[i:i+2]
encrypted_bytes += self._encrypt_bytes(data_bytes_slice)
return ''.join([chr(x) for x in encrypted_bytes])
def decrypt(self, string_data_enc):
data_bytes = [ord(x) for x in string_data_enc]
decrypted_bytes = []
n = len(data_bytes)
for i in range(0, n, 2):
data_bytes_slice = data_bytes[i:i+2]
decrypted_bytes += self._decrypt_bytes(data_bytes_slice)
return ''.join([chr(x) for x in decrypted_bytes])