from random import randrange

def octets(n:int) -> int:
    """
    renvoie le nombre d'octets nécessaires pour écrire n
    """
    assert n >= 0
    s = 0
    while n > 0:
        n //= 256
        s += 1
    return s

BMIN = 1 # nombre minimum d'octets de bourrage
def cypher(mb:bytes, Kpu) -> bytes:
    """
    mb: message à chiffrer
    Kpu:(n,e) clé publique
    renvoie le message chiffré en découpant mb autant que nécessaire
    et en faisant un bourrage
    """
    n, e = Kpu
    s = octets(n)
    sbloc = s - 2 - BMIN # taille d'un bloc, en octets
    assert sbloc >= 1, "Clé trop petite"
    mbc = bytes()
    for i in range(0, len(mb), sbloc):
        bloc = mb[i:i+sbloc]
        rb = s - 2 - len(bloc)
        bloc = bytes([rb+1]+[randrange(256) for _ in range(rb)]) + bloc
        m = int.from_bytes(bloc, 'big')
        mc = pow(m, e, n)
        cbloc = mc.to_bytes(s, 'big')
        mbc += cbloc
    return mbc

def decypher(mbc:bytes, Kpr) -> bytes:
    """
    mbc: message à déchiffrer
    Kpr:(n,d) clé privée
    renvoie le message déchiffré
    """
    n, d = Kpr
    s = octets(n)
    mb = bytes()
    for i in range(0, len(mbc), s):
        cbloc = mbc[i:i+s]
        mc = int.from_bytes(cbloc, 'big')
        m = pow(mc, d, n)
        bloc = m.to_bytes(s-1, 'big')
        rb = bloc[0]
        mb += bloc[rb:]
    return mb

if __name__ == '__main__':
    Kpu = (191677976481073, 17)
    Kpr = (191677976481073, 56375867282673)
    mb = "Le petit éléphant se promenait en agitant sa trompe.".encode('utf8')
    mbc = cypher(mb, Kpu)
    mb2 = decypher(mbc, Kpr)
    assert mb2 == mb
