diff --git a/receiver/challenges/Sheesh.py b/receiver/challenges/Sheesh.py index ef6496a..31f0da7 100644 --- a/receiver/challenges/Sheesh.py +++ b/receiver/challenges/Sheesh.py @@ -4,6 +4,7 @@ import subprocess import time import re import os +import binascii class Sheesh(Challenge): flag_location = 'flags/sheesh.txt' @@ -14,8 +15,10 @@ class Sheesh(Challenge): _HEX_RE = re.compile(r'^[0-9a-fA-F]+$') def _read_container_flag(self) -> str: - out = subprocess.run(["docker", "exec", self._CONTAINER, "cat", "/flag.txt"], - capture_output=True, text=True) + out = subprocess.run( + ["docker", "exec", self._CONTAINER, "cat", "/flag.txt"], + capture_output=True, text=True + ) if out.returncode != 0 or not out.stdout.strip(): raise FileNotFoundError("Flag not found in container (/flag.txt)") return out.stdout.strip() @@ -30,23 +33,34 @@ class Sheesh(Challenge): bufsize=0, ) - def _read_until(self, proc, token, timeout=5.0, max_bytes=1_000_000): + def _read_exact_line(self, proc, timeout=5.0): start = time.time() buf = [] r = proc.stdout.read while True: if time.time() - start > timeout: - tail = ''.join(buf)[-500:] - raise TimeoutError(f"Timeout waiting for '{token}'. Got so far:\n{tail}") + raise TimeoutError("Timeout waiting for a line") ch = r(1) if ch == "" and proc.poll() is not None: - raise RuntimeError(f"Process ended while waiting for '{token}'. Output:\n{''.join(buf)}") + raise RuntimeError("Process ended unexpectedly") buf.append(ch) - if len(buf) > max_bytes: - raise RuntimeError("Exceeded max read size") - if token in "".join(buf): + if ch == "\n": return "".join(buf) + def _read_until_prompt(self, proc, token: str, timeout=5.0): + start = time.time() + acc = [] + r = proc.stdout.read + while True: + if time.time() - start > timeout: + raise TimeoutError(f"Timeout waiting for '{token}'") + ch = r(1) + if ch == "" and proc.poll() is not None: + raise RuntimeError("Process ended while waiting for prompt") + acc.append(ch) + if "".join(acc).endswith(token): + return "".join(acc) + def _send_line(self, proc, s: str): proc.stdin.write(s + "\n") proc.stdin.flush() @@ -76,6 +90,7 @@ class Sheesh(Challenge): def check(self): try: + # parity check: host vs container flag with open(self.flag_location, 'r') as f: host_flag = f.read().strip() container_flag = self._read_container_flag() @@ -84,96 +99,66 @@ class Sheesh(Challenge): proc = self._spawn() - def menu(): - self._read_until(proc, "> ", timeout=5.0) + line1 = self._read_exact_line(proc, timeout=5.0) + line2 = self._read_exact_line(proc, timeout=5.0) + blank = self._read_exact_line(proc, timeout=5.0) + assert line1.startswith("iv:") and line2.startswith("ct:"), "Missing banner" + iv_hex = self._expect_hex_field(line1, "iv") + ct_hex = self._expect_hex_field(line2, "ct") + iv0 = bytes.fromhex(iv_hex) + ct0 = bytes.fromhex(ct_hex) + assert len(iv0) == 16 and len(ct0) >= 16 and len(ct0) % 16 == 0, "Bad banner sizes" + self.logger.info("[ok] banner parsed") - menu() + self._read_until_prompt(proc, "blob: ", timeout=5.0) - self._send_line(proc, "1") - self._read_until(proc, "pt: ", timeout=3.0) - pt_hex = "414243444546" - self._send_line(proc, pt_hex) - out = self._read_until(proc, "\n\n", timeout=3.0) - ct_hex = self._expect_hex_field(out, "ct") - ct = bytes.fromhex(ct_hex) - assert len(ct) == 16 + len(bytes.fromhex(pt_hex)), "CFB ct length mismatch (IV+PT)" - iv_a = ct[:16] - self.logger.info("[ok] encrypt(1) basic") + qA = "00" * 16 + self._send_line(proc, qA) + lA = self._read_exact_line(proc, timeout=3.0) + assert lA.startswith("blk:"), "Expected blk: for 16-byte blob" + blk1 = self._expect_hex_field(lA, "blk") + _ = self._read_exact_line(proc, timeout=3.0) - menu() - self._send_line(proc, "1") - self._read_until(proc, "pt: ", timeout=3.0) - pt_hex2 = "00" * 8 - self._send_line(proc, pt_hex2) - out2 = self._read_until(proc, "\n\n", timeout=3.0) - ct2_hex = self._expect_hex_field(out2, "ct") - ct2 = bytes.fromhex(ct2_hex) - assert len(ct2) == 16 + len(bytes.fromhex(pt_hex2)), "Second CFB ct length mismatch" - iv_b = ct2[:16] - assert iv_a != iv_b, "CFB IV appears reused" - self.logger.info("[ok] encrypt(1) IV rotates") + self._read_until_prompt(proc, "blob: ", timeout=5.0) + self._send_line(proc, qA) + lA2 = self._read_exact_line(proc, timeout=3.0) + blk2 = self._expect_hex_field(lA2, "blk") + _ = self._read_exact_line(proc, timeout=3.0) + assert blk1 == blk2, "ECB-masked oracle not deterministic for same input" + self.logger.info("[ok] blk determinism") - menu() - self._send_line(proc, "3") - self._read_until(proc, "pt: ", timeout=3.0) - self._send_line(proc, "00" * 15) - out3a = self._read_until(proc, "\n", timeout=3.0) - assert "hmmm" in out3a.lower(), "Option 3 must reject non-16-byte input (15B)" + self._read_until_prompt(proc, "blob: ", timeout=5.0) + qB = "01" * 16 + self._send_line(proc, qB) + lB = self._read_exact_line(proc, timeout=3.0) + blk3 = self._expect_hex_field(lB, "blk") + _ = self._read_exact_line(proc, timeout=3.0) + assert blk3 != blk1, "ECB-masked oracle returned same block for different input" + self.logger.info("[ok] blk diversity") - menu() - self._send_line(proc, "3") - self._read_until(proc, "pt: ", timeout=3.0) - self._send_line(proc, "00" * 17) - out3b = self._read_until(proc, "\n", timeout=3.0) - assert "hmmm" in out3b.lower(), "Option 3 must reject non-16-byte input (17B)" + self._read_until_prompt(proc, "blob: ", timeout=5.0) + self._send_line(proc, "zzzz") # invalid hex + line_err = self._read_exact_line(proc, timeout=3.0) + assert "err" in line_err.lower(), "Expected 'err' for non-hex input" + self.logger.info("[ok] error handling on non-hex") - menu() - self._send_line(proc, "3") - self._read_until(proc, "pt: ", timeout=3.0) - self._send_line(proc, "11" * 16) - out3 = self._read_until(proc, "\n\n", timeout=3.0) - ct3_hex = self._expect_hex_field(out3, "ct") - ct3 = bytes.fromhex(ct3_hex) - assert len(ct3) in (16, 32), f"Unexpected CT length for option 3: {len(ct3)} (expected 16 or 32)" - if len(ct3) == 16: - self.logger.info("[ok] get third(3): vuln flavor (16-byte CT)") - else: - self.logger.info("[ok] get third(3): patched flavor (32-byte CT)") + self._read_until_prompt(proc, "blob: ", timeout=5.0) + self._send_line(proc, (iv_hex + ct_hex)) + ans1 = self._read_exact_line(proc, timeout=3.0).strip().lower() + assert ans1 in ("ok", "no"), "Padding oracle did not return ok/no for iv||ct" + self.logger.info(f"[ok] padding oracle responds ({ans1}) on banner bundle") - menu() - self._send_line(proc, "3") - out4 = self._read_until(proc, "\n", timeout=3.0) - assert "sheesh" in out4.lower(), "Second call to option 3 should be rejected with 'sheesh'" - self.logger.info("[ok] get third(3) lockout") + self._read_until_prompt(proc, "blob: ", timeout=5.0) + iv_synth = os.urandom(16).hex() + ct_synth = ("00" * (6 * 16)) + self._send_line(proc, iv_synth + ct_synth) + ans2 = self._read_exact_line(proc, timeout=3.0).strip().lower() + assert ans2 in ("ok", "no"), "Padding oracle did not return ok/no on synthetic 6-block ct" + self.logger.info(f"[ok] padding oracle responds ({ans2}) on synthetic payload") - menu() - self._send_line(proc, "2") - out5 = self._read_until(proc, "\n\n", timeout=3.0) - iv1_hex = self._expect_hex_field(out5, "iv1") - iv2_hex = self._expect_hex_field(out5, "iv2") - ct5_hex = self._expect_hex_field(out5, "ct") - assert len(bytes.fromhex(iv1_hex)) == 16, "iv1 length invalid" - assert len(bytes.fromhex(iv2_hex)) == 16, "iv2 length invalid" - assert len(bytes.fromhex(ct5_hex)) >= 16 and len(bytes.fromhex(ct5_hex)) % 16 == 0, "CBC ct invalid length" - self.logger.info("[ok] profit(2) first call") - - menu() - self._send_line(proc, "2") - out6 = self._read_until(proc, "\n\n", timeout=3.0) - iv1_hex_2 = self._expect_hex_field(out6, "iv1") - iv2_hex_2 = self._expect_hex_field(out6, "iv2") - assert iv1_hex_2 != iv1_hex or iv2_hex_2 != iv2_hex, "IVs did not change between consecutive '2' calls" - self.logger.info("[ok] profit(2) fresh IVs") - - menu() - self._send_line(proc, "4") - try: - proc.wait(timeout=2.0) - except subprocess.TimeoutExpired: - proc.kill() - raise AssertionError("Program did not exit after option 4") - self.logger.info("[ok] service exit on 4") + self._read_until_prompt(proc, "blob: ", timeout=5.0) + proc.kill() self.logger.info('Check passed for sheesh') return True diff --git a/services/sheesh/dist/chall.py b/services/sheesh/dist/chall.py index e7e30b7..032ef46 100644 --- a/services/sheesh/dist/chall.py +++ b/services/sheesh/dist/chall.py @@ -1,109 +1,98 @@ #!/usr/bin/env python3 -import os -import binascii -import hashlib -import signal +import os, sys, signal, binascii, random from Crypto.Cipher import AES from Crypto.Util.Padding import pad, unpad +from Crypto.Util.number import bytes_to_long, long_to_bytes -signal.alarm(50) -seed_bits = 23 -seed_max = 1 << seed_bits -seed_len = (seed_bits + 7) // 8 -key = os.urandom(16) +random.seed(os.urandom(16)) +K0 = os.urandom(16) +K1 = os.urandom(16) +S0 = os.urandom(16) +S1 = os.urandom(16) +M0 = os.urandom(16) +M1 = os.urandom(16) -def hash_seed(seed_int: int) -> bytes: - sb = seed_int.to_bytes(seed_len, "big") - return hashlib.sha256(sb).digest()[:16] - -seed = int.from_bytes(os.urandom(4), "big") % seed_max -seed2 = int.from_bytes(os.urandom(4), "big") % seed_max -K1 = hash_seed(seed) -K2 = hash_seed(seed2) - -with open("/flag.txt", "rb") as f: +with open("/flag.txt","rb") as f: flag = f.read() -def read_hex(prompt: str): - s = input(prompt).strip() - try: - return binascii.unhexlify(s) - except Exception: - print("hmm") +def hex_input(q): + s = input(q).strip() + try: return binascii.unhexlify(s) + except: print("err"); return None + +def enc1(b16: bytes) -> bytes: + x = AES.new(K1, AES.MODE_ECB).encrypt(b16) + return bytes(a ^ b for a, b in zip(x, b16)) + +def enc2(iv: bytes, m: bytes) -> bytes: + return AES.new(K0, AES.MODE_CBC, iv=iv).encrypt(pad(m, 16)) + +def enc3(m: bytes) -> bytes: + return AES.new(K1, AES.MODE_CBC, iv=b"\x00"*16).encrypt(pad(m, 16))[-16:] + +def T(iv: bytes, ct: bytes): + n = len(ct) + if n < 96 or (n & 15): + return None + v = memoryview(ct) + W = [bytes(v[i:i+16]) for i in range(0, n, 16)] + m = len(W) + + r = ((iv[0] & 7) + 2) % m + if r: + W = W[r:] + W[:r] + + j = 1 + (W[0][0] & 1) + if len(W) <= j: + return None + del W[j] + if len(W) < 2: return None -def enc_cfb(pt: bytes) -> bytes: - iv = os.urandom(16) - aes = AES.new(key, AES.MODE_CFB, iv=iv, segment_size=128) - ct = aes.encrypt(pt) - return iv + ct + a0 = long_to_bytes(bytes_to_long(W[0]) ^ bytes_to_long(M0)) + a1 = long_to_bytes(bytes_to_long(W[1]) ^ bytes_to_long(M1)) + return b"".join((S0, a0, S1, a1, *W[2:])) -def enc_cbc(data: bytes, iv1: bytes, iv2: bytes, padd: bool) -> bytes: - x = pad(data, 16) if padd else data - c1 = AES.new(K1, AES.MODE_CBC, iv=iv1).encrypt(x) - c2 = AES.new(K2, AES.MODE_CBC, iv=iv2).encrypt(c1) - return c2 +def C(iv: bytes, ct: bytes) -> bool: + z = T(iv, ct) + if z is None: + ok = False + else: + try: + x = AES.new(K0, AES.MODE_CBC, iv=iv).decrypt(z) + unpad(x, 16) + ok = True + except: + ok = False + if random.random() < 0.08: + ok = not ok + return ok -def menu(): - print(""" -1. encrypt -2. profit -3. get third -4. exit - """) - -third = 0 -iv11 = None -iv22 = None +iv = os.urandom(16) +MK = enc3(iv + iv) +H0 = (flag + b"\x00"*16)[:16] +H1 = bytes(a ^ b for a, b in zip(H0, MK)) +pt = H1 + flag[16:] +ct = enc2(iv, pt) +print("iv:", iv.hex()) +print("ct:", ct.hex()) +print() while True: - menu() - op = input("> ").strip() - - if op == "1": - data = read_hex("pt: ") - if data is None: - print() - continue - out = enc_cfb(data) - print("ct: ", out.hex()) - print() - - elif op == "2": - if iv11 is not None and iv22 is not None: - iv1, iv2 = iv11, iv22 - iv11 = iv22 = None + try: + blob = hex_input("blob: ") + if blob is None: + print("err\n"); continue + L = len(blob) + if L == 16: + y = enc1(blob) + print("blk:", y.hex()); print() + elif L >= 32 and (L % 16) == 0: + iv, ct = blob[:16], blob[16:] + print("ok\n" if C(iv, ct) else "no\n") else: - iv1 = os.urandom(16) - iv2 = os.urandom(16) - ct = enc_cbc(flag, iv1, iv2, padd=True) - print("iv1: ", iv1.hex()) - print("iv2: ", iv2.hex()) - print("ct: ", ct.hex()) - print() - - elif op == "3": - if third: - print("sheesh") - continue - block = read_hex("pt: ") - if block is None: - print() - continue - if len(block) != 16: - print("hmmm\n") - continue - iv1 = os.urandom(16) - iv2 = os.urandom(16) - ct = enc_cbc(block, iv1, iv2, padd=False) - iv11, iv22 = iv1, iv2 - print("ct: ", ct.hex()) - third = 1 - print() - - elif op == "4": + print("err\n") + except EOFError: break - else: - print("mabokkkk?") diff --git a/services/sheesh/src/chall.py b/services/sheesh/src/chall.py index 2ad1df3..2080e20 100644 --- a/services/sheesh/src/chall.py +++ b/services/sheesh/src/chall.py @@ -1,109 +1,98 @@ #!/usr/bin/env python3 -import os -import binascii -import hashlib -import signal +import os, sys, signal, binascii, random from Crypto.Cipher import AES from Crypto.Util.Padding import pad, unpad +from Crypto.Util.number import bytes_to_long, long_to_bytes -signal.alarm(50) -seed_bits = 23 -seed_max = 1 << seed_bits -seed_len = (seed_bits + 7) // 8 -key = os.urandom(16) +random.seed(os.urandom(16)) +K0 = os.urandom(16) +K1 = os.urandom(16) +S0 = os.urandom(16) +S1 = os.urandom(16) +M0 = os.urandom(16) +M1 = os.urandom(16) -def hash_seed(seed_int: int) -> bytes: - sb = seed_int.to_bytes(seed_len, "big") - return hashlib.sha256(sb).digest()[:16] - -seed = int.from_bytes(os.urandom(4), "big") % seed_max -seed2 = int.from_bytes(os.urandom(4), "big") % seed_max -K1 = hash_seed(seed) -K2 = hash_seed(seed2) - -with open("/flag.txt", "rb") as f: +with open("/flag.txt","rb") as f: flag = f.read() -def read_hex(prompt: str): - s = input(prompt).strip() - try: - return binascii.unhexlify(s) - except Exception: - print("hmm") +def hex_input(q): + s = input(q).strip() + try: return binascii.unhexlify(s) + except: print("err"); return None + +def enc1(b16: bytes) -> bytes: + x = AES.new(K1, AES.MODE_ECB).encrypt(b16) + return bytes(a ^ b for a, b in zip(x, b16)) + +def enc2(iv: bytes, m: bytes) -> bytes: + return AES.new(K0, AES.MODE_CBC, iv=iv).encrypt(pad(m, 16)) + +def enc3(m: bytes) -> bytes: + return AES.new(K1, AES.MODE_CBC, iv=b"\x00"*16).encrypt(pad(m, 16))[-16:] + +def T(iv: bytes, ct: bytes): + n = len(ct) + if n < 96 or (n & 15): + return None + v = memoryview(ct) + W = [bytes(v[i:i+16]) for i in range(0, n, 16)] + m = len(W) + + r = ((iv[0] & 7) + 2) % m + if r: + W = W[r:] + W[:r] + + j = 1 + (W[0][0] & 1) + if len(W) <= j: + return None + del W[j] + if len(W) < 2: return None -def enc_cfb(pt: bytes) -> bytes: - iv = os.urandom(16) - aes = AES.new(key, AES.MODE_CFB, iv=iv, segment_size=128) - ct = aes.encrypt(pt) - return iv + ct + a0 = long_to_bytes(bytes_to_long(W[0]) ^ bytes_to_long(M0)) + a1 = long_to_bytes(bytes_to_long(W[1]) ^ bytes_to_long(M1)) + return b"".join((S0, a0, S1, a1, *W[2:])) -def enc_cbc(data: bytes, iv1: bytes, iv2: bytes, padd: bool) -> bytes: - x = pad(data, 16) if padd else data - c1 = AES.new(K1, AES.MODE_CBC, iv=iv1).encrypt(x) - c2 = AES.new(K2, AES.MODE_CBC, iv=iv2).encrypt(c1) - return c2 +def C(iv: bytes, ct: bytes) -> bool: + z = T(iv, ct) + if z is None: + ok = False + else: + try: + x = AES.new(K0, AES.MODE_CBC, iv=iv).decrypt(z) + unpad(x, 16) + ok = True + except: + ok = False + if random.random() < 0.08: + ok = not ok + return ok -def menu(): - print(""" -1. encrypt -2. profit -3. get third -4. exit - """) - -third = 0 -iv11 = None -iv22 = None +iv = os.urandom(16) +MK = enc3(iv + iv) +H0 = (flag + b"\x00"*16)[:16] +H1 = bytes(a ^ b for a, b in zip(H0, MK)) +pt = H1 + flag[16:] +ct = enc2(iv, pt) +print("iv:", iv.hex()) +print("ct:", ct.hex()) +print() while True: - menu() - op = input("> ").strip() - - if op == "1": - data = read_hex("pt: ") - if data is None: - print() - continue - out = enc_cfb(data) - print("ct: ", out.hex()) - print() - - elif op == "2": - if iv11 is not None and iv22 is not None: - iv1, iv2 = iv11, iv22 - iv11 = iv22 = None + try: + blob = hex_input("blob: ") + if blob is None: + print("err\n"); continue + L = len(blob) + if L == 16: + y = enc1(blob) + print("blk:", y.hex()); print() + elif L >= 32 and (L % 16) == 0: + iv, ct = blob[:16], blob[16:] + print("ok\n" if C(iv, ct) else "no\n") else: - iv1 = os.urandom(16) - iv2 = os.urandom(16) - ct = enc_cbc(flag, iv1, iv2, padd=True) - print("iv1: ", iv1.hex()) - print("iv2: ", iv2.hex()) - print("ct: ", ct.hex()) - print() - - elif op == "3": - if third: - print("sheesh") - continue - block = read_hex("pt: ") - if block is None: - print() - continue - if len(block) != 16: - print("hmmm\n") - continue - iv1 = os.urandom(16) - iv2 = os.urandom(16) - ct = enc_cbc(block, iv1, iv2, padd=False) - iv11, iv22 = iv1, iv2 - print("ct: ", ct.hex()) - third = 1 - print() - - elif op == "4": + print("err\n") + except EOFError: break - else: - print("mabokkkk?")