UMDCTF 2026 - Crypto

weave

注:最近的UMDCTF比赛上,出了一道Gabidulin code的题目,要用Welch-Berlekamp 解码算法(目前还不太会这个方向的问题,用ai解出来的…)这篇文章是在ai的辅助下的粗浅理解

Gabidulin code是一种特殊的纠错码 (看来要找个时间学一下纠错码了)

题目:

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
#!/usr/bin/env sage
import json, os
from hashlib import sha256
from Crypto.Cipher import AES

Q = 2
M = 43
N = 40
K = 8
FRAYS = 5
FIBER_D = 3

R.<x> = PolynomialRing(GF(Q))
MODPOLY = x**43 + x**21 + x**3 + x + 1
assert MODPOLY.is_irreducible()

Fq = GF(Q)
Fqm.<a> = GF(Q**M, modulus=MODPOLY) # 2的43次扩域

set_random_seed(int.from_bytes(os.urandom(32), 'big'))

def qpow(e, j):
return e ** (Q ** j) # 幂塔

def pick_independent(count): # 在有限域扩张上生成随机线性无关组
while True:
vs = [Fqm.random_element() for _ in range(count)]
rows = []
for v in vs:
cs = list(v.polynomial()) if v else []
cs = [Fq(c) for c in cs] + [Fq(0)] * (M - len(cs))
rows.append(cs)
if Matrix(Fq, rows).rank() == count: # 满秩
return vs

def pick_invertible_Fqm(sz):
while True:
M_ = Matrix(Fqm, sz, sz, [Fqm.random_element() for _ in range(sz * sz)])
if M_.is_invertible():
return M_

def pick_shuttle(sz, fibers):
while True:
entries = []
for _ in range(sz * sz):
cs = [Fq.random_element() for _ in range(FIBER_D)]
entries.append(sum(c * v for c, v in zip(cs, fibers)))
M_ = Matrix(Fqm, sz, sz, entries)
if M_.is_invertible():
return M_

pegs = pick_independent(N)
loom = Matrix(Fqm, K, N, lambda j, i: qpow(pegs[i], j))
knot = pick_invertible_Fqm(K)
fibers = pick_independent(FIBER_D)
shuttle = pick_shuttle(N, fibers)
warp = knot * loom * shuttle.inverse()

secret = vector(Fqm, [Fqm.random_element() for _ in range(K)])

def pack(v):
out = []
for e in v:
cs = list(e.polynomial()) if e else []
cs = [int(c) for c in cs] + [0] * (M - len(cs))
val = 0
for i, c in enumerate(cs):
val |= c << i
out.append(int(val))
return out

def pack_mat(M_):
return [pack(row) for row in M_.rows()]

secret_bytes = b''.join(int(v).to_bytes((M + 7) // 8, 'big') for v in pack(secret))
wrap_key = sha256(secret_bytes).digest()[:16]

try:
flag = open('flag.txt', 'rb').read().strip()
except FileNotFoundError:
flag = b'UMDCTF{test_flag}'

iv = os.urandom(12)
body, tag = AES.new(wrap_key, AES.MODE_GCM, nonce=iv).encrypt_and_digest(flag)

def pick_frays():
while True:
B = Matrix(Fq, 5, 40, [Fq.random_element() for _ in range(5 * 40)])
if B.rank() == 5:
break
u = vector(Fqm, [Fqm.random_element() for _ in range(5)])
return u * B.change_ring(Fqm)

frays = pick_frays()
bolt = secret * warp + frays

handout = {
'spec': {
'q': int(Q),
'm': int(M),
'n': int(N),
'k': int(K),
'frays': int(FRAYS),
'modulus': [int(c) for c in MODPOLY.list()],
},
'warp': pack_mat(warp),
'bolt': pack(bolt),
'loom': {
'pegs': pack(vector(Fqm, pegs)),
'knot': pack_mat(knot),
'shuttle': pack_mat(shuttle),
'fibers': pack(vector(Fqm, fibers)),
},
'vault': {
'iv': iv.hex(),
'body': body.hex(),
'tag': tag.hex(),
},
}

with open('output.json', 'w') as fh:
json.dump(handout, fh)

print('Output.json is created!')

题目关键:
引入了一个错误定位多项式 $\Lambda(x)$
$$
\begin{aligned}
& warp = knot loomshuttle^{-1} \
& bolt = secretwarp+frays \
& 题目已知的是warp、bolt、knot、shuttle \
& shuttle的秩是5,frays的秩是3,secret的长度是8 \
&注意:frays的具体是不知道的 \
\
&上面俩个式子化为(左右同乘shuttle):\
& shuttle
bolt=secretknotloom+frays*shuttle \
& 记为:B’ = S *loom + E \
\
& 规定 W(x)=\Lambda(L(x))、\Lambda(E)=0、L(x)是S的线性多项式(有8个未知量) \
&
\end{aligned}
$$

脚本:
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
import json
import hashlib
from Crypto.Cipher import AES

with open('output.json', 'r') as fh:
handout = json.load(fh)

spec = handout['spec']
Q = spec['q']
M = spec['m']
N = spec['n']
K = spec['k']
FRAYS = spec['frays']
modulus_coeffs = spec['modulus']

P.<x> = PolynomialRing(GF(Q))
MODPOLY = P(modulus_coeffs)
Fq = GF(Q)
Fqm.<a> = GF(Q**M, modulus=MODPOLY, name='a')

def unpack(data):
elements = []
for val in data:
coeffs = [Fq((val >> i) & 1) for i in range(M)]
elements.append(Fqm(P(coeffs)))
return vector(Fqm, elements)

def unpack_mat(data):
rows = [unpack(row) for row in data]
return Matrix(Fqm, rows)

warp = unpack_mat(handout['warp'])
bolt = unpack(handout['bolt'])

loom_data = handout['loom']
pegs = list(unpack(loom_data['pegs']))
knot = unpack_mat(loom_data['knot'])
shuttle = unpack_mat(loom_data['shuttle'])
fibers = list(unpack(loom_data['fibers']))
FIBER_D = len(fibers)

vault = handout['vault']
iv = bytes.fromhex(vault['iv'])
body = bytes.fromhex(vault['body'])
tag = bytes.fromhex(vault['tag'])

# 1. 转换到 B' = S * loom + E
B_prime = bolt * shuttle

# 2. 构造 Welch-Berlekamp 线性方程组
# W(x) 阶数 = (K - 1) + 15 = 22 (23 个未知数 w_0 ... w_22)
# Lambda(x) 阶数 = 15 (假设 lambda_0 = 1, 有 15 个未知数 lambda_1 ... lambda_15)
# 总共 38 个未知数,利用 N=40 个点进行评估,必然有唯一解
mat_rows = []
vec_rhs = []

for i in range(N):
row = []
# 填充 W(x) 的系数项 w_j
for j in range(23):
row.append(pegs[i] ** (2**j))
# 填充 Lambda(x) 的系数项 lambda_k (由于在GF(2)上减法等于加法,直接用加法)
for k in range(1, 16):
row.append( B_prime[i] ** (2**k) )

mat_rows.append(row)
# 等式右边为 lambda_0 * B'_i,由于我们设 lambda_0 = 1,所以常数项就是 B'_i
vec_rhs.append(B_prime[i])

Mat = Matrix(Fqm, mat_rows)
Vec = vector(Fqm, vec_rhs)

X = Mat.solve_right(Vec)

W_coeffs = list(X[:23])
Lambda_coeffs = [Fqm(1)] + list(X[23:])

# 3. 递归求解 L(x) 的系数 S_0 ... S_7
# 由关系式 W(x) = Lambda(L(x)) 可知可以由低到高逐级求出 S 的各项系数
S = []
for m in range(8):
sm = W_coeffs[m]
for k in range(1, m + 1):
sm -= Lambda_coeffs[k] * (S[m - k] ** (2**k))
S.append(sm)

# 4. 恢复真正的 secret 向量
S_vec = vector(Fqm, S)
secret = S_vec * knot.inverse()

# 5. 打包并计算 AES Key 进行解密
def pack(v):
out = []
for e in v:
cs = list(e.polynomial()) if e else []
cs = [int(c) for c in cs] + [0] * (M - len(cs))
val = 0
for i, c in enumerate(cs):
val |= c << i
out.append(int(val))
return out

secret_packed = pack(secret)
secret_bytes = b''.join(int(v).to_bytes((M + 7) // 8, 'big') for v in secret_packed)
wrap_key = hashlib.sha256(secret_bytes).digest()[:16]

cipher = AES.new(wrap_key, AES.MODE_GCM, nonce=iv)

m = cipher.decrypt_and_verify(body, tag)
print(m)

no-brainrot-allowed

参考wp:https://medium.com/@mazenmagdy1598/no-brainrot-allowed-umdctf-crypto-0c38205cbc26

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
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
#!/usr/bin/env python3
from pwn import *

HOST = "challs.umdctf.io"
PORT = 32767

context.log_level = "info"

# RSA 公钥参数
n = 89496838321330017124211425752928111009238414395285545597372895783391482460166014550795440784240669454038164776392492949832230406030665778241454645944939829559549747525412818621247626093163657213524408194055221128159991890855776297338418179985226639927931716465641085590302394062423554511419578835789906477703
e = 65537

# 已知 flag 的密文 ct = flag^e mod n
ct = 7754782549233547741892262011884269269634473224225227064848605234096464292342695844400918832869742989785685496372442722948589824059885664742180188925993430350247652395812127146595142859972102395302095473677093880196683037670451512853001503582104512714892761518926915267957380484576367984853786495267989619184

# 服务端判断:
# pt = hex(m)
# if pt.startswith("0x67"):
# 触发报错
#
# 所以命中条件是:解密出来的整数 m 的十六进制最高字节是 0x67
# 若总 hex 长度固定为 256 位(不含 0x),那么满足:
# A <= m < B
# 其中:
# A = 0x67 * 16^254
# B = 0x68 * 16^254
A = 103 * 16**254
B = 104 * 16**254

# 已知 flag 的前缀和总长度 这个111的长度应该是wp的作者通过测试得到flag后,直接设置为111长度的
# 所以初始时可以把 flag 限定在 [L, U) 这个整数区间里
PREFIX = b"UMDCTF{"
FLAG_LEN = 111

# 经验参数:
# ALPHA 用来控制 s_target 的大小
# 它不是严格数学常数,而是命中率 / 收缩效率之间的折中参数
ALPHA = 16

# 每轮在候选 s 区间里先粗采样 32 个点
NUM_SAMPLES = 32

# 如果一个都没命中,再更密一点采样 48 个点
DENSE_SAMPLES = 48

# 最多做多少轮 batch
MAX_BATCHES = 220


def connect():
p = remote(HOST, PORT)
p.recvuntil(b"Your messages:")
return p


def query_batch(p, s_values):
"""
一次性测试多个 s 值。

对于每个 s,我们发送:
ct' = ct * s^e mod n

因为 RSA 乘法同态:
Dec(ct') = Dec(ct) * Dec(s^e) mod n
= flag * s mod n

服务端实际上在检查:
hex(flag * s mod n).startswith("0x67")
"""

# 把多个变形后的密文用逗号拼起来,一次发给服务端
payload = ",".join(str((ct * pow(s, e, n)) % n) for s in s_values)
p.sendline(payload.encode())

# 服务端会对每个输入给一行响应,最后再次打印 "Your messages:"
text = p.recvuntil(b"Your messages:", drop=False).decode(errors="replace")

responses = []
for line in text.splitlines():
# 如果命中,服务端返回报错
if "ERROR: BRAINROT DETECTED" in line:
responses.append(True)

# 否则表示没有命中
elif "thanks you for your message" in line:
responses.append(False)

# 只保留那些命中的 s
return [s for s, ok in zip(s_values, responses) if ok]


def to_bytes(x):
"""
把整数转回字节串。
"""
hx = hex(x)[2:]
if len(hx) % 2:
hx = "0" + hx
return bytes.fromhex(hx)


# ========================
# 主逻辑(基本不变)
# ========================

def main():
# prefix_int 是前缀 "UMDCTF{" 对应的大整数
prefix_int = int.from_bytes(PREFIX, "big")

# 构造初始区间 [L, U)
#
# L: 前缀固定,后面全补 0x00
# U: 相当于"前缀 + 1",后面全补 0x00
#
# 因此所有以 PREFIX 开头、总长为 FLAG_LEN 的字符串,
# 转成整数后都会落在 [L, U) 中
L = prefix_int * 256 ** (FLAG_LEN - len(PREFIX))
U = (prefix_int + 1) * 256 ** (FLAG_LEN - len(PREFIX))

# 建立远程连接
p = connect()

# 进行多轮区间收缩
for batch in range(1, MAX_BATCHES + 1):
# 当前 flag 仍可能落在 [L, U)
width = U - L

# 如果区间宽度已经 <= 1,说明基本定位完成
if width <= 1:
break

# 用中点近似真实 flag
mid = (L + U) // 2

# 选择一个"理想量级"的 s
#
# 设计思路:
# 希望当前候选区间 [L, U) 乘上 s 后,
# 宽度约为 ALPHA * (B - A)
#
# 即:
# s * (U - L) ≈ ALPHA * (B - A)
#
# 解得:
# s ≈ ALPHA * (B - A) / width
#
# 而这里 B - A = 16^254
s_target = ALPHA * 16**254 // width

# 估计对应的 k
#
# 命中时满足:
# A <= flag*s - k*n < B
#
# 即:
# flag*s ≈ k*n
#
# 因为 flag 不知道,就用 mid 近似:
# mid*s ≈ k*n
# k ≈ mid*s / n
#
# 取最近整数
k = round(s_target * mid / n)

# 对于固定的 k,希望存在某个 flag ∈ [L, U) 使得:
# A <= flag*s - k*n < B
#
# 改写为:
# k*n + A <= flag*s < k*n + B
#
# 又因为 flag ∈ [L, U),所以 flag*s ∈ [L*s, U*s)
#
# 要让这两个区间有交集,需要 s 落在某个范围 [s_lo, s_hi]
# 这就是下面两个式子
s_lo = (k * n + A + U - 1) // U # 向上取整
s_hi = (k * n + B - 1) // L # 向下取整

# 如果范围为空,说明当前估计出的 k 不合理
if s_hi < s_lo:
raise RuntimeError(f"empty s-interval at batch {batch}")

# 候选 s 区间总长度
total = s_hi - s_lo

# 在 [s_lo, s_hi] 中均匀采样 NUM_SAMPLES 个点
#
# 注意这不是枚举所有 s,只是抽样测试
s_values = [s_lo + (total * i) // (NUM_SAMPLES - 1) for i in range(NUM_SAMPLES)]

try:
# 批量测试这些 s,看看哪些会触发 oracle
hits = query_batch(p, s_values)
except EOFError:
# 远程有时会断连,重连后继续
log.warning("reconnecting...")
p.close()
p = connect()
hits = query_batch(p, s_values)

# 如果粗采样一个都没中,就更密一点再试一次
if not hits:
s_values = [s_lo + (total * i) // (DENSE_SAMPLES - 1) for i in range(DENSE_SAMPLES)]

try:
hits = query_batch(p, s_values)
except EOFError:
log.warning("reconnecting...")
p.close()
p = connect()
hits = query_batch(p, s_values)

# 如果还是一个命中都没有,说明这一轮策略失败
if not hits:
raise RuntimeError(f"no positive hits at batch {batch}")

# ========================
# 利用命中的 s 收缩区间
# ========================
#
# 对每个命中的 s,都有:
# A <= flag*s - k*n < B
#
# 于是可推出:
# (k*n + A)/s <= flag < (k*n + B)/s
#
# 再和原区间 [L, U) 取交,得到更小的新范围
for s in hits:
# 这里重新估计 k
#
# 因为真实 flag >= L,所以:
# k = floor(flag*s / n)
# 至少可以先用 floor(L*s / n) 作为对应分支
#
# 这一步是题解/脚本里的经验写法。
k = (s * L) // n

# 从
# flag >= (k*n + A)/s
# 得到下界,注意要向上取整
new_L = max(L, (k * n + A + s - 1) // s)

# 从
# flag < (k*n + B)/s
# 得到上界
#
# 这里写成:
# floor((k*n + B - 1)/s) + 1
# 是为了保持区间仍然是左闭右开 [new_L, new_U)
new_U = min(U, (k * n + B - 1) // s + 1)

# 如果新区间空了,说明哪步估计出了问题
if not (new_L < new_U):
raise RuntimeError("interval collapsed")

# 更新当前区间
L, U = new_L, new_U

# 定期打印当前进度
if batch == 1 or batch % 5 == 0:
log.info(f"batch={batch} hits={len(hits)} bits={(U - L).bit_length()}")

p.close()

# 最后区间已经很小
log.success(f"final width = {U - L}")
log.success(f"final bits = {(U - L).bit_length()}")

# 在最终区间里暴力枚举,找出真正满足 m^e mod n = ct 的那个 m
print(f'[L,U) = [{L},{U})')
for m in range(L, U):
if pow(m, e, n) == ct:
flag = to_bytes(m)
log.success(f"FLAG = {flag.decode()}")
return

log.error("flag not found")
exit(1)


if __name__ == "__main__":
main()

UMDCTF 2026 - Crypto
https://baymax-fools.github.io/2026/04/27/crypto/UMDCTF 2026 - Crypto/
Author
Baymax
Posted on
April 27, 2026
Updated on
June 17, 2026
Licensed under