In [1]:
# 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])
In [4]:
saes = SAES('Yz')
plain_text = 'Ovaj tekst ce biti EnKrIpToVaN simple SAES.'
print(f'Plain text: {plain_text}')
encrypted_data = saes.encrypt(plain_text)
print('Encrypted text (SAES):')
print(encrypted_data)
decrypted_data = saes.decrypt(encrypted_data)
print(f'Decrypted text (SAES):\n{decrypted_data}')

# print('Cracking the password (Plain text attack)...')
# for i in range(256):
#     for j in range(256):
#         found = False
#         key = chr(i) + chr(j)
#         saes_intruder = SAES(key)
#         if saes_intruder.encrypt('Kriptografija@MATF') == encrypted_data:
#             print('Password cracked: ',key)
#             found = True
#             break
#     if found:
#         break
Plain text: Ovaj tekst ce biti EnKrIpToVaN simple SAES.
Encrypted text (SAES):
 Ñ·Ý$÷~ð‘tùHã2¤dx3æ~@‹(œ¤ÕJ<Éôø )Jã2„!Ô		
Decrypted text (SAES):
Ovaj tekst ce biti EnKrIpToVaN simple SAES. 
In [ ]: