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 }