lib/reticulum/src/crypto/hkdf.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const crypto = @import("root.zig");
3
4 pub const hash_length: u8 = crypto.hmac.tag_length;
5 pub const max_output_length: u16 = std.math.maxInt(u16);
6
7 pub const HkdfError = error{
8 InvalidLength,
9 OutputTooLong,
10 EmptyKeyMaterial,
11 OutputTooSmall,
12 OverlappingBuffers,
13 };
14
15 /// Writes `length` derived bytes into `out` and returns them, so a caller turns
16 /// a shared secret into exactly as many key bytes as it needs, following
17 /// Reticulum@1.5.0 RNS/Cryptography/HKDF.py:35-62. The first stage signs the
18 /// key material under the salt, and the second stage builds the output one
19 /// 32-byte block at a time, each block signed over the block before it, then
20 /// the context, then a counter byte that starts at one and wraps at 256. A salt
21 /// the caller leaves out, or passes empty, becomes 32 zero bytes. The call
22 /// returns `error.InvalidLength` for a length of zero, `error.OutputTooLong`
23 /// past 65,535 bytes, `error.EmptyKeyMaterial` for empty key material,
24 /// `error.OutputTooSmall` when `out` is shorter than the length, and
25 /// `error.OverlappingBuffers` when the key material, the salt, or the context
26 /// overlaps the output.
27 pub fn derive(
28 length: u17,
29 key_material: []const u8,
30 salt: ?[]const u8,
31 context: ?[]const u8,
32 out: []u8,
33 ) HkdfError![]u8 {
34 if (length == 0) return error.InvalidLength;
35 if (length > max_output_length) return error.OutputTooLong;
36 if (key_material.len == 0) return error.EmptyKeyMaterial;
37 const output_length: usize = @intCast(length);
38 if (out.len < output_length) return error.OutputTooSmall;
39 const result = out[0..output_length];
40 if (overlaps(key_material, result)) return error.OverlappingBuffers;
41 if (salt) |value| if (overlaps(value, result)) return error.OverlappingBuffers;
42 if (context) |value| if (overlaps(value, result)) return error.OverlappingBuffers;
43
44 const zero_salt: [hash_length]u8 = @splat(0);
45 const selected_salt = effectiveSalt(salt, &zero_salt);
46 const selected_context = context orelse "";
47 const pseudorandom_key = crypto.hmac.sign(selected_salt, key_material);
48 expand(@intCast(length), pseudorandom_key, selected_context, result);
49 return result;
50 }
51
52 fn effectiveSalt(salt: ?[]const u8, zero_salt: *const [hash_length]u8) []const u8 {
53 const value = salt orelse return zero_salt;
54 if (value.len == 0) return zero_salt;
55 return value;
56 }
57
58 fn overlaps(input: []const u8, output: []u8) bool {
59 if (input.len == 0) return false;
60 const input_start = @intFromPtr(input.ptr);
61 const output_start = @intFromPtr(output.ptr);
62 const input_end = input_start + input.len;
63 const output_end = output_start + output.len;
64 return input_start < output_end and output_start < input_end;
65 }
66
67 fn expand(
68 length: u16,
69 pseudorandom_key: [hash_length]u8,
70 context: []const u8,
71 out: []u8,
72 ) void {
73 const output_length: usize = length;
74 std.debug.assert(output_length == out.len);
75 std.debug.assert(length <= max_output_length);
76 var previous: [hash_length]u8 = undefined;
77 var previous_length: usize = 0;
78 var written: usize = 0;
79 var block_index: u16 = 0;
80 while (written < output_length) : (block_index += 1) {
81 const counter = [1]u8{@intCast((block_index + 1) % 256)};
82 var signer = crypto.hmac.Signer.init(&pseudorandom_key);
83 signer.update(previous[0..previous_length]);
84 signer.update(context);
85 signer.update(&counter);
86 previous = signer.final();
87 previous_length = hash_length;
88 const take = @min(hash_length, output_length - written);
89 @memcpy(out[written..][0..take], previous[0..take]);
90 written += take;
91 }
92 std.debug.assert(written == output_length);
93 }