lib/stun/src/message.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const attribute = @import("attribute.zig");
  3 const integrity = @import("integrity.zig");
  4 
  5 pub const magic_cookie: u32 = 0x2112a442;
  6 pub const header_bytes: usize = 20;
  7 pub const max_payload_bytes: u16 = 65_532;
  8 pub const max_message_bytes: u17 = header_bytes + max_payload_bytes;
  9 pub const max_attribute_count: u14 = max_payload_bytes / 4;
 10 
 11 pub const TransactionId = [12]u8;
 12 
 13 pub const Class = enum(u2) {
 14     request = 0,
 15     indication = 1,
 16     success_response = 2,
 17     error_response = 3,
 18 };
 19 
 20 pub const Method = enum(u12) {
 21     reserved = 0x000,
 22     binding = 0x001,
 23     _,
 24 };
 25 
 26 pub const Header = struct {
 27     class: Class,
 28     method: Method,
 29     length: u16,
 30     transaction_id: TransactionId,
 31 
 32     pub fn messageType(self: Header) u16 {
 33         return encodeType(self.method, self.class);
 34     }
 35 };
 36 
 37 const MessageDecodeError = error{
 38     FingerprintNotLast,
 39     InvalidMagicCookie,
 40     InvalidMessageLength,
 41     InvalidMessageType,
 42     MessageTooShort,
 43     TooManyAttributes,
 44     TrailingBytes,
 45     TruncatedAttribute,
 46 };
 47 
 48 pub const DecodeError = attribute.DecodeError || MessageDecodeError;
 49 
 50 const AuthenticationError = error{
 51     FingerprintMismatch,
 52     IntegrityMismatch,
 53     MissingFingerprint,
 54     MissingIntegrity,
 55 };
 56 
 57 pub const VerifyError = DecodeError || AuthenticationError;
 58 
 59 const MessageEncodeError = error{
 60     AttributeLimitExceeded,
 61     CapacityOverflow,
 62     InvalidAttributePadding,
 63     InvalidAttributeType,
 64     InvalidBuilderState,
 65     InvalidLimits,
 66     MessageTooLarge,
 67     StorageTooSmall,
 68     UseFinishFunction,
 69 };
 70 
 71 pub const EncodeError = integrity.CredentialError || MessageEncodeError;
 72 
 73 pub const Limits = struct {
 74     message_bytes: u17,
 75     attributes: u14,
 76 };
 77 
 78 pub const Capacity = struct {
 79     message_bytes: u17,
 80     payload_bytes: u16,
 81     attributes: u14,
 82 
 83     pub fn derive(limits: Limits) EncodeError!Capacity {
 84         if (limits.message_bytes < header_bytes) return error.InvalidLimits;
 85         if (limits.message_bytes > max_message_bytes) return error.InvalidLimits;
 86         if ((limits.message_bytes & 3) != 0) return error.InvalidLimits;
 87         const payload = std.math.sub(u17, limits.message_bytes, header_bytes) catch
 88             return error.CapacityOverflow;
 89         if (payload > max_payload_bytes) return error.InvalidLimits;
 90         return .{
 91             .message_bytes = limits.message_bytes,
 92             .payload_bytes = @intCast(payload),
 93             .attributes = limits.attributes,
 94         };
 95     }
 96 };
 97 
 98 pub const Message = struct {
 99     header: Header,
100     bytes: []const u8,
101     attribute_count: u14,
102 
103     pub fn decode(bytes: []const u8) DecodeError!Message {
104         if (bytes.len < header_bytes) return error.MessageTooShort;
105         const message_type = std.mem.readInt(u16, bytes[0..2], .big);
106         if ((message_type & 0xc000) != 0) return error.InvalidMessageType;
107         const payload_len = std.mem.readInt(u16, bytes[2..4], .big);
108         if ((payload_len & 3) != 0) return error.InvalidMessageLength;
109         const total = std.math.add(usize, header_bytes, payload_len) catch
110             return error.InvalidMessageLength;
111         if (total > bytes.len) return error.InvalidMessageLength;
112         if (total < bytes.len) return error.TrailingBytes;
113         if (std.mem.readInt(u32, bytes[4..8], .big) != magic_cookie) {
114             return error.InvalidMagicCookie;
115         }
116         const header = Header{
117             .class = decodeClass(message_type),
118             .method = decodeMethod(message_type),
119             .length = payload_len,
120             .transaction_id = bytes[8..20].*,
121         };
122         const count = try validateAttributes(bytes, header.transaction_id);
123         return .{ .header = header, .bytes = bytes, .attribute_count = count };
124     }
125 
126     pub fn attributes(self: Message) Iterator {
127         return .{
128             .bytes = self.bytes,
129             .transaction_id = self.header.transaction_id,
130         };
131     }
132 
133     pub fn verifyIntegrity(self: Message, key: []const u8) VerifyError!void {
134         const item = try self.integrityItem() orelse return error.MissingIntegrity;
135         const adjusted_len: u16 = @intCast(item.end - header_bytes);
136         switch (item.attribute.value) {
137             .message_integrity => |expected| {
138                 var actual: [integrity.sha1_bytes]u8 = undefined;
139                 integrity.hmacSha1Adjusted(
140                     &actual,
141                     self.bytes,
142                     item.start,
143                     adjusted_len,
144                     key,
145                 );
146                 if (!integrity.equalSha1(actual, expected)) return error.IntegrityMismatch;
147             },
148             .message_integrity_sha256 => |expected| {
149                 var actual: [integrity.sha256_bytes]u8 = undefined;
150                 integrity.hmacSha256Adjusted(
151                     &actual,
152                     self.bytes,
153                     item.start,
154                     adjusted_len,
155                     key,
156                 );
157                 if (!integrity.equalSha256(actual, expected)) return error.IntegrityMismatch;
158             },
159             else => unreachable,
160         }
161     }
162 
163     pub fn verifyFingerprint(self: Message) VerifyError!void {
164         var iterator = self.attributes();
165         for (0..max_attribute_count) |_| {
166             const item = try iterator.nextItem() orelse break;
167             if (item.attribute.value == .fingerprint) {
168                 const expected = item.attribute.value.fingerprint;
169                 if (integrity.fingerprint(self.bytes[0..item.start]) != expected) {
170                     return error.FingerprintMismatch;
171                 }
172                 return;
173             }
174         }
175         return error.MissingFingerprint;
176     }
177 
178     fn integrityItem(self: Message) DecodeError!?Item {
179         var iterator = self.attributes();
180         var sha1: ?Item = null;
181         for (0..max_attribute_count) |_| {
182             const item = try iterator.nextItem() orelse break;
183             switch (item.attribute.value) {
184                 .message_integrity_sha256 => return item,
185                 .message_integrity => if (sha1 == null) {
186                     sha1 = item;
187                 },
188                 else => {},
189             }
190         }
191         return sha1;
192     }
193 };
194 
195 pub const Item = struct {
196     attribute: attribute.Attribute,
197     start: usize,
198     end: usize,
199 };
200 
201 pub const Iterator = struct {
202     bytes: []const u8,
203     transaction_id: TransactionId,
204     offset: usize = header_bytes,
205 
206     pub fn next(self: *Iterator) DecodeError!?attribute.Attribute {
207         const item = try self.nextItem() orelse return null;
208         return item.attribute;
209     }
210 
211     pub fn nextItem(self: *Iterator) DecodeError!?Item {
212         if (self.offset == self.bytes.len) return null;
213         if (self.bytes.len - self.offset < 4) return error.TruncatedAttribute;
214         const start = self.offset;
215         const header = self.bytes[start..][0..4];
216         const type_code = std.mem.readInt(u16, header[0..2], .big);
217         const value_len = std.mem.readInt(u16, header[2..4], .big);
218         const value_start = std.math.add(usize, start, 4) catch
219             return error.TruncatedAttribute;
220         const value_end = std.math.add(usize, value_start, value_len) catch
221             return error.TruncatedAttribute;
222         if (value_end > self.bytes.len) return error.TruncatedAttribute;
223         const end = std.math.add(usize, value_end, attribute.paddingLength(value_len)) catch
224             return error.TruncatedAttribute;
225         if (end > self.bytes.len) return error.TruncatedAttribute;
226         const decoded = try attribute.decode(
227             type_code,
228             self.bytes[value_start..value_end],
229             self.bytes[value_end..end],
230             magic_cookie,
231             self.transaction_id,
232         );
233         self.offset = end;
234         return .{ .attribute = decoded, .start = start, .end = end };
235     }
236 };
237 
238 fn validateAttributes(bytes: []const u8, transaction_id: TransactionId) DecodeError!u14 {
239     var iterator = Iterator{ .bytes = bytes, .transaction_id = transaction_id };
240     var count: usize = 0;
241     var fingerprint_seen = false;
242     for (0..max_attribute_count) |_| {
243         const item = try iterator.nextItem() orelse return @intCast(count);
244         if (fingerprint_seen) return error.FingerprintNotLast;
245         if (item.attribute.value == .fingerprint) fingerprint_seen = true;
246         count += 1;
247     }
248     if (iterator.offset != bytes.len) return error.TooManyAttributes;
249     return @intCast(count);
250 }
251 
252 const BuilderState = enum {
253     attributes,
254     sha1,
255     sha256,
256     fingerprint,
257 };
258 
259 const BuilderRegion = struct {
260     start: usize,
261     value_start: usize,
262     value_end: usize,
263     end: usize,
264 };
265 
266 pub const Builder = struct {
267     storage: []u8,
268     capacity: Capacity,
269     transaction_id: TransactionId,
270     index: usize = header_bytes,
271     attribute_count: u14 = 0,
272     state: BuilderState = .attributes,
273 
274     pub fn init(
275         storage: []u8,
276         limits: Limits,
277         class: Class,
278         method: Method,
279         transaction_id: TransactionId,
280     ) EncodeError!Builder {
281         const capacity = try Capacity.derive(limits);
282         if (storage.len < capacity.message_bytes) return error.StorageTooSmall;
283         var builder = Builder{
284             .storage = storage[0..capacity.message_bytes],
285             .capacity = capacity,
286             .transaction_id = transaction_id,
287         };
288         std.mem.writeInt(u16, builder.storage[0..2], encodeType(method, class), .big);
289         std.mem.writeInt(u16, builder.storage[2..4], 0, .big);
290         std.mem.writeInt(u32, builder.storage[4..8], magic_cookie, .big);
291         @memcpy(builder.storage[8..20], &transaction_id);
292         return builder;
293     }
294 
295     pub fn appendRaw(
296         self: *Builder,
297         type_code: u16,
298         value: []const u8,
299     ) EncodeError!void {
300         const zeros = [_]u8{ 0, 0, 0 };
301         const padding_len = attribute.paddingLength(value.len);
302         try self.appendRawPadded(type_code, value, zeros[0..padding_len]);
303     }
304 
305     pub fn appendRawPadded(
306         self: *Builder,
307         type_code: u16,
308         value: []const u8,
309         padding: []const u8,
310     ) EncodeError!void {
311         if (self.state != .attributes) return error.InvalidBuilderState;
312         if (isIntegrityOrFingerprint(type_code)) return error.UseFinishFunction;
313         if (padding.len != attribute.paddingLength(value.len)) {
314             return error.InvalidAttributePadding;
315         }
316         const region = try self.reserve(value.len, padding.len);
317         std.mem.writeInt(u16, self.storage[region.start..][0..2], type_code, .big);
318         std.mem.writeInt(u16, self.storage[region.start + 2 ..][0..2], @intCast(value.len), .big);
319         @memcpy(self.storage[region.value_start..region.value_end], value);
320         @memcpy(self.storage[region.value_end..region.end], padding);
321         self.commit(region.end);
322     }
323 
324     pub fn appendDecoded(self: *Builder, value: attribute.Attribute) EncodeError!void {
325         try self.appendRawPadded(value.type_code, value.raw_value, value.padding);
326     }
327 
328     pub fn appendUsername(self: *Builder, username: []const u8) EncodeError!void {
329         _ = try integrity.shortTermKey(username);
330         try self.appendRaw(@backingInt(attribute.Type.username), username);
331     }
332 
333     pub fn appendAddress(
334         self: *Builder,
335         attribute_type: attribute.Type,
336         address: attribute.Address,
337     ) EncodeError!void {
338         const xor_cookie: ?u32 = switch (attribute_type) {
339             .xor_mapped_address => magic_cookie,
340             .mapped_address, .alternate_server, .response_origin, .other_address => null,
341             else => return error.InvalidAttributeType,
342         };
343         var encoded: [20]u8 = undefined;
344         const length = attribute.encodeAddress(
345             &encoded,
346             address,
347             xor_cookie,
348             self.transaction_id,
349         );
350         try self.appendRaw(@backingInt(attribute_type), encoded[0..length]);
351     }
352 
353     pub fn appendChangeRequest(
354         self: *Builder,
355         request: attribute.ChangeRequest,
356     ) EncodeError!void {
357         var encoded: [4]u8 = @splat(0);
358         if (request.change_ip) encoded[3] |= 0x04;
359         if (request.change_port) encoded[3] |= 0x02;
360         try self.appendRaw(@backingInt(attribute.Type.change_request), &encoded);
361     }
362 
363     pub fn appendResponsePort(self: *Builder, port: u16) EncodeError!void {
364         var encoded: [2]u8 = undefined;
365         std.mem.writeInt(u16, &encoded, port, .big);
366         try self.appendRaw(@backingInt(attribute.Type.response_port), &encoded);
367     }
368 
369     pub fn finishIntegrity(self: *Builder, key: []const u8) EncodeError!void {
370         if (self.state != .attributes) return error.InvalidBuilderState;
371         const region = try self.reserve(integrity.sha1_bytes, 0);
372         writeAttributeHeader(
373             self.storage,
374             region.start,
375             @backingInt(attribute.Type.message_integrity),
376             integrity.sha1_bytes,
377         );
378         @memset(self.storage[region.value_start..region.value_end], 0);
379         const adjusted_len: u16 = @intCast(region.end - header_bytes);
380         var digest: [integrity.sha1_bytes]u8 = undefined;
381         integrity.hmacSha1Adjusted(&digest, self.storage, region.start, adjusted_len, key);
382         @memcpy(self.storage[region.value_start..region.value_end], &digest);
383         self.state = .sha1;
384         self.commit(region.end);
385     }
386 
387     pub fn finishIntegritySha256(self: *Builder, key: []const u8) EncodeError!void {
388         if (self.state != .attributes and self.state != .sha1) {
389             return error.InvalidBuilderState;
390         }
391         const region = try self.reserve(integrity.sha256_bytes, 0);
392         writeAttributeHeader(
393             self.storage,
394             region.start,
395             @backingInt(attribute.Type.message_integrity_sha256),
396             integrity.sha256_bytes,
397         );
398         @memset(self.storage[region.value_start..region.value_end], 0);
399         const adjusted_len: u16 = @intCast(region.end - header_bytes);
400         var digest: [integrity.sha256_bytes]u8 = undefined;
401         integrity.hmacSha256Adjusted(&digest, self.storage, region.start, adjusted_len, key);
402         @memcpy(self.storage[region.value_start..region.value_end], &digest);
403         self.state = .sha256;
404         self.commit(region.end);
405     }
406 
407     pub fn finishFingerprint(self: *Builder) EncodeError!void {
408         if (self.state == .fingerprint) return error.InvalidBuilderState;
409         const region = try self.reserve(4, 0);
410         writeAttributeHeader(
411             self.storage,
412             region.start,
413             @backingInt(attribute.Type.fingerprint),
414             4,
415         );
416         @memset(self.storage[region.value_start..region.value_end], 0);
417         self.patchLength(region.end);
418         const value = integrity.fingerprint(self.storage[0..region.start]);
419         std.mem.writeInt(u32, self.storage[region.value_start..][0..4], value, .big);
420         self.state = .fingerprint;
421         self.commit(region.end);
422     }
423 
424     pub fn finish(self: *Builder) []const u8 {
425         self.patchLength(self.index);
426         return self.storage[0..self.index];
427     }
428 
429     fn reserve(self: *Builder, value_len: usize, padding_len: usize) EncodeError!BuilderRegion {
430         if (self.attribute_count == self.capacity.attributes) {
431             return error.AttributeLimitExceeded;
432         }
433         if (value_len > std.math.maxInt(u16)) return error.MessageTooLarge;
434         const value_start = std.math.add(usize, self.index, 4) catch
435             return error.CapacityOverflow;
436         const value_end = std.math.add(usize, value_start, value_len) catch
437             return error.CapacityOverflow;
438         const end = std.math.add(usize, value_end, padding_len) catch
439             return error.CapacityOverflow;
440         if (end > self.capacity.message_bytes) return error.MessageTooLarge;
441         if (end - header_bytes > max_payload_bytes) return error.MessageTooLarge;
442         return .{
443             .start = self.index,
444             .value_start = value_start,
445             .value_end = value_end,
446             .end = end,
447         };
448     }
449 
450     fn commit(self: *Builder, end: usize) void {
451         std.debug.assert(end > self.index);
452         std.debug.assert(end <= self.capacity.message_bytes);
453         self.index = end;
454         self.attribute_count += 1;
455         self.patchLength(end);
456     }
457 
458     fn patchLength(self: *Builder, end: usize) void {
459         std.debug.assert(end >= header_bytes);
460         std.debug.assert(end - header_bytes <= max_payload_bytes);
461         std.mem.writeInt(u16, self.storage[2..4], @intCast(end - header_bytes), .big);
462     }
463 };
464 
465 fn writeAttributeHeader(
466     storage: []u8,
467     start: usize,
468     type_code: u16,
469     value_len: u16,
470 ) void {
471     std.mem.writeInt(u16, storage[start..][0..2], type_code, .big);
472     std.mem.writeInt(u16, storage[start + 2 ..][0..2], value_len, .big);
473 }
474 
475 fn isIntegrityOrFingerprint(type_code: u16) bool {
476     return type_code == @backingInt(attribute.Type.message_integrity) or
477         type_code == @backingInt(attribute.Type.message_integrity_sha256) or
478         type_code == @backingInt(attribute.Type.fingerprint);
479 }
480 
481 pub fn encodeType(method: Method, class: Class) u16 {
482     const method_bits: u16 = @backingInt(method);
483     const class_bits: u16 = @backingInt(class);
484     return ((method_bits & 0x0f80) << 2) |
485         ((method_bits & 0x0070) << 1) |
486         (method_bits & 0x000f) |
487         ((class_bits & 0x0002) << 7) |
488         ((class_bits & 0x0001) << 4);
489 }
490 
491 pub fn decodeMethod(message_type: u16) Method {
492     const value = ((message_type & 0x3e00) >> 2) |
493         ((message_type & 0x00e0) >> 1) |
494         (message_type & 0x000f);
495     return @fromBackingInt(@intCast(@as(u12, @truncate(value))));
496 }
497 
498 pub fn decodeClass(message_type: u16) Class {
499     const value = ((message_type & 0x0100) >> 7) |
500         ((message_type & 0x0010) >> 4);
501     return @fromBackingInt(@intCast(@as(u2, @truncate(value))));
502 }