# -*- coding: utf-8 -*-
"""
Created on Wed Feb 10 10:50:56 2021

@author: G5A
"""


def cesar_encode(message, key):
    if (len(message) == 0):
        return False
    key %= 26
    message_encode = ""
    ordi = 0
    for letter in message:
        ordi = ord(letter)
        if (ordi >= 97 and ordi <= 122):
            ordi += key
            if (ordi > 122):
                ordi -= 26
        if (ordi >= 65 and ordi <= 90):
            ordi += key
            if (ordi > 90):
                ordi -= 26
        message_encode += chr(ordi)
    return message_encode


def cesar_decode(message, key):
    if (len(message) == 0):
        return False
    key %= 26
    message_decode = ""
    ordi = 0
    for letter in message:
        ordi = ord(letter)
        if (ordi >= 97 and ordi <= 122):
            ordi -= key
            if (ordi < 97):
                ordi += 26
        if (ordi >= 65 and ordi <= 90):
            ordi -= key
            if (ordi < 65):
                ordi += 26
        message_decode += chr(ordi)
    return message_decode


def wegner_encode(message, key):
    if (not message or not key):
        return False
    key = key.lower()
    i = 0
    message_encode = ""
    ordi = 0
    for letter in message:
        if(key[i] == ' '):   
            i+=1
        ordi = ord(letter)
        if (ordi >= 97 and ordi <= 122):
            ordi += ord(key[i]) - 97
            if (ordi > 122):
                ordi -= 26
        if (ordi >= 65 and ordi <= 90):
            ordi += ord(key[i]) - 97
            if (ordi > 90):
                ordi -= 26
        if (letter.isalpha()):
            i += 1
            i = i % len(key)
        message_encode += chr(ordi)
    return message_encode


def wegner_decode(message, key):
    if (not message or not key):
        return False
    key = key.lower()
    i = 0
    message_decode = ""
    ordi = 0
    for letter in message:
        if(key[i] == ' '):
            i+=1
        ordi = ord(letter)
        if (ordi >= 97 and ordi <= 122):
            ordi -= ord(key[i]) - 97
            if (ordi < 97):
                ordi += 26
        if (ordi >= 65 and ordi <= 90):
            ordi -= ord(key[i]) - 97
            if (ordi < 65):
                ordi += 26
        if (letter.isalpha()):
            i += 1
            i = i % len(key)
        message_decode += chr(ordi)
    return message_decode


def resultat(type_code, encode_decode, key, code, key2 = None):
    result = ""
    if (type_code == "Cesar"):
        key = int(key)
        if (encode_decode == "encode"):
            result = cesar_encode(code, key)
        elif (encode_decode == "decode"):
            result = cesar_decode(code, key)
        else:
            print("erreur code cesar")
    elif (type_code == "Wegner"):
        if (encode_decode == "encode"):
            result = wegner_encode(code, key)
        elif (encode_decode == "decode"):
            result = wegner_decode(code, key)
        else:
            print("erreur code wegner")
    elif (type_code == "RC4"):
        if (encode_decode == "encode"):
            result = RC4_encode(code, key)
        elif (encode_decode == "decode"):
            result = RC4_decode(code, key)
        else:
            print("erreur code RC4")
    elif(type_code == "Affine"):
        if(encode_decode == "encode"):
            result = affine_encode(code, key, key2)
        elif(encode_decode == "decode"):
            result = affine_decode(code, key, key2)
        else:
            print("erreur code affine")
    else:
        print("erreur code choice")
    return result

def getKeyStream(size,K):
    K = [ord(x) for x in K]
    # Initialisation : (Key Schedule Algorithm)
    S = [i for i in range(256)]
    T = [K[i % len(K)] for i in range(256)]
    j = 0
    for i in range(256):
        j = (j + S[i] + T[i]) % 256
        S[i], S[j] = S[j], S[i]
    # Keystream : (Pseudo-Random Generation Algorithm)
    i, j = 0, 0
    KS = []
    for k in range(size):
        i = (i + 1) % 256
        j = (j + S[i]) % 256
        S[i], S[j] = S[j], S[i]
        t = (S[i] + S[j]) % 256
        KS.append(S[t])
    return KS

def RC4_encode(M, K):
    if (not M or not K):
        return False
    M = [ord(x) for x in M]
    KS = getKeyStream(len(M),K)
    MC = [M[k] ^ KS[k] for k in range(len(M))]
    return " ".join([str(x) for x in MC])

def RC4_decode(MC,K):
    if (not MC or not K):
        return False
    MC = [int(x) for x in MC.split(" ")]
    KS = getKeyStream(len(MC),K)
    M = [MC[k] ^ KS[k] for k in range(len(MC))]
    return "".join([chr(x) for x in M])


def check_inversible(val):
    inverse = 0
    for i in range (1,26):
        if (i*val) % 26 == 1:
            inverse = i
    return inverse

def affine_encode(message, a, b):
    if (not message or not a or not b):
        return False
    inverse_a = check_inversible(a)
    message_chiffre = ""
    if inverse_a:
        ordi = 0
        for letter in message:
            ordi = ord(letter)
            if (ordi >= 97 and ordi <= 122):
                z = ordi - 97
                z = (a*z+b) % 26
                c = chr(z + 97)
            if (ordi >= 65 and ordi <= 90):
                z = ordi - 65
                z = (a*z+b) % 26
                c = chr(z + 65)
            if letter.isalpha():
                message_chiffre += c
            elif letter == ' ':
                message_chiffre += letter
        return message_chiffre
    else:
        print('{}'.format(a) + " n'est pas une valeur inversible modulo 26")
        return message_chiffre
    
def affine_decode(message_chiffre, a, b):
    if (not message_chiffre or not a or not b):
        return False
    inverse_a = check_inversible(a)
    message = ""
    if inverse_a:
        ordi = 0
        for letter in message_chiffre:
            ordi = ord(letter)
            if (ordi >= 97 and ordi <= 122):
                z = ordi - 97
                z = (inverse_a * z - inverse_a * b) % 26
                c = chr(z + 97)
            if (ordi >= 65 and ordi <= 90):
                z = ordi - 65
                z = (inverse_a * z - inverse_a * b) % 26
                c = chr(z + 65)
            if letter.isalpha():
                message += c
            elif letter == ' ':
                message += letter            
        return message
    else:
        print('{}'.format(a) + " n'est pas une valeur inversible modulo 26")
        return message