This repository was archived by the owner on Nov 23, 2025. It is now read-only.
-
Notifications
You must be signed in to change notification settings - Fork 6
Expand file tree
/
Copy pathchall.py
More file actions
180 lines (157 loc) 路 5.61 KB
/
Copy pathchall.py
File metadata and controls
180 lines (157 loc) 路 5.61 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
from Crypto.Util.number import bytes_to_long
from hashlib import shake_128
from ast import literal_eval
from secrets import token_bytes
from math import floor, ceil, log2
import os
FLAG = os.getenv("FLAG", "SEKAI{}")
m = 256
w = 21
n = 128
l1 = ceil(m / log2(w))
l2 = floor(log2(l1*(w-1)) / log2(w)) + 1
l = l1 + l2
class WOTS:
def __init__(self):
self.sk = [token_bytes(n // 8) for _ in range(l)]
self.pk = [WOTS.chain(sk, w - 1) for sk in self.sk]
def sign(self, digest: bytes) -> list[bytes]:
assert 8 * len(digest) == m
d1 = WOTS.pack(bytes_to_long(digest), l1, w)
checksum = sum(w-1-i for i in d1)
d2 = WOTS.pack(checksum, l2, w)
d = d1 + d2
sig = [WOTS.chain(self.sk[i], w - d[i] - 1) for i in range(l)]
return sig
def get_pubkey_hash(self) -> bytes:
hasher = shake_128(b"\x04")
for i in range(l):
hasher.update(self.pk[i])
return hasher.digest(16)
@staticmethod
def pack(num: int, length: int, base: int) -> list[int]:
packed = []
while num > 0:
packed.append(num % base)
num //= base
if len(packed) < length:
packed += [0] * (length - len(packed))
return packed
@staticmethod
def chain(x: bytes, n: int) -> bytes:
if n == 0:
return x
x = shake_128(b"\x03" + x).digest(16)
return WOTS.chain(x, n - 1)
@staticmethod
def verify(digest: bytes, sig: list[bytes]) -> bytes:
d1 = WOTS.pack(bytes_to_long(digest), l1, w)
checksum = sum(w-1-i for i in d1)
d2 = WOTS.pack(checksum, l2, w)
d = d1 + d2
sig_pk = [WOTS.chain(sig[i], d[i]) for i in range(l)]
hasher = shake_128(b"\x04")
for i in range(len(sig_pk)):
hasher.update(sig_pk[i])
sig_hash = hasher.digest(16)
return sig_hash
class MerkleTree:
def __init__(self, height: int = 8):
self.h = height
self.keys = [WOTS() for _ in range(2**height)]
self.tree = []
self.root = self.build_tree([key.get_pubkey_hash() for key in self.keys])
def build_tree(self, leaves: list[bytes]) -> bytes:
self.tree.append(leaves)
if len(leaves) == 1:
return leaves[0]
parents = []
for i in range(0, len(leaves), 2):
left = leaves[i]
if i + 1 < len(leaves):
right = leaves[i + 1]
else:
right = leaves[i]
hasher = shake_128(b"\x02" + left + right).digest(16)
parents.append(hasher)
return self.build_tree(parents)
def sign(self, index: int, digest: bytes) -> list:
assert 0 <= index < len(self.keys)
key = self.keys[index]
wots_sig = key.sign(digest)
sig = [wots_sig]
for i in range(self.h):
leaves = self.tree[i]
u = index >> i
if u % 2 == 0:
if u + 1 < len(leaves):
sig.append((0, leaves[u + 1]))
else:
sig.append((0, leaves[u]))
else:
sig.append((1, leaves[u - 1]))
return sig
@staticmethod
def verify(sig: list, digest: bytes) -> bytes:
wots_sig = sig[0]
sig = sig[1:]
pk_hash = WOTS.verify(digest, wots_sig)
root_hash = pk_hash
for (side, leaf) in sig:
if side == 0:
root_hash = shake_128(b"\x02" + root_hash + leaf).digest(16)
else:
root_hash = shake_128(b"\x02" + leaf + root_hash).digest(16)
return root_hash
class Challenge:
def __init__(self, h: int = 8):
self.h = h
self.max_signs = 2 ** h - 1
self.tree = MerkleTree(h)
self.root = self.tree.root
self.used = set()
self.before_input = f"public key: {self.root.hex()}"
def sign(self, num_sign: int, inds: list, messages: list):
assert num_sign + len(self.used) <= self.max_signs
assert len(inds) == len(set(inds)) == len(messages) == num_sign
assert self.used.isdisjoint(inds)
assert all(b"flag" not in msg for msg in messages)
sigs = []
for i in range(num_sign):
digest = shake_128(b"\x00" + messages[i]).digest(32)
sigs.append(self.tree.sign(inds[i], digest))
self.used.update(inds)
return sigs
def next(self):
new_tree = MerkleTree(self.h)
digest = shake_128(b"\x01" + new_tree.root).digest(32)
index = next(i for i in range(2 ** self.h) if i not in self.used)
sig = new_tree.sign(index, digest)
self.tree = new_tree
return {
"root": new_tree.root,
"sig": sig,
"index": index,
}
def verify(self, sig: list, message: bytes):
digest = shake_128(b"\x00" + message).digest(32)
for i, s in enumerate(reversed(sig)):
if i != 0:
digest = shake_128(b"\x01" + digest).digest(32)
digest = MerkleTree.verify(s, digest)
return digest == self.root
def get_flag(self, sig: list):
if not self.verify(sig, b"Give me the flag"):
return {"message": "Invalid signature"}
else:
return {"message": f"Congratulations! Here is your flag: {FLAG}"}
def __call__(self, type: str, **kwargs):
assert type in ["sign", "next", "get_flag"]
return getattr(self, type)(**kwargs)
challenge = Challenge()
print(challenge.before_input)
try:
while True:
print(challenge(**literal_eval(input("input: "))))
except:
exit(1)