Skip to content

RSA Implementation

shinywaterjeong edited this page Nov 17, 2020 · 1 revision

RSA

the RSA Algorithm

RSA는 아래와 같은 방식을 통해 구현이 되어있다.

Key Generation

  1. 소수이면서 서로 다른 값을 가지는 p, q를 찾는다.
    1. 임의의 홀수를 생성하여 Miller-Rabin test 를 통해 소수인지 확인한다.
    2. Miller-Rabin test를 통과할 때까지 임의의 수를 생성하여 소수를 만든다.
  2. 두 p, q를 곱하여 n 을 계산한다.
  3. Euler Totient Function을 이용하여 ϕ(n)을 구한다.
    1. n은 p와 q, 두 소수의 곱으로 표현된다.
    2. ϕ(n) = (p - 1) * (q - 1)가 성립한다.
  4. ϕ(n)과 서로소이면서 1 < e < ϕ(n)을 만족하는 e 를 찾는다.
    1. e가 ϕ(n)와 서로소 라면 gcd(ϕ(n), e) = 1이다.
    2. Euclide Algorithm 을 이용해 값을 구할 수 있다.
  5. mod ϕ(n)에 대해 e의 역원 d 를 찾는다.
    1. d = e-1 mod ϕ(n)을 만족한다.
    2. 즉, de mod ϕ(n) = 1을 만족한다.
    3. Extended Euclide Algorithm 을 통해 값을 구할 수 있다.
  6. 위 모든 계산의 결과로 public key와 private key를 얻을 수 있다.
    1. PU = {e, n}
    2. PR = {d, n}

Encryption

  1. 보내고 싶은 plaintext인 M은 n보다 크기가 작다.
    1. M < n
  2. ciphertext C는 상대방의 public key, PU = {e, n} 를 이용하여 얻는다.
    1. C = Me mod n을 계산한다.

Decryption

  1. ciphertext C는 자신의 private key, PR = {d, n} 를 이용하여 decryption 가능하다.
    1. M = Cd mod n
  2. M = Cd mod n은 아래와 같이 증명 가능하다.
    Cd mod n
    = (Me mod n)d mod n, (C = Me mod n)
    = Med mod n
    = Mkϕ(n)+1 mod n
    = Mkϕ(n) * M mod n
    = M mod n (Mkϕ(n) mod d = 1 by 중국인의 나머지 정리)
    = M (by M < n)

Public-key Cryptosystem Implementation

import random
import sys

sys.setrecursionlimit(2048)


def gcd(a, b):
    if a < b:
        a, b = b, a
    if a == b:
        return a
    if b == 0:
        return a
    return gcd(b, a % b)


def extended_euclid(a, b):
    if a == b:
        return 1, 0, a
    if b == 0:
        return 1, 0, a
    x_1, y_1, r_1 = 1, 0, a
    x_2, y_2, r_2 = 0, 1, b
    while r_2 != 0:
        q = r_1 // r_2

        r_t = r_1 - q * r_2
        x_t = x_1 - q * x_2
        y_t = y_1 - q * y_2

        x_1, y_1, r_1 = x_2, y_2, r_2
        x_2, y_2, r_2 = x_t, y_t, r_t
    return x_1, y_1, r_1


def m_inv(a, n):
    x, y, r = extended_euclid(n, a % n)
    if r != 1:
        print("No multiplicative inverse")
        return
    return y % n


def int_to_bin(num):
    return list(bin(num))[2:]


def exp(a, b, n):
    c, f = 0, 1
    bin_b = int_to_bin(b);
    k = len(bin_b)
    for i in range(k):
        c = 2 * c
        f = (f * f) % n
        if bin_b[i] == '1':
            c = c + 1
            f = (f * a) % n
    return f


Prime = 0
Composite = 1


def miller_rabin(n, s):
    if n == 2:
        return Prime
    if n % 2 == 0:
        return Composite
    for _ in range(s):
        a = random.randint(1, n - 1)
        if test(a, n) == Composite:
            return Composite
    return Prime


def test(a, n):
    bits = int_to_bin(n - 1)
    k = len(bits) - 1
    t = 0
    while bits[k] == '0':
        t = t + 1
        k = k - 1
    u = (n - 1) >> t
    x = exp(a, u, n)

    for _ in range(t):
        _x = x
        x = (_x * _x) % n
        if x == 1 & _x != 1 & _x != n - 1:
            return Composite
    if x != 1:
        return Composite
    return Prime


def keygen(len):
    bound = 1 << (len // 2)
    p = 2 * random.randint(bound // 4 + 1, bound // 2) - 1
    while miller_rabin(p, 50) == Composite:
        p = 2 * random.randint(bound // 4 + 1, bound // 2) - 1
    q = 2 * random.randint(bound // 4 + 1, bound // 2) - 1
    while miller_rabin(q, 50) == Composite:
        q = 2 * random.randint(bound // 4 + 1, bound // 2) - 1

    n = p * q
    phi_n = (p - 1) * (q - 1)
    e = 2 * random.randint(1, phi_n // 2) - 1
    while gcd(phi_n, e) != 1:
        e = 2 * random.randint(1, phi_n // 2) - 1

    d = m_inv(e, phi_n)
    return e, d, n


def encrypt(M, e, n):
    return exp(M, e, n)


def decrypt(C, d, n):
    return exp(C, d, n)


def rsa_enc_test(bitlenlst, itr1, itr2):
    for bitlen in bitlenlst:
        for i in range(itr1):
            e, d, n = keygen(bitlen)
            for j in range(itr2):
                M = random.randint(1, n - 1)
                C = encrypt(M, e, n)
                # M = (M + 1) % n
                MM = decrypt(C, d, n)
                print("%dbit-RSA Secrecy [%04d-%04d] M, C, "
                      "MM: %d, %d, %d" % (bitlen, i + 1, j + 1, M, C, MM))


def sign(M, d, n):
    return exp(M, d, n)


def verify(M, S, e, n):
    return M == exp(S, e, n)


def rsa_sign_test(bitlenlst, itr1, itr2):
    for bitlen in bitlenlst:
        for i in range(itr1):
            e, d, n = keygen(bitlen)
            for j in range(itr2):
                M = random.randint(1, n - 1)
                S = sign(M, d, n)

                # M = (M + 1) % n
                # S = (S + 1) % n
                result = verify(M, S, e, n)
                print("%dbit-RSA Auth [%04d-%04d] "
                      "M, S: %d, %d" % (bitlen, i + 1, j + 1, M, S))
                if result:
                    print("  verification success.")
                else:
                    print("  verification failed.")


def rsa_enc_and_sign_test(bitlenlst, itr1, itr2):
    for bitlen in bitlenlst:
        for i in range(itr1):
            eA, dA, nA = keygen(bitlen)
            # 2 * nA < nB
            sys.setrecursionlimit(2 * bitlen + 10)
            eB, dB, nB = keygen(2 * bitlen + 10)

            for j in range(itr2):
                M = random.randint(1, nA)
                S = sign(M, dA, nA)
                MS = M * (2 ** bitlen) + S
                C = encrypt(MS, eB, nB)
                # C = (C + 1) % nB
                DMS = decrypt(C, dB, nB)
                DM, DS = DMS // (2 ** bitlen), DMS % (2 ** bitlen)

                result = verify(DM, DS, eA, nA)
                print("[%02d] M, S, C, MD, DS: %d, %d, %d, %d, %d" % (i + 1, M, S, C, DM, DS))
                if result:
                    print("  verification success.")
                else:
                    print("  verification failed.")


if __name__ == "__main__":
    e, d, n = keygen(128)
    M = 88
    C = encrypt(M, e, n)
    MM = decrypt(C, d, n)
    if M == MM:
        print("M is same from MM")
    else:
        print("M is not same from MM")
    print("M={}, PU=({}, {}), PR=({}, {}), C={}, MM={}".format(M, e, n, d, n, C, MM))

    rsa_enc_test([128, 256, 1024, 2048], 4, 4)

    rsa_sign_test([128], 3, 4)

    rsa_enc_and_sign_test([128], 4, 4)

Clone this wiki locally