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 }