Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_crypto/tools/linkring_oracle.py

15.0 KiB, 1 run

created by r1870400018:61087, which is this file's identity for as long as the history lasts, whatever it is later renamed to

download · who wrote it · its history

1#!/usr/bin/env python3
2"""Independent reference implementation of linkring/1, used only as a test
3oracle for fe2o3_crypto::linkring. Never a runtime dependency.
4
5Everything here is written from the specifications, not from the Rust code:
6ristretto255 from RFC 9496, hash_to_ristretto255 from RFC 9380, and the
7one-out-of-many proof from Bootle et al. (ESORICS 2015) with the tag relation
8described in the Rust module header. The prover computes each G_k by plain
9polynomial expansion over every padded entry, with none of the Rust prover's
10subset expansion or padding fold, so agreement is evidence rather than
11self-consistency.
12
13Usage:
14 linkring_oracle.py selftest
15 linkring_oracle.py sign < cases (case blocks, see parse_blocks)
16 linkring_oracle.py verify < cases
17
18The vectors in tests/data/linkring_vectors.txt are `sign` output.
19"""
20import hashlib
21import sys
22
23# ── Field and group (RFC 9496) ──────────────────────────────────────────────
24
25P = 2**255 - 19
26L = 2**252 + 27742317777372353535851937790883648493
27
28
29def inv(a):
30 return pow(a % P, P - 2, P)
31
32
33D = (-121665 * inv(121666)) % P
34SQRT_M1 = pow(2, (P - 1) // 4, P)
35
36
37def is_neg(a):
38 return (a % P) & 1
39
40
41def ct_abs(a):
42 a %= P
43 return (P - a) % P if is_neg(a) else a
44
45
46def sqrt_ratio_m1(u, v):
47 u %= P
48 v %= P
49 v3 = v * v % P * v % P
50 v7 = v3 * v3 % P * v % P
51 r = (u * v3) % P * pow(u * v7 % P, (P - 5) // 8, P) % P
52 check = v * r % P * r % P
53 correct = check == u
54 flipped = check == (-u) % P
55 flipped_i = check == (-u * SQRT_M1) % P
56 if flipped or flipped_i:
57 r = r * SQRT_M1 % P
58 r = ct_abs(r)
59 return (correct or flipped), r
60
61
62def _sqrt(a):
63 ok, r = sqrt_ratio_m1(a, 1)
64 assert ok
65 return r
66
67
68# RFC 9496 section 4.1 constants, checked against their defining relations.
69SQRT_AD_MINUS_ONE = 25063068953384623474111414158702152701244531502492656460079210482610430750235
70INVSQRT_A_MINUS_D = 54469307008909316920995813868745141605393597292927456921205312896311721017578
71ONE_MINUS_D_SQ = (1 - D * D) % P
72D_MINUS_ONE_SQ = (D - 1) * (D - 1) % P
73assert SQRT_AD_MINUS_ONE * SQRT_AD_MINUS_ONE % P == (-D - 1) % P
74assert INVSQRT_A_MINUS_D * INVSQRT_A_MINUS_D % P * ((-1 - D) % P) % P == 1
75
76IDENT = (0, 1, 1, 0)
77
78
79def add(p1, p2):
80 x1, y1, z1, t1 = p1
81 x2, y2, z2, t2 = p2
82 a = (y1 - x1) * (y2 - x2) % P
83 b = (y1 + x1) * (y2 + x2) % P
84 c = t1 * 2 * D % P * t2 % P
85 d = z1 * 2 * z2 % P
86 e, f, g, h = b - a, d - c, d + c, b + a
87 return (e * f % P, g * h % P, f * g % P, e * h % P)
88
89
90def neg(p1):
91 x, y, z, t = p1
92 return ((-x) % P, y, z, (-t) % P)
93
94
95def mul(k, pt):
96 k %= L
97 acc = IDENT
98 for bit in reversed(range(k.bit_length())):
99 acc = add(acc, acc)
100 if (k >> bit) & 1:
101 acc = add(acc, pt)
102 return acc
103
104
105def eq(p1, p2):
106 x1, y1, _, _ = p1
107 x2, y2, _, _ = p2
108 return (x1 * y2 - y1 * x2) % P == 0 or (y1 * y2 - x1 * x2) % P == 0
109
110
111def is_ident(pt):
112 return eq(pt, IDENT)
113
114
115def encode(pt):
116 x0, y0, z0, t0 = pt
117 u1 = (z0 + y0) * (z0 - y0) % P
118 u2 = x0 * y0 % P
119 _, invsqrt = sqrt_ratio_m1(1, u1 * u2 % P * u2 % P)
120 den1 = invsqrt * u1 % P
121 den2 = invsqrt * u2 % P
122 z_inv = den1 * den2 % P * t0 % P
123 ix0 = x0 * SQRT_M1 % P
124 iy0 = y0 * SQRT_M1 % P
125 ench = den1 * INVSQRT_A_MINUS_D % P
126 rotate = is_neg(t0 * z_inv)
127 if rotate:
128 x, y, den_inv = iy0, ix0, ench
129 else:
130 x, y, den_inv = x0, y0, den2
131 if is_neg(x * z_inv):
132 y = (-y) % P
133 s = ct_abs(den_inv * (z0 - y))
134 return s.to_bytes(32, "little")
135
136
137def decode(b):
138 if len(b) != 32:
139 return None
140 s = int.from_bytes(b, "little")
141 if s >= P or is_neg(s):
142 return None
143 ss = s * s % P
144 u1 = (1 - ss) % P
145 u2 = (1 + ss) % P
146 u2s = u2 * u2 % P
147 v = (-(D * u1 % P * u1) - u2s) % P
148 ok, invsqrt = sqrt_ratio_m1(1, v * u2s % P)
149 den_x = invsqrt * u2 % P
150 den_y = invsqrt * den_x % P * v % P
151 x = ct_abs(2 * s * den_x)
152 y = u1 * den_y % P
153 t = x * y % P
154 if not ok or is_neg(t) or y == 0:
155 return None
156 return (x, y, 1, t)
157
158
159def _map(t):
160 r = SQRT_M1 * t % P * t % P
161 u = (r + 1) * ONE_MINUS_D_SQ % P
162 v = (-1 - r * D) * (r + D) % P
163 was_square, s = sqrt_ratio_m1(u, v)
164 s_prime = (-ct_abs(s * t)) % P
165 if not was_square:
166 s = s_prime
167 c = P - 1 if was_square else r
168 n = (c * (r - 1) % P * D_MINUS_ONE_SQ - v) % P
169 w0 = 2 * s * v % P
170 w1 = n * SQRT_AD_MINUS_ONE % P
171 w2 = (1 - s * s) % P
172 w3 = (1 + s * s) % P
173 return (w0 * w3 % P, w2 * w1 % P, w1 * w3 % P, w0 * w2 % P)
174
175
176def from_uniform(b64):
177 r0 = int.from_bytes(b64[:32], "little") & ((1 << 255) - 1)
178 r1 = int.from_bytes(b64[32:], "little") & ((1 << 255) - 1)
179 return add(_map(r0 % P), _map(r1 % P))
180
181
182def _base():
183 y = 4 * inv(5) % P
184 x = _sqrt((y * y - 1) * inv(D * y * y + 1))
185 if is_neg(x):
186 x = P - x
187 return (x, y, 1, x * y % P)
188
189
190G = _base()
191
192# ── Hashing (RFC 9380) ──────────────────────────────────────────────────────
193
194
195def xmd_sha512_64(msg, dst):
196 dst_prime = dst + bytes([len(dst)])
197 b0 = hashlib.sha512(bytes(128) + msg + (64).to_bytes(2, "big") + b"\x00" + dst_prime).digest()
198 b1 = hashlib.sha512(b0 + b"\x01" + dst_prime).digest()
199 return b1
200
201
202def hash_to_group(dst, msg):
203 return from_uniform(xmd_sha512_64(msg, dst))
204
205
206def wide(b):
207 return int.from_bytes(b, "little") % L
208
209
210def sc_bytes(s):
211 return (s % L).to_bytes(32, "little")
212
213
214def u32(v):
215 return v.to_bytes(4, "little")
216
217
218def u64(v):
219 return v.to_bytes(8, "little")
220
221# ── linkring/1 ──────────────────────────────────────────────────────────────
222
223
224N_RADIX = 16
225
226
227def digits(n):
228 m, cap = 1, N_RADIX
229 while cap < n:
230 m, cap = m + 1, cap * N_RADIX
231 return m
232
233
234def secret_from_seed(seed):
235 return wide(hashlib.sha512(b"linkring/1:secret" + seed).digest())
236
237
238def gens(m):
239 return [hash_to_group(b"linkring/1:gen", u32(i)) for i in range(m * N_RADIX)]
240
241
242def scope_base(scope):
243 return hash_to_group(b"linkring/1:scope", scope)
244
245
246def ring_digest(keys):
247 return hashlib.sha256(b"".join(keys)).digest()
248
249
250def challenge(n, digest, m, scope, tag, msg, pts):
251 h = hashlib.sha512()
252 h.update(b"linkring/1" + u32(N_RADIX) + u32(m) + u64(n) + digest)
253 h.update(u64(len(scope)) + scope + tag + u64(len(msg)) + msg)
254 for p in pts:
255 h.update(encode(p))
256 return wide(h.digest())
257
258
259def msm(pairs):
260 acc = IDENT
261 for s, p in pairs:
262 acc = add(acc, mul(s, p))
263 return acc
264
265
266def poly_mul_linear(poly, c1, c0):
267 # poly · (c1·x + c0), coefficients low to high.
268 out = [0] * (len(poly) + 1)
269 for k, a in enumerate(poly):
270 out[k] = (out[k] + a * c0) % L
271 out[k + 1] = (out[k + 1] + a * c1) % L
272 return out
273
274
275def sign(keys, secret, l, scope, msg, aux):
276 n = len(keys)
277 m = digits(n)
278 pts = [decode(k) for k in keys]
279 digest = ring_digest(keys)
280 u = scope_base(scope)
281 tau = mul(secret, u)
282 tag = encode(tau)
283 hs = gens(m)
284 seed = hashlib.sha512(b"linkring/1:nonce" + sc_bytes(secret) + digest + u64(n) + u64(l)
285 + u64(len(scope)) + scope + u64(len(msg)) + msg + u64(len(aux)) + aux).digest()
286 ctr = [0]
287
288 def nxt():
289 v = wide(hashlib.sha512(seed + u32(ctr[0])).digest())
290 ctr[0] += 1
291 return v
292
293 ld = [(l >> (4 * j)) & 15 for j in range(m)]
294 a = [[0] * N_RADIX for _ in range(m)]
295 for j in range(m):
296 for i in range(1, N_RADIX):
297 a[j][i] = nxt()
298 a[j][0] = (-sum(a[j][1:])) % L
299 r_a, r_b, r_c, r_d = nxt(), nxt(), nxt(), nxt()
300 rho = [nxt() for _ in range(m)]
301 sig = [[1 if ld[j] == i else 0 for i in range(N_RADIX)] for j in range(m)]
302
303 def com(r, vals):
304 return msm([(r, G)] + [(vals[j][i], hs[j * N_RADIX + i]) for j in range(m) for i in range(N_RADIX)])
305
306 A = com(r_a, a)
307 B = com(r_b, sig)
308 C = com(r_c, [[a[j][i] * (1 - 2 * sig[j][i]) for i in range(N_RADIX)] for j in range(m)])
309 Dc = com(r_d, [[-a[j][i] * a[j][i] for i in range(N_RADIX)] for j in range(m)])
310
311 # Plain expansion over every padded entry; entry i ≥ n is keys[n-1].
312 coef = [[0] * m for _ in range(n)]
313 for i in range(N_RADIX ** m):
314 poly = [1]
315 for j in range(m):
316 d = (i >> (4 * j)) & 15
317 poly = poly_mul_linear(poly, 1 if d == ld[j] else 0, a[j][d])
318 tgt = min(i, n - 1)
319 for k in range(m):
320 coef[tgt][k] = (coef[tgt][k] + poly[k]) % L
321 if i == l:
322 assert poly[m] == 1
323 else:
324 assert poly[m] == 0
325 Gk = [msm([(coef[i][k], pts[i]) for i in range(n)] + [(rho[k], G)]) for k in range(m)]
326 Yk = [mul(rho[k], u) for k in range(m)]
327 x = challenge(n, digest, m, scope, tag, msg, [A, B, C, Dc] + Gk + Yk)
328 body = bytes([m]) + b"".join(encode(p) for p in [A, B, C, Dc] + Gk + Yk)
329 for j in range(m):
330 for i in range(1, N_RADIX):
331 body += sc_bytes(sig[j][i] * x + a[j][i])
332 z = secret * pow(x, m, L) - sum(rho[k] * pow(x, k, L) for k in range(m))
333 body += sc_bytes(r_b * x + r_a) + sc_bytes(r_c * x + r_d) + sc_bytes(z)
334 return tag, body
335
336
337def verify(keys, scope, msg, tag, body):
338 n = len(keys)
339 m = digits(n)
340 if len(body) != 1 + 32 * (7 + 17 * m) or body[0] != m:
341 return False
342 pts = [decode(k) for k in keys]
343 if any(p is None or is_ident(p) for p in pts):
344 return False
345 tau = decode(tag)
346 if tau is None or is_ident(tau):
347 return False
348 off = 1
349 cp = []
350 for _ in range(4 + 2 * m):
351 p = decode(body[off:off + 32])
352 if p is None:
353 return False
354 cp.append(p)
355 off += 32
356 sv = []
357 for _ in range(m * 15 + 3):
358 s = int.from_bytes(body[off:off + 32], "little")
359 if s >= L:
360 return False
361 sv.append(s)
362 off += 32
363 A, B, C, Dc = cp[:4]
364 Gk, Yk = cp[4:4 + m], cp[4 + m:]
365 z_a, z_c, z = sv[m * 15:]
366 x = challenge(n, ring_digest(keys), m, scope, tag, msg, cp)
367 f = [[0] * N_RADIX for _ in range(m)]
368 for j in range(m):
369 for i in range(1, N_RADIX):
370 f[j][i] = sv[j * 15 + i - 1]
371 f[j][0] = (x - sum(f[j][1:])) % L
372 hs = gens(m)
373
374 def com(r, vals):
375 return msm([(r, G)] + [(vals[j][i], hs[j * N_RADIX + i]) for j in range(m) for i in range(N_RADIX)])
376
377 if not eq(add(A, mul(x, B)), com(z_a, f)):
378 return False
379 if not eq(add(mul(x, C), Dc), com(z_c, [[f[j][i] * (x - f[j][i]) for i in range(N_RADIX)] for j in range(m)])):
380 return False
381 u = scope_base(scope)
382 lhs = mul(pow(x, m, L), tau)
383 for k in range(m):
384 lhs = add(lhs, neg(mul(pow(x, k, L), Yk[k])))
385 if not eq(lhs, mul(z, u)):
386 return False
387 per = [0] * n
388 for i in range(N_RADIX ** m):
389 p = 1
390 for j in range(m):
391 p = p * f[j][(i >> (4 * j)) & 15] % L
392 per[min(i, n - 1)] = (per[min(i, n - 1)] + p) % L
393 lhs = msm([(per[i], pts[i]) for i in range(n)])
394 for k in range(m):
395 lhs = add(lhs, neg(mul(pow(x, k, L), Gk[k])))
396 return eq(lhs, mul(z, G))
397
398# ── Self-test (RFC 9496 A.1 published vectors) ──────────────────────────────
399
400
401RFC9496_MULTIPLES = [
402 "0000000000000000000000000000000000000000000000000000000000000000",
403 "e2f2ae0a6abc4e71a884a961c500515f58e30b6aa582dd8db6a65945e08d2d76",
404 "6a493210f7499cd17fecb510ae0cea23a110e8d5b901f8acadd3095c73a3b919",
405 "94741f5d5d52755ece4f23f044ee27d5d1ea1e2bd196b462166b16152a9d0259",
406 "da80862773358b466ffadfe0b3293ab3d9fd53c5ea6c955358f568322daf6a57",
407 "e882b131016b52c1d3337080187cf768423efccbb517bb495ab812c4160ff44e",
408 "f64746d3c92b13050ed8d80236a7f0007c3b3f962f5ba793d19a601ebb1df403",
409 "44f53520926ec81fbd5a387845beb7df85a96a24ece18738bdcfa6a7822a176d",
410 "903293d8f2287ebe10e2374dc1a53e0bc887e592699f02d077d5263cdd55601c",
411 "02622ace8f7303a31cafc63f8fc48fdc16e1c8c8d234b2f0d6685282a9076031",
412 "20706fd788b2720a1ed2a5dad4952b01f413bcf0e7564de8cdc816689e2db95f",
413 "bce83f8ba5dd2fa572864c24ba1810f9522bc6004afe95877ac73241cafdab42",
414 "e4549ee16b9aa03099ca208c67adafcafa4c3f3e4e5303de6026e3ca8ff84460",
415 "aa52e000df2e16f55fb1032fc33bc42742dad6bd5a8fc0be0167436c5948501f",
416 "46376b80f409b29dc2b5f6f0c52591990896e5716f41477cd30085ab7f10301e",
417 "e0c418f7c8d9c4cdd7395b93ea124f3ad99021bb681dfc3302a9d99a2e53e64e",
418]
419
420
421def selftest(out=sys.stdout):
422 acc = IDENT
423 for i, want in enumerate(RFC9496_MULTIPLES):
424 got = encode(acc).hex()
425 if got != want:
426 print(f"FAIL multiple {i}: {got} != {want}", file=out)
427 return 1
428 rt = decode(bytes.fromhex(want))
429 if rt is None or not eq(rt, acc):
430 print(f"FAIL decode {i}", file=out)
431 return 1
432 acc = add(acc, G)
433 print("selftest ok", file=out)
434 return 0
435
436# ── Block I/O ───────────────────────────────────────────────────────────────
437
438
439def parse_blocks(text):
440 cases, cur = [], None
441 for line in text.splitlines():
442 line = line.strip()
443 if not line or line.startswith("#"):
444 continue
445 key, _, val = line.partition(" ")
446 if key == "case":
447 cur = {"case": val}
448 elif key == "end":
449 cases.append(cur)
450 cur = None
451 else:
452 cur[key] = val
453 return cases
454
455
456def hx(v):
457 return bytes.fromhex(v) if v else b""
458
459
460def do_sign(c):
461 seeds = [hx(s) for s in c["seeds"].split(",")]
462 secrets = [secret_from_seed(s) for s in seeds]
463 keys = [encode(mul(s, G)) for s in secrets]
464 l = int(c["signer"])
465 tag, body = sign(keys, secrets[l], l, hx(c.get("scope", "")), hx(c.get("msg", "")), hx(c.get("aux", "")))
466 return keys, tag, body
467
468
469def main():
470 mode = sys.argv[1] if len(sys.argv) > 1 else "selftest"
471 if mode == "selftest":
472 return selftest()
473 if selftest(sys.stderr) != 0:
474 return 1
475 if mode == "sign":
476 for c in parse_blocks(sys.stdin.read()):
477 keys, tag, body = do_sign(c)
478 print(f"case {c['case']}")
479 print(f"ring {b''.join(keys).hex()}")
480 print(f"digest {ring_digest(keys).hex()}")
481 print(f"tag {tag.hex()}")
482 print(f"body {body.hex()}")
483 print("end")
484 return 0
485 if mode == "verify":
486 for c in parse_blocks(sys.stdin.read()):
487 ring = hx(c["ring"])
488 keys = [ring[i:i + 32] for i in range(0, len(ring), 32)]
489 ok = verify(keys, hx(c.get("scope", "")), hx(c.get("msg", "")), hx(c["tag"]), hx(c["body"]))
490 print(f"case {c['case']}")
491 print(f"ok {'true' if ok else 'false'}")
492 print("end")
493 return 0
494 print(f"unknown mode {mode}")
495 return 2
496
497
498if __name__ == "__main__":
499 sys.exit(main())