a lot
This commit is contained in:
@@ -0,0 +1,76 @@
|
||||
#include "UMNSystem.h"
|
||||
#include <errno.h>
|
||||
#include <string.h>
|
||||
#include <sys/socket.h>
|
||||
#include <sys/ioctl.h>
|
||||
#include <unistd.h>
|
||||
#ifdef __APPLE__
|
||||
#include <net/if.h>
|
||||
#include <net/if_utun.h>
|
||||
#include <sys/kern_control.h>
|
||||
#include <sys/sys_domain.h>
|
||||
#endif
|
||||
|
||||
int umn_utun_open(char *name, size_t name_capacity) {
|
||||
#ifdef __APPLE__
|
||||
int fd = socket(PF_SYSTEM, SOCK_DGRAM, SYSPROTO_CONTROL);
|
||||
if (fd < 0) return -1;
|
||||
struct ctl_info info;
|
||||
memset(&info, 0, sizeof(info));
|
||||
strlcpy(info.ctl_name, UTUN_CONTROL_NAME, sizeof(info.ctl_name));
|
||||
if (ioctl(fd, CTLIOCGINFO, &info) < 0) { close(fd); return -1; }
|
||||
struct sockaddr_ctl address;
|
||||
memset(&address, 0, sizeof(address));
|
||||
address.sc_len = sizeof(address); address.sc_family = AF_SYSTEM;
|
||||
address.ss_sysaddr = AF_SYS_CONTROL; address.sc_id = info.ctl_id; address.sc_unit = 0;
|
||||
if (connect(fd, (struct sockaddr *)&address, sizeof(address)) < 0) { close(fd); return -1; }
|
||||
socklen_t length = (socklen_t)name_capacity;
|
||||
if (getsockopt(fd, SYSPROTO_CONTROL, UTUN_OPT_IFNAME, name, &length) < 0) { close(fd); return -1; }
|
||||
return fd;
|
||||
#else
|
||||
(void)name; (void)name_capacity; errno = ENOTSUP; return -1;
|
||||
#endif
|
||||
}
|
||||
|
||||
int umn_get_peer_eid(int socket_fd, uid_t *uid, gid_t *gid) {
|
||||
#ifdef __APPLE__
|
||||
return getpeereid(socket_fd, uid, gid);
|
||||
#else
|
||||
struct ucred credential; socklen_t size = sizeof(credential);
|
||||
if (getsockopt(socket_fd, SOL_SOCKET, SO_PEERCRED, &credential, &size) < 0) return -1;
|
||||
*uid = credential.uid; *gid = credential.gid; return 0;
|
||||
#endif
|
||||
}
|
||||
|
||||
int umn_send_fd(int socket_fd, int passed_fd, const void *data, size_t length) {
|
||||
struct iovec io = { .iov_base = (void *)data, .iov_len = length };
|
||||
char control[CMSG_SPACE(sizeof(int))]; memset(control, 0, sizeof(control));
|
||||
struct msghdr message; memset(&message, 0, sizeof(message));
|
||||
message.msg_iov = &io; message.msg_iovlen = 1;
|
||||
message.msg_control = control; message.msg_controllen = sizeof(control);
|
||||
struct cmsghdr *header = CMSG_FIRSTHDR(&message);
|
||||
header->cmsg_level = SOL_SOCKET; header->cmsg_type = SCM_RIGHTS; header->cmsg_len = CMSG_LEN(sizeof(int));
|
||||
memcpy(CMSG_DATA(header), &passed_fd, sizeof(int));
|
||||
ssize_t result = sendmsg(socket_fd, &message, 0);
|
||||
if (result < 0) return -1;
|
||||
if ((size_t)result != length) { errno = EIO; return -1; }
|
||||
return 0;
|
||||
}
|
||||
|
||||
ssize_t umn_recv_fd(int socket_fd, int *passed_fd, void *data, size_t capacity) {
|
||||
struct iovec io = { .iov_base = data, .iov_len = capacity };
|
||||
char control[CMSG_SPACE(sizeof(int))]; memset(control, 0, sizeof(control));
|
||||
struct msghdr message; memset(&message, 0, sizeof(message));
|
||||
message.msg_iov = &io; message.msg_iovlen = 1;
|
||||
message.msg_control = control; message.msg_controllen = sizeof(control);
|
||||
ssize_t result = recvmsg(socket_fd, &message, 0);
|
||||
if (result <= 0) return result;
|
||||
*passed_fd = -1;
|
||||
for (struct cmsghdr *header = CMSG_FIRSTHDR(&message); header; header = CMSG_NXTHDR(&message, header)) {
|
||||
if (header->cmsg_level == SOL_SOCKET && header->cmsg_type == SCM_RIGHTS && header->cmsg_len >= CMSG_LEN(sizeof(int))) {
|
||||
memcpy(passed_fd, CMSG_DATA(header), sizeof(int)); break;
|
||||
}
|
||||
}
|
||||
if (*passed_fd < 0) { errno = EBADMSG; return -1; }
|
||||
return result;
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
#ifndef UMN_SYSTEM_H
|
||||
#define UMN_SYSTEM_H
|
||||
#include <stddef.h>
|
||||
#include <sys/types.h>
|
||||
|
||||
int umn_utun_open(char *name, size_t name_capacity);
|
||||
int umn_get_peer_eid(int socket_fd, uid_t *uid, gid_t *gid);
|
||||
int umn_send_fd(int socket_fd, int passed_fd, const void *data, size_t length);
|
||||
ssize_t umn_recv_fd(int socket_fd, int *passed_fd, void *data, size_t capacity);
|
||||
|
||||
#endif
|
||||
@@ -1,9 +1,24 @@
|
||||
import Foundation
|
||||
|
||||
public enum FirewallProtocol: String, Codable, CaseIterable, Sendable {
|
||||
case overlay, tcp, udp
|
||||
}
|
||||
|
||||
public struct FirewallRule: Codable, Hashable, Sendable {
|
||||
public let protocolKind: FirewallProtocol
|
||||
public let port: UInt16
|
||||
public let source: MeshAddress?
|
||||
public init(port: UInt16, source: MeshAddress?) { self.port = port; self.source = source }
|
||||
public init(protocolKind: FirewallProtocol = .overlay, port: UInt16, source: MeshAddress?) {
|
||||
self.protocolKind = protocolKind; self.port = port; self.source = source
|
||||
}
|
||||
|
||||
private enum CodingKeys: String, CodingKey { case protocolKind = "protocol", port, source }
|
||||
public init(from decoder: Decoder) throws {
|
||||
let c = try decoder.container(keyedBy: CodingKeys.self)
|
||||
protocolKind = try c.decodeIfPresent(FirewallProtocol.self, forKey: .protocolKind) ?? .overlay
|
||||
port = try c.decode(UInt16.self, forKey: .port)
|
||||
source = try c.decodeIfPresent(MeshAddress.self, forKey: .source)
|
||||
}
|
||||
}
|
||||
|
||||
public final class MeshFirewall: @unchecked Sendable {
|
||||
@@ -11,11 +26,20 @@ public final class MeshFirewall: @unchecked Sendable {
|
||||
private let lock = NSLock()
|
||||
public init(rules: Set<FirewallRule> = []) { self.rules = rules }
|
||||
|
||||
public func allows(port: UInt16, source: MeshAddress) -> Bool {
|
||||
public func allows(protocol protocolKind: FirewallProtocol = .overlay, port: UInt16, source: MeshAddress) -> Bool {
|
||||
lock.lock(); defer { lock.unlock() }
|
||||
return rules.contains(FirewallRule(port: port, source: nil)) || rules.contains(FirewallRule(port: port, source: source))
|
||||
return rules.contains(FirewallRule(protocolKind: protocolKind, port: port, source: nil)) ||
|
||||
rules.contains(FirewallRule(protocolKind: protocolKind, port: port, source: source))
|
||||
}
|
||||
public func allow(protocol protocolKind: FirewallProtocol = .overlay, port: UInt16, source: MeshAddress?) {
|
||||
lock.lock(); rules.insert(.init(protocolKind: protocolKind, port: port, source: source)); lock.unlock()
|
||||
}
|
||||
public func revoke(protocol protocolKind: FirewallProtocol = .overlay, port: UInt16, source: MeshAddress?) {
|
||||
lock.lock(); rules.remove(.init(protocolKind: protocolKind, port: port, source: source)); lock.unlock()
|
||||
}
|
||||
public func allRules() -> [FirewallRule] {
|
||||
lock.lock(); defer { lock.unlock() }
|
||||
return rules.sorted { ($0.protocolKind.rawValue, $0.port, $0.source?.description ?? "") <
|
||||
($1.protocolKind.rawValue, $1.port, $1.source?.description ?? "") }
|
||||
}
|
||||
public func allow(port: UInt16, source: MeshAddress?) { lock.lock(); rules.insert(.init(port: port, source: source)); lock.unlock() }
|
||||
public func revoke(port: UInt16, source: MeshAddress?) { lock.lock(); rules.remove(.init(port: port, source: source)); lock.unlock() }
|
||||
public func allRules() -> [FirewallRule] { lock.lock(); defer { lock.unlock() }; return rules.sorted { $0.port < $1.port } }
|
||||
}
|
||||
|
||||
@@ -25,7 +25,7 @@ public enum FrameCodec {
|
||||
}
|
||||
|
||||
public enum IPCOperation: String, Codable, Sendable {
|
||||
case status, address, peers, routes, ping
|
||||
case status, address, peers, routes, ping, interfaceStatus
|
||||
case aliasSet, aliasRemove, aliasList
|
||||
case firewallAllow, firewallRevoke, firewallList
|
||||
case textSend, bindText, bindStream, openStream, streamData, streamClose
|
||||
@@ -39,6 +39,13 @@ public struct IPCMessage: Codable, Sendable {
|
||||
public var source: String?
|
||||
public var name: String?
|
||||
public var port: UInt16?
|
||||
public var firewallProtocol: FirewallProtocol?
|
||||
public var interfaceName: String?
|
||||
public var interfaceAddress: String?
|
||||
public var interfaceMTU: UInt16?
|
||||
public var helperState: String?
|
||||
public var synchronizedRouteCount: Int?
|
||||
public var dnsState: String?
|
||||
public var streamID: UInt64?
|
||||
public var data: Data?
|
||||
public var ok: Bool?
|
||||
@@ -46,6 +53,6 @@ public struct IPCMessage: Codable, Sendable {
|
||||
public var values: [String]?
|
||||
|
||||
public init(operation: IPCOperation, requestID: UUID = UUID()) {
|
||||
self.version = 1; self.operation = operation; self.requestID = requestID
|
||||
self.version = 2; self.operation = operation; self.requestID = requestID
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
import Foundation
|
||||
|
||||
public struct RouteSetDiff: Equatable, Sendable {
|
||||
public let additions: Set<MeshAddress>
|
||||
public let removals: Set<MeshAddress>
|
||||
}
|
||||
|
||||
public enum HelperRouteSet {
|
||||
public static func validate(_ desired: Set<MeshAddress>, local: MeshAddress, maximum: Int = 31) throws {
|
||||
guard desired.count <= maximum, !desired.contains(local) else { throw UMNError.invalidFrame }
|
||||
}
|
||||
public static func diff(current: Set<MeshAddress>, desired: Set<MeshAddress>) -> RouteSetDiff {
|
||||
.init(additions: desired.subtracting(current), removals: current.subtracting(desired))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
import Foundation
|
||||
|
||||
public enum IPv6Transport: UInt8, Sendable {
|
||||
case tcp = 6
|
||||
case udp = 17
|
||||
case icmpv6 = 58
|
||||
}
|
||||
|
||||
public struct IPv6PacketInfo: Sendable {
|
||||
public let source: MeshAddress
|
||||
public let destination: MeshAddress
|
||||
public let transport: IPv6Transport
|
||||
public let sourcePort: UInt16?
|
||||
public let destinationPort: UInt16?
|
||||
public let tcpFlags: UInt8?
|
||||
public let icmpType: UInt8?
|
||||
public let transportOffset: Int
|
||||
}
|
||||
|
||||
public enum IPv6PacketParser {
|
||||
public static let maximumPacketSize = 1280
|
||||
|
||||
public static func parse(_ packet: Data) throws -> IPv6PacketInfo {
|
||||
guard packet.count >= 40, packet.count <= maximumPacketSize, packet[0] >> 4 == 6 else { throw UMNError.invalidFrame }
|
||||
let payloadLength = Int(read16(packet, 4))
|
||||
guard payloadLength == packet.count - 40 else { throw UMNError.invalidFrame }
|
||||
let source = try MeshAddress(bytes: packet.subdata(in: 8..<24))
|
||||
let destinationBytes = packet.subdata(in: 24..<40)
|
||||
guard destinationBytes.first != 0xff else { throw UMNError.invalidFrame }
|
||||
let destination = try MeshAddress(bytes: destinationBytes)
|
||||
|
||||
var next = packet[6]
|
||||
var offset = 40
|
||||
var extensions = 0
|
||||
while [UInt8(0), 43, 60].contains(next) {
|
||||
guard extensions < 8, offset + 2 <= packet.count else { throw UMNError.invalidFrame }
|
||||
let following = packet[offset]
|
||||
let length: Int
|
||||
length = (Int(packet[offset + 1]) + 1) * 8
|
||||
guard length >= 8, offset + length <= packet.count else { throw UMNError.invalidFrame }
|
||||
if next == 43, packet[offset + 3] != 0 { throw UMNError.invalidFrame }
|
||||
if next == 0 || next == 60 { try validateOptions(packet, range: (offset + 2)..<(offset + length)) }
|
||||
next = following; offset += length; extensions += 1
|
||||
}
|
||||
guard next != 44, let transport = IPv6Transport(rawValue: next) else { throw UMNError.invalidFrame }
|
||||
|
||||
switch transport {
|
||||
case .tcp:
|
||||
guard offset + 20 <= packet.count else { throw UMNError.invalidFrame }
|
||||
let headerLength = Int(packet[offset + 12] >> 4) * 4
|
||||
guard headerLength >= 20, offset + headerLength <= packet.count else { throw UMNError.invalidFrame }
|
||||
return .init(source: source, destination: destination, transport: transport,
|
||||
sourcePort: read16(packet, offset), destinationPort: read16(packet, offset + 2),
|
||||
tcpFlags: packet[offset + 13], icmpType: nil, transportOffset: offset)
|
||||
case .udp:
|
||||
guard offset + 8 <= packet.count else { throw UMNError.invalidFrame }
|
||||
let length = Int(read16(packet, offset + 4))
|
||||
guard length >= 8, offset + length == packet.count else { throw UMNError.invalidFrame }
|
||||
return .init(source: source, destination: destination, transport: transport,
|
||||
sourcePort: read16(packet, offset), destinationPort: read16(packet, offset + 2),
|
||||
tcpFlags: nil, icmpType: nil, transportOffset: offset)
|
||||
case .icmpv6:
|
||||
guard offset + 8 <= packet.count else { throw UMNError.invalidFrame }
|
||||
return .init(source: source, destination: destination, transport: transport,
|
||||
sourcePort: nil, destinationPort: nil, tcpFlags: nil,
|
||||
icmpType: packet[offset], transportOffset: offset)
|
||||
}
|
||||
}
|
||||
|
||||
private static func read16(_ data: Data, _ offset: Int) -> UInt16 {
|
||||
UInt16(data[offset]) << 8 | UInt16(data[offset + 1])
|
||||
}
|
||||
private static func validateOptions(_ data: Data, range: Range<Int>) throws {
|
||||
var index = range.lowerBound
|
||||
while index < range.upperBound {
|
||||
if data[index] == 0 { index += 1; continue } // Pad1
|
||||
guard index + 2 <= range.upperBound else { throw UMNError.invalidFrame }
|
||||
let length = Int(data[index + 1])
|
||||
guard index + 2 + length <= range.upperBound else { throw UMNError.invalidFrame }
|
||||
index += 2 + length
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
public final class NativePacketFirewall: @unchecked Sendable {
|
||||
private struct Flow: Hashable {
|
||||
let kind: FirewallProtocol
|
||||
let local: MeshAddress
|
||||
let remote: MeshAddress
|
||||
let localPort: UInt16
|
||||
let remotePort: UInt16
|
||||
}
|
||||
private struct Entry { var lastSeen: Date; var established: Bool }
|
||||
private let rules: MeshFirewall
|
||||
private let lock = NSLock()
|
||||
private var flows: [Flow: Entry] = [:]
|
||||
private var echoWindows: [MeshAddress: (Date, Int)] = [:]
|
||||
public let maximumFlows: Int
|
||||
|
||||
public init(rules: MeshFirewall, maximumFlows: Int = 4096) {
|
||||
self.rules = rules; self.maximumFlows = maximumFlows
|
||||
}
|
||||
|
||||
/// Records an OS-originated packet. Outbound connections are always permitted.
|
||||
public func allowOutbound(_ info: IPv6PacketInfo, now: Date = Date()) -> Bool {
|
||||
guard let kind = protocolKind(info), let localPort = info.sourcePort, let remotePort = info.destinationPort else {
|
||||
return info.transport == .icmpv6 && [1, 2, 3, 4, 128, 129].contains(info.icmpType ?? 255)
|
||||
}
|
||||
lock.lock(); defer { lock.unlock() }
|
||||
expireLocked(now)
|
||||
let key = Flow(kind: kind, local: info.source, remote: info.destination,
|
||||
localPort: localPort, remotePort: remotePort)
|
||||
insert(key, entry: .init(lastSeen: now, established: kind == .udp || ((info.tcpFlags ?? 0) & 0x02) != 0))
|
||||
return true
|
||||
}
|
||||
|
||||
public func allowInbound(_ info: IPv6PacketInfo, packet: Data? = nil, now: Date = Date()) -> Bool {
|
||||
if info.transport == .icmpv6 { return allowICMP(info, packet: packet, now: now) }
|
||||
guard let kind = protocolKind(info), let remotePort = info.sourcePort, let localPort = info.destinationPort else { return false }
|
||||
lock.lock(); defer { lock.unlock() }
|
||||
expireLocked(now)
|
||||
let key = Flow(kind: kind, local: info.destination, remote: info.source,
|
||||
localPort: localPort, remotePort: remotePort)
|
||||
if var entry = flows[key] {
|
||||
entry.lastSeen = now
|
||||
if kind == .tcp, ((info.tcpFlags ?? 0) & 0x12) == 0x12 { entry.established = true }
|
||||
flows[key] = entry
|
||||
if kind == .tcp, ((info.tcpFlags ?? 0) & 0x05) != 0 { flows.removeValue(forKey: key) }
|
||||
return true
|
||||
}
|
||||
guard rules.allows(protocol: kind, port: localPort, source: info.source) else { return false }
|
||||
// A new TCP flow must begin with SYN and not ACK. UDP is connectionless.
|
||||
if kind == .tcp, ((info.tcpFlags ?? 0) & 0x12) != 0x02 { return false }
|
||||
insert(key, entry: .init(lastSeen: now, established: kind == .udp))
|
||||
return true
|
||||
}
|
||||
|
||||
private func allowICMP(_ info: IPv6PacketInfo, packet: Data?, now: Date) -> Bool {
|
||||
guard let type = info.icmpType else { return false }
|
||||
if (1...4).contains(type), let packet, packet.count >= info.transportOffset + 8 + 48 {
|
||||
var quoted = Data(packet[(info.transportOffset + 8)...])
|
||||
let actualPayload = quoted.count - 40
|
||||
quoted[4] = UInt8(actualPayload >> 8); quoted[5] = UInt8(actualPayload & 255)
|
||||
guard let original = try? IPv6PacketParser.parse(quoted),
|
||||
let kind = protocolKind(original), let localPort = original.sourcePort,
|
||||
let remotePort = original.destinationPort else { return false }
|
||||
lock.lock(); defer { lock.unlock() }; expireLocked(now)
|
||||
return flows[Flow(kind: kind, local: original.source, remote: original.destination,
|
||||
localPort: localPort, remotePort: remotePort)] != nil
|
||||
}
|
||||
guard type == 128 || type == 129 else { return false }
|
||||
lock.lock(); defer { lock.unlock() }
|
||||
let current = echoWindows[info.source]
|
||||
if let current, now.timeIntervalSince(current.0) < 1 {
|
||||
guard current.1 < 10 else { return false }
|
||||
echoWindows[info.source] = (current.0, current.1 + 1)
|
||||
} else {
|
||||
guard echoWindows[info.source] != nil || echoWindows.count < 128 else { return false }
|
||||
echoWindows[info.source] = (now, 1)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
private func protocolKind(_ info: IPv6PacketInfo) -> FirewallProtocol? {
|
||||
info.transport == .tcp ? .tcp : (info.transport == .udp ? .udp : nil)
|
||||
}
|
||||
private func insert(_ flow: Flow, entry: Entry) {
|
||||
if flows[flow] == nil, flows.count >= maximumFlows,
|
||||
let oldest = flows.min(by: { $0.value.lastSeen < $1.value.lastSeen })?.key { flows.removeValue(forKey: oldest) }
|
||||
flows[flow] = entry
|
||||
}
|
||||
private func expireLocked(_ now: Date) {
|
||||
flows = flows.filter { now.timeIntervalSince($0.value.lastSeen) <= ($0.key.kind == .tcp ? 300 : 60) }
|
||||
echoWindows = echoWindows.filter { now.timeIntervalSince($0.value.0) < 60 }
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
import Foundation
|
||||
|
||||
public enum MeshDNSCodec {
|
||||
public static func response(to query: Data, aliases: [String: MeshAddress], ttl: UInt32 = 30) throws -> Data {
|
||||
guard query.count >= 12, read16(query, 4) == 1 else { throw UMNError.invalidFrame }
|
||||
var offset = 12; var visited = Set<Int>()
|
||||
let labels = try decodeName(query, offset: &offset, visited: &visited)
|
||||
guard offset + 4 <= query.count else { throw UMNError.invalidFrame }
|
||||
let type = read16(query, offset), qclass = read16(query, offset + 2); offset += 4
|
||||
guard qclass == 1 else { throw UMNError.invalidFrame }
|
||||
let question = query.subdata(in: 12..<offset)
|
||||
let fqdn = labels.joined(separator: ".").lowercased()
|
||||
let aliasName = fqdn.hasSuffix(".mesh") ? String(fqdn.dropLast(5)) : ""
|
||||
let address = aliases[aliasName]
|
||||
let answer = address != nil && type == 28
|
||||
let rcode: UInt16 = address == nil ? 3 : 0
|
||||
var output = Data()
|
||||
append16(read16(query, 0), to: &output)
|
||||
append16(0x8400 | rcode, to: &output) // response + authoritative
|
||||
append16(1, to: &output); append16(answer ? 1 : 0, to: &output)
|
||||
append16(0, to: &output); append16(0, to: &output); output.append(question)
|
||||
if answer, let address {
|
||||
append16(0xc00c, to: &output); append16(28, to: &output); append16(1, to: &output)
|
||||
append32(ttl, to: &output); append16(16, to: &output); output.append(address.bytes)
|
||||
}
|
||||
guard output.count <= 512 else { throw UMNError.invalidFrame }
|
||||
return output
|
||||
}
|
||||
|
||||
private static func decodeName(_ data: Data, offset: inout Int, visited: inout Set<Int>, depth: Int = 0) throws -> [String] {
|
||||
guard depth < 16 else { throw UMNError.invalidFrame }
|
||||
var labels: [String] = []
|
||||
while true {
|
||||
guard offset < data.count else { throw UMNError.invalidFrame }
|
||||
let length = Int(data[offset]); offset += 1
|
||||
if length == 0 { return labels }
|
||||
if length & 0xc0 == 0xc0 {
|
||||
guard offset < data.count else { throw UMNError.invalidFrame }
|
||||
let pointer = ((length & 0x3f) << 8) | Int(data[offset]); offset += 1
|
||||
guard pointer >= 12, pointer < data.count, visited.insert(pointer).inserted else { throw UMNError.invalidFrame }
|
||||
var nested = pointer
|
||||
labels += try decodeName(data, offset: &nested, visited: &visited, depth: depth + 1)
|
||||
return labels
|
||||
}
|
||||
guard length <= 63, offset + length <= data.count,
|
||||
let label = String(data: data.subdata(in: offset..<(offset + length)), encoding: .utf8),
|
||||
!label.isEmpty else { throw UMNError.invalidFrame }
|
||||
labels.append(label); offset += length
|
||||
}
|
||||
}
|
||||
private static func read16(_ data: Data, _ offset: Int) -> UInt16 { UInt16(data[offset]) << 8 | UInt16(data[offset + 1]) }
|
||||
private static func append16(_ value: UInt16, to data: inout Data) { data.append(UInt8(value >> 8)); data.append(UInt8(value & 255)) }
|
||||
private static func append32(_ value: UInt32, to data: inout Data) {
|
||||
data.append(UInt8(value >> 24)); data.append(UInt8((value >> 16) & 255)); data.append(UInt8((value >> 8) & 255)); data.append(UInt8(value & 255))
|
||||
}
|
||||
}
|
||||
@@ -102,7 +102,7 @@ public struct LinkState: Codable, Hashable, Sendable {
|
||||
}
|
||||
|
||||
public enum PacketKind: String, Codable, Sendable {
|
||||
case pingRequest, pingReply, text, streamOpen, streamData, streamAck, streamClose, streamReset, error
|
||||
case pingRequest, pingReply, text, streamOpen, streamData, streamAck, streamClose, streamReset, ipv6Packet, error
|
||||
}
|
||||
|
||||
public struct InnerFrame: Codable, Sendable {
|
||||
@@ -113,13 +113,15 @@ public struct InnerFrame: Codable, Sendable {
|
||||
public let acknowledgment: UInt64
|
||||
public let payload: Data
|
||||
public let sourceRecord: NodeRecord
|
||||
/// Required for ipv6Packet frames. It is covered by the end-to-end signature.
|
||||
public let replayID: UUID?
|
||||
|
||||
public init(kind: PacketKind, streamID: UInt64 = 0, port: UInt16? = nil,
|
||||
sequence: UInt64 = 0, acknowledgment: UInt64 = 0,
|
||||
payload: Data = Data(), sourceRecord: NodeRecord) {
|
||||
payload: Data = Data(), sourceRecord: NodeRecord, replayID: UUID? = nil) {
|
||||
self.kind = kind; self.streamID = streamID; self.port = port
|
||||
self.sequence = sequence; self.acknowledgment = acknowledgment
|
||||
self.payload = payload; self.sourceRecord = sourceRecord
|
||||
self.payload = payload; self.sourceRecord = sourceRecord; self.replayID = replayID
|
||||
}
|
||||
}
|
||||
|
||||
@@ -134,16 +136,18 @@ public struct SealedPayload: Codable, Sendable {
|
||||
}
|
||||
|
||||
public struct RoutedPacket: Codable, Sendable {
|
||||
public enum TrafficClass: String, Codable, Sendable { case overlay, nativeIPv6 }
|
||||
public let id: UUID
|
||||
public let source: MeshAddress
|
||||
public let destination: MeshAddress
|
||||
public var hopLimit: UInt8
|
||||
public let sealed: SealedPayload
|
||||
public let trafficClass: TrafficClass
|
||||
|
||||
public init(id: UUID = UUID(), source: MeshAddress, destination: MeshAddress,
|
||||
hopLimit: UInt8 = 16, sealed: SealedPayload) {
|
||||
hopLimit: UInt8 = 16, sealed: SealedPayload, trafficClass: TrafficClass = .overlay) {
|
||||
self.id = id; self.source = source; self.destination = destination
|
||||
self.hopLimit = hopLimit; self.sealed = sealed
|
||||
self.hopLimit = hopLimit; self.sealed = sealed; self.trafficClass = trafficClass
|
||||
}
|
||||
}
|
||||
|
||||
@@ -178,7 +182,7 @@ public enum WireMessage: Codable, Sendable {
|
||||
}
|
||||
|
||||
public struct WireEnvelope: Codable, Sendable {
|
||||
public static let currentVersion: UInt16 = 1
|
||||
public static let currentVersion: UInt16 = 2
|
||||
public let version: UInt16
|
||||
public let message: WireMessage
|
||||
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
import Foundation
|
||||
import Darwin
|
||||
import UltraMeshCore
|
||||
import UMNSystem
|
||||
|
||||
private struct HelperOptions {
|
||||
var uid: uid_t
|
||||
var socket = "/var/run/ultramesh-helper.sock"
|
||||
init() throws {
|
||||
let args = Array(CommandLine.arguments.dropFirst())
|
||||
guard let uidIndex = args.firstIndex(of: "--uid"), uidIndex + 1 < args.count,
|
||||
let parsed = UInt32(args[uidIndex + 1]) else { throw UMNError.message("umn-helper requires --uid") }
|
||||
uid = parsed
|
||||
if let index = args.firstIndex(of: "--socket"), index + 1 < args.count { socket = args[index + 1] }
|
||||
}
|
||||
}
|
||||
|
||||
private func runCommand(_ executable: String, _ arguments: [String]) throws {
|
||||
let process = Process(); process.executableURL = URL(fileURLWithPath: executable); process.arguments = arguments
|
||||
process.standardOutput = FileHandle.nullDevice; process.standardError = FileHandle.standardError
|
||||
try process.run(); process.waitUntilExit()
|
||||
guard process.terminationStatus == 0 else {
|
||||
throw UMNError.message("\(executable) failed with status \(process.terminationStatus)")
|
||||
}
|
||||
}
|
||||
|
||||
private func readLine(_ fd: Int32, limit: Int = 32_768) throws -> String? {
|
||||
var bytes: [UInt8] = []
|
||||
while bytes.count < limit {
|
||||
var byte: UInt8 = 0; let count = Darwin.read(fd, &byte, 1)
|
||||
if count == 0 { return bytes.isEmpty ? nil : String(decoding: bytes, as: UTF8.self) }
|
||||
if count < 0 { if errno == EINTR { continue }; throw POSIXError(.init(rawValue: errno) ?? .EIO) }
|
||||
if byte == 10 { return String(decoding: bytes, as: UTF8.self) }
|
||||
bytes.append(byte)
|
||||
}
|
||||
throw UMNError.invalidFrame
|
||||
}
|
||||
|
||||
private func handle(_ client: Int32, allowedUID: uid_t) {
|
||||
var device: Int32 = -1; var interface = ""; var routes = Set<MeshAddress>()
|
||||
defer {
|
||||
for route in routes { try? runCommand("/sbin/route", ["-n", "delete", "-inet6", route.description]) }
|
||||
if !interface.isEmpty { try? runCommand("/sbin/ifconfig", [interface, "down"]) }
|
||||
if device >= 0 { Darwin.close(device) }
|
||||
UnixIPC.close(client)
|
||||
}
|
||||
do {
|
||||
var peerUID: uid_t = 0; var peerGID: gid_t = 0
|
||||
guard umn_get_peer_eid(client, &peerUID, &peerGID) == 0, peerUID == allowedUID else {
|
||||
throw UMNError.message("helper authentication failed")
|
||||
}
|
||||
guard let hello = try readLine(client), hello.hasPrefix("UMN2 "),
|
||||
let local = try? MeshAddress(String(hello.dropFirst(5))) else { throw UMNError.invalidFrame }
|
||||
var name = [CChar](repeating: 0, count: Int(IFNAMSIZ))
|
||||
device = umn_utun_open(&name, name.count)
|
||||
guard device >= 0 else { throw POSIXError(.init(rawValue: errno) ?? .EIO) }
|
||||
interface = String(cString: name)
|
||||
try runCommand("/sbin/ifconfig", [interface, "inet6", local.description, "prefixlen", "128", "mtu", "1280", "up"])
|
||||
let ready = Data("READY \(interface) 1280\n".utf8)
|
||||
let sent = ready.withUnsafeBytes { umn_send_fd(client, device, $0.baseAddress, ready.count) }
|
||||
guard sent == 0 else { throw POSIXError(.init(rawValue: errno) ?? .EIO) }
|
||||
|
||||
while let line = try readLine(client) {
|
||||
guard line.hasPrefix("ROUTES") else { throw UMNError.invalidFrame }
|
||||
let fields = line.split(separator: " ").dropFirst()
|
||||
var desired = Set<MeshAddress>()
|
||||
for field in fields {
|
||||
guard let address = try? MeshAddress(String(field)) else { throw UMNError.invalidFrame }
|
||||
desired.insert(address)
|
||||
}
|
||||
try HelperRouteSet.validate(desired, local: local)
|
||||
let diff = HelperRouteSet.diff(current: routes, desired: desired)
|
||||
for address in diff.removals {
|
||||
try runCommand("/sbin/route", ["-n", "delete", "-inet6", address.description]); routes.remove(address)
|
||||
}
|
||||
for address in diff.additions {
|
||||
try runCommand("/sbin/route", ["-n", "add", "-inet6", address.description, "-interface", interface]); routes.insert(address)
|
||||
}
|
||||
}
|
||||
} catch {
|
||||
FileHandle.standardError.write(Data("umn-helper: \(error.localizedDescription)\n".utf8))
|
||||
}
|
||||
}
|
||||
|
||||
do {
|
||||
guard geteuid() == 0 else { throw UMNError.message("umn-helper must run as root") }
|
||||
let options = try HelperOptions()
|
||||
let server = try UnixIPC.listen(path: options.socket)
|
||||
guard chown(options.socket, options.uid, gid_t.max) == 0, chmod(options.socket, 0o600) == 0 else {
|
||||
throw POSIXError(.init(rawValue: errno) ?? .EIO)
|
||||
}
|
||||
signal(SIGPIPE, SIG_IGN)
|
||||
while true { let client = try UnixIPC.accept(server); handle(client, allowedUID: options.uid) }
|
||||
} catch {
|
||||
FileHandle.standardError.write(Data("umn-helper: \(error.localizedDescription)\n".utf8)); exit(1)
|
||||
}
|
||||
@@ -1,5 +1,7 @@
|
||||
import Foundation
|
||||
import Darwin
|
||||
import UltraMeshCore
|
||||
import UMNSystem
|
||||
|
||||
struct TestFailure: Error, LocalizedError {
|
||||
let errorDescription: String?
|
||||
@@ -12,6 +14,31 @@ func check(_ condition: @autoclosure () throws -> Bool, _ name: String) throws {
|
||||
passed += 1; print("ok \(passed) - \(name)")
|
||||
}
|
||||
|
||||
func ipv6Packet(source: MeshAddress, destination: MeshAddress, next: UInt8, payload: Data) -> Data {
|
||||
var packet = Data(repeating: 0, count: 40)
|
||||
packet[0] = 0x60; packet[4] = UInt8(payload.count >> 8); packet[5] = UInt8(payload.count & 255)
|
||||
packet[6] = next; packet[7] = 64
|
||||
packet.replaceSubrange(8..<24, with: source.bytes); packet.replaceSubrange(24..<40, with: destination.bytes)
|
||||
packet.append(payload); return packet
|
||||
}
|
||||
func tcpHeader(source: UInt16, destination: UInt16, flags: UInt8) -> Data {
|
||||
var data = Data(repeating: 0, count: 20)
|
||||
data[0] = UInt8(source >> 8); data[1] = UInt8(source & 255)
|
||||
data[2] = UInt8(destination >> 8); data[3] = UInt8(destination & 255)
|
||||
data[12] = 0x50; data[13] = flags; return data
|
||||
}
|
||||
func udpHeader(source: UInt16, destination: UInt16, payload: Data = Data()) -> Data {
|
||||
let length = 8 + payload.count; var data = Data(repeating: 0, count: 8)
|
||||
data[0] = UInt8(source >> 8); data[1] = UInt8(source & 255)
|
||||
data[2] = UInt8(destination >> 8); data[3] = UInt8(destination & 255)
|
||||
data[4] = UInt8(length >> 8); data[5] = UInt8(length & 255); data.append(payload); return data
|
||||
}
|
||||
func dnsQuery(_ name: String, type: UInt16) -> Data {
|
||||
var data = Data([0x12, 0x34, 0x01, 0x00, 0, 1, 0, 0, 0, 0, 0, 0])
|
||||
for label in name.split(separator: ".") { data.append(UInt8(label.utf8.count)); data.append(contentsOf: label.utf8) }
|
||||
data.append(0); data.append(UInt8(type >> 8)); data.append(UInt8(type & 255)); data.append(contentsOf: [0, 1]); return data
|
||||
}
|
||||
|
||||
do {
|
||||
let alice = try NodeIdentity(); let bob = try NodeIdentity(); let carol = try NodeIdentity()
|
||||
try check(try MeshAddress(alice.record.address.description) == alice.record.address, "mesh address round trip")
|
||||
@@ -35,6 +62,10 @@ do {
|
||||
var rejected = false
|
||||
do { _ = try bob.open(tampered, expectedSource: alice.record.address) } catch { rejected = true }
|
||||
try check(rejected, "tampered ciphertext rejected")
|
||||
let replayID = UUID()
|
||||
let nativeFrame = InnerFrame(kind: .ipv6Packet, payload: Data([1, 2]), sourceRecord: alice.record, replayID: replayID)
|
||||
let nativeOpened = try bob.open(alice.seal(nativeFrame, to: bob.record), expectedSource: alice.record.address)
|
||||
try check(nativeOpened.replayID == replayID, "signed native packet replay UUID round trip")
|
||||
|
||||
let router = LinkStateRouter(local: alice.record.address)
|
||||
try check(router.ingest(try alice.makeLinkState(sequence: 1, neighbors: [bob.record.address])), "local topology accepted")
|
||||
@@ -53,11 +84,109 @@ do {
|
||||
try check(!firewall.allows(port: 80, source: bob.record.address), "scoped firewall rule denies other source")
|
||||
firewall.allow(port: 443, source: nil)
|
||||
try check(firewall.allows(port: 443, source: bob.record.address), "any-source firewall rule")
|
||||
let migrated = try JSONDecoder().decode(FirewallRule.self, from: Data("{\"port\":7000}".utf8))
|
||||
try check(migrated.protocolKind == .overlay, "protocol-less firewall rule migrates to overlay")
|
||||
firewall.allow(protocol: .tcp, port: 8080, source: nil)
|
||||
try check(firewall.allows(protocol: .tcp, port: 8080, source: carol.record.address) &&
|
||||
!firewall.allows(protocol: .udp, port: 8080, source: carol.record.address), "firewall protocols remain distinct")
|
||||
|
||||
let syn = ipv6Packet(source: alice.record.address, destination: bob.record.address, next: 6,
|
||||
payload: tcpHeader(source: 50_000, destination: 8080, flags: 0x02))
|
||||
let synInfo = try IPv6PacketParser.parse(syn)
|
||||
try check(synInfo.transport == .tcp && synInfo.destinationPort == 8080, "TCP packet parsed")
|
||||
let udp = ipv6Packet(source: alice.record.address, destination: bob.record.address, next: 17,
|
||||
payload: udpHeader(source: 50_001, destination: 5353, payload: Data([1, 2, 3])))
|
||||
try check(try IPv6PacketParser.parse(udp).transport == .udp, "UDP packet parsed")
|
||||
let echo = ipv6Packet(source: alice.record.address, destination: bob.record.address, next: 58,
|
||||
payload: Data([128, 0, 0, 0, 0, 1, 0, 1]))
|
||||
try check(try IPv6PacketParser.parse(echo).icmpType == 128, "ICMPv6 echo parsed")
|
||||
let extensionHeader = Data([6, 0, 0, 0, 0, 0, 0, 0]) + tcpHeader(source: 1, destination: 2, flags: 0x02)
|
||||
let extended = ipv6Packet(source: alice.record.address, destination: bob.record.address, next: 0, payload: extensionHeader)
|
||||
try check(try IPv6PacketParser.parse(extended).transportOffset == 48, "IPv6 extension header parsed safely")
|
||||
var badOption = extensionHeader; badOption[2] = 5; badOption[3] = 10
|
||||
var badOptionRejected = false
|
||||
do { _ = try IPv6PacketParser.parse(ipv6Packet(source: alice.record.address, destination: bob.record.address,
|
||||
next: 0, payload: badOption)) } catch { badOptionRejected = true }
|
||||
try check(badOptionRejected, "malformed IPv6 option rejected")
|
||||
var fragmentRejected = false
|
||||
do { _ = try IPv6PacketParser.parse(ipv6Packet(source: alice.record.address, destination: bob.record.address,
|
||||
next: 44, payload: Data(repeating: 0, count: 8))) } catch { fragmentRejected = true }
|
||||
try check(fragmentRejected, "IPv6 fragments rejected")
|
||||
var malformedRejected = false; var malformedUDP = udp; malformedUDP[44] = 0; malformedUDP[45] = 8
|
||||
do { _ = try IPv6PacketParser.parse(malformedUDP) } catch { malformedRejected = true }
|
||||
try check(malformedRejected, "malformed UDP length rejected")
|
||||
try check(synInfo.source != carol.record.address, "packet source exposes spoof mismatch")
|
||||
|
||||
let packetFrame = InnerFrame(kind: .ipv6Packet, payload: syn, sourceRecord: alice.record, replayID: UUID())
|
||||
let routedNative = RoutedPacket(source: alice.record.address, destination: carol.record.address,
|
||||
sealed: try alice.seal(packetFrame, to: carol.record), trafficClass: .nativeIPv6)
|
||||
try check(router.routes()[carol.record.address]?.nextHop == bob.record.address, "native packet selects three-node next hop")
|
||||
var relayCannotDecrypt = false
|
||||
do { _ = try bob.open(routedNative.sealed, expectedSource: alice.record.address) } catch { relayCannotDecrypt = true }
|
||||
let destinationFrame = try carol.open(routedNative.sealed, expectedSource: alice.record.address)
|
||||
try check(relayCannotDecrypt && destinationFrame.payload == syn, "native packet relays end-to-end encrypted")
|
||||
var replayCache = Set<UUID>(); let packetReplayID = destinationFrame.replayID!
|
||||
try check(replayCache.insert(packetReplayID).inserted && !replayCache.insert(packetReplayID).inserted,
|
||||
"native replay UUID rejects duplicate")
|
||||
var nativeBytes = 0
|
||||
for _ in 0..<900 {
|
||||
let large = ipv6Packet(source: alice.record.address, destination: carol.record.address, next: 17,
|
||||
payload: udpHeader(source: 40_000, destination: 8080,
|
||||
payload: Data(repeating: 0x5a, count: 1_160)))
|
||||
_ = try IPv6PacketParser.parse(large); nativeBytes += large.count
|
||||
}
|
||||
try check(nativeBytes > 1_048_576, "native packet path exceeds legacy 1 MiB stream cap")
|
||||
try check(router.ingest(try alice.makeLinkState(sequence: 2, neighbors: [carol.record.address])) &&
|
||||
router.ingest(try carol.makeLinkState(sequence: 2, neighbors: [alice.record.address])), "route-change topology accepted")
|
||||
try check(router.routes()[carol.record.address]?.nextHop == carol.record.address,
|
||||
"retransmitted native packet uses current route")
|
||||
|
||||
let nativeRules = MeshFirewall(); let stateful = NativePacketFirewall(rules: nativeRules, maximumFlows: 4)
|
||||
try check(stateful.allowOutbound(synInfo), "outbound TCP allowed and tracked")
|
||||
let synAck = ipv6Packet(source: bob.record.address, destination: alice.record.address, next: 6,
|
||||
payload: tcpHeader(source: 8080, destination: 50_000, flags: 0x12))
|
||||
try check(stateful.allowInbound(try IPv6PacketParser.parse(synAck)), "TCP return packet admitted")
|
||||
let unsolicited = ipv6Packet(source: bob.record.address, destination: alice.record.address, next: 6,
|
||||
payload: tcpHeader(source: 9000, destination: 9001, flags: 0x02))
|
||||
try check(!stateful.allowInbound(try IPv6PacketParser.parse(unsolicited)), "new inbound TCP denied by default")
|
||||
nativeRules.allow(protocol: .tcp, port: 9001, source: bob.record.address)
|
||||
try check(stateful.allowInbound(try IPv6PacketParser.parse(unsolicited)), "source-scoped inbound TCP admitted")
|
||||
|
||||
let aaaaResponse = try MeshDNSCodec.response(to: dnsQuery("alice.mesh", type: 28), aliases: ["alice": alice.record.address])
|
||||
try check(aaaaResponse[7] == 1 && aaaaResponse.suffix(16) == alice.record.address.bytes, "DNS AAAA response")
|
||||
let noData = try MeshDNSCodec.response(to: dnsQuery("alice.mesh", type: 1), aliases: ["alice": alice.record.address])
|
||||
try check(noData[3] & 0x0f == 0 && noData[7] == 0, "DNS A query returns NODATA")
|
||||
let notFound = try MeshDNSCodec.response(to: dnsQuery("missing.mesh", type: 28), aliases: [:])
|
||||
try check(notFound[3] & 0x0f == 3, "DNS unknown alias returns NXDOMAIN")
|
||||
var compressionRejected = false
|
||||
var compressed = dnsQuery("x.mesh", type: 28); compressed.replaceSubrange(12..<20, with: Data([0xc0, 0x0c, 0, 28, 0, 1]))
|
||||
do { _ = try MeshDNSCodec.response(to: compressed, aliases: [:]) } catch { compressionRejected = true }
|
||||
try check(compressionRejected, "DNS compression loop rejected")
|
||||
|
||||
let desired: Set<MeshAddress> = [bob.record.address, carol.record.address]
|
||||
try HelperRouteSet.validate(desired, local: alice.record.address)
|
||||
let routeDiff = HelperRouteSet.diff(current: [bob.record.address], desired: desired)
|
||||
try check(routeDiff.additions == [carol.record.address] && routeDiff.removals.isEmpty, "helper route-set diff is idempotent")
|
||||
var localRouteRejected = false
|
||||
do { try HelperRouteSet.validate([alice.record.address], local: alice.record.address) } catch { localRouteRejected = true }
|
||||
try check(localRouteRejected, "helper rejects a route to the local identity")
|
||||
|
||||
var sockets = [Int32](repeating: -1, count: 2); var pipeFDs = [Int32](repeating: -1, count: 2)
|
||||
guard socketpair(AF_UNIX, SOCK_STREAM, 0, &sockets) == 0, pipe(&pipeFDs) == 0 else { throw TestFailure("socket setup") }
|
||||
defer { sockets.forEach { Darwin.close($0) }; pipeFDs.forEach { Darwin.close($0) } }
|
||||
var peerUID: uid_t = 0, peerGID: gid_t = 0
|
||||
try check(umn_get_peer_eid(sockets[0], &peerUID, &peerGID) == 0 && peerUID == getuid(), "helper peer UID authentication primitive")
|
||||
let marker = Data("FD\n".utf8)
|
||||
let fdSent = marker.withUnsafeBytes { umn_send_fd(sockets[0], pipeFDs[0], $0.baseAddress, marker.count) }
|
||||
var receivedFD: Int32 = -1; var fdBuffer = [UInt8](repeating: 0, count: 8)
|
||||
let fdCount = umn_recv_fd(sockets[1], &receivedFD, &fdBuffer, fdBuffer.count)
|
||||
defer { if receivedFD >= 0 { Darwin.close(receivedFD) } }
|
||||
try check(fdSent == 0 && fdCount == marker.count && receivedFD >= 0, "utun descriptor passing primitive")
|
||||
|
||||
let framed = try FrameCodec.encode(WireEnvelope(message: .keepalive))
|
||||
try check(try FrameCodec.bodyLength(from: framed.prefix(4)) == framed.count - 4, "frame length prefix")
|
||||
let decoded = try FrameCodec.decode(WireEnvelope.self, body: framed.dropFirst(4))
|
||||
if decoded.version == 1, case .keepalive = decoded.message { try check(true, "wire frame round trip") }
|
||||
if decoded.version == 2, case .keepalive = decoded.message { try check(true, "wire v2 frame round trip") }
|
||||
else { throw TestFailure("wire frame round trip") }
|
||||
var zeroRejected = false
|
||||
do { _ = try FrameCodec.bodyLength(from: Data([0, 0, 0, 0])) } catch { zeroRejected = true }
|
||||
|
||||
@@ -6,9 +6,9 @@ let args = Array(CommandLine.arguments.dropFirst())
|
||||
func usage() -> Never {
|
||||
print("""
|
||||
Ultra Mesh Network
|
||||
umn status | address | peers | routes | ping <target>
|
||||
umn status | address | peers | routes | interface status | ping <target>
|
||||
umn alias set <name> <address> | remove <name> | list
|
||||
umn firewall allow|revoke <port> --from <address|any> | list
|
||||
umn firewall allow|revoke <overlay|tcp|udp> <port> --from <address|any> | list
|
||||
umn text send <target> <port> <message> | listen <port>
|
||||
umn expose <mesh-port> --to <host:port>
|
||||
umn proxy [--listen 127.0.0.1:<port>]
|
||||
@@ -44,6 +44,9 @@ do {
|
||||
case "address": request = IPCMessage(operation: .address)
|
||||
case "peers": request = IPCMessage(operation: .peers)
|
||||
case "routes": request = IPCMessage(operation: .routes)
|
||||
case "interface":
|
||||
guard args.count == 2, args[1] == "status" else { usage() }
|
||||
request = IPCMessage(operation: .interfaceStatus)
|
||||
case "ping":
|
||||
guard args.count == 2 else { usage() }
|
||||
request = IPCMessage(operation: .ping); request.target = args[1]
|
||||
@@ -63,9 +66,10 @@ do {
|
||||
guard args.count >= 2 else { usage() }
|
||||
if args[1] == "list" { request = IPCMessage(operation: .firewallList) }
|
||||
else {
|
||||
guard args.count == 5, args[3] == "--from", ["allow", "revoke"].contains(args[1]) else { usage() }
|
||||
guard args.count == 6, args[4] == "--from", ["allow", "revoke"].contains(args[1]),
|
||||
let protocolKind = FirewallProtocol(rawValue: args[2]) else { usage() }
|
||||
request = IPCMessage(operation: args[1] == "allow" ? .firewallAllow : .firewallRevoke)
|
||||
request.port = parsePort(args[2]); request.source = args[4]
|
||||
request.firewallProtocol = protocolKind; request.port = parsePort(args[3]); request.source = args[5]
|
||||
}
|
||||
case "text":
|
||||
guard args.count >= 3 else { usage() }
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
import Foundation
|
||||
import Darwin
|
||||
import UltraMeshCore
|
||||
|
||||
final class MeshDNSServer {
|
||||
static let port: UInt16 = 53535
|
||||
private let aliases: () -> [String: MeshAddress]
|
||||
private var udp: Int32 = -1, tcp: Int32 = -1
|
||||
private var running = false
|
||||
|
||||
init(aliases: @escaping () -> [String: MeshAddress]) { self.aliases = aliases }
|
||||
|
||||
func start() throws {
|
||||
udp = try bindSocket(type: SOCK_DGRAM); tcp = try bindSocket(type: SOCK_STREAM)
|
||||
guard Darwin.listen(tcp, 32) == 0 else { throw POSIXError(.init(rawValue: errno) ?? .EIO) }
|
||||
running = true
|
||||
DispatchQueue.global(qos: .utility).async { [weak self] in self?.udpLoop() }
|
||||
DispatchQueue.global(qos: .utility).async { [weak self] in self?.tcpLoop() }
|
||||
}
|
||||
func stop() {
|
||||
running = false
|
||||
if udp >= 0 { Darwin.close(udp); udp = -1 }
|
||||
if tcp >= 0 { Darwin.close(tcp); tcp = -1 }
|
||||
}
|
||||
|
||||
private func bindSocket(type: Int32) throws -> Int32 {
|
||||
let fd = socket(AF_INET, type, 0); guard fd >= 0 else { throw POSIXError(.init(rawValue: errno) ?? .EIO) }
|
||||
var yes: Int32 = 1; setsockopt(fd, SOL_SOCKET, SO_REUSEADDR, &yes, socklen_t(MemoryLayout<Int32>.size))
|
||||
var address = sockaddr_in(); address.sin_len = UInt8(MemoryLayout<sockaddr_in>.size)
|
||||
address.sin_family = sa_family_t(AF_INET); address.sin_port = Self.port.bigEndian
|
||||
address.sin_addr = in_addr(s_addr: inet_addr("127.0.0.1"))
|
||||
let result = withUnsafePointer(to: &address) { pointer in
|
||||
pointer.withMemoryRebound(to: sockaddr.self, capacity: 1) { Darwin.bind(fd, $0, socklen_t(MemoryLayout<sockaddr_in>.size)) }
|
||||
}
|
||||
guard result == 0 else { let code = errno; Darwin.close(fd); throw POSIXError(.init(rawValue: code) ?? .EIO) }
|
||||
return fd
|
||||
}
|
||||
private func udpLoop() {
|
||||
while running {
|
||||
var bytes = [UInt8](repeating: 0, count: 4096); var peer = sockaddr_storage(); var length = socklen_t(MemoryLayout.size(ofValue: peer))
|
||||
let count = withUnsafeMutablePointer(to: &peer) { pointer in
|
||||
pointer.withMemoryRebound(to: sockaddr.self, capacity: 1) { recvfrom(udp, &bytes, bytes.count, 0, $0, &length) }
|
||||
}
|
||||
guard count > 0 else { if errno == EINTR { continue }; return }
|
||||
guard let response = try? MeshDNSCodec.response(to: Data(bytes.prefix(count)), aliases: aliases()) else { continue }
|
||||
response.withUnsafeBytes { body in
|
||||
withUnsafePointer(to: &peer) { pointer in
|
||||
pointer.withMemoryRebound(to: sockaddr.self, capacity: 1) { _ = sendto(udp, body.baseAddress, response.count, 0, $0, length) }
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
private func tcpLoop() {
|
||||
while running {
|
||||
let client = Darwin.accept(tcp, nil, nil)
|
||||
if client < 0 { if errno == EINTR { continue }; return }
|
||||
DispatchQueue.global(qos: .utility).async { [weak self] in self?.handleTCP(client) }
|
||||
}
|
||||
}
|
||||
private func handleTCP(_ fd: Int32) {
|
||||
defer { Darwin.close(fd) }
|
||||
guard let prefix = readExact(2, fd), prefix.count == 2 else { return }
|
||||
let length = Int(prefix[0]) << 8 | Int(prefix[1]); guard length > 0, length <= 4096,
|
||||
let query = readExact(length, fd), let response = try? MeshDNSCodec.response(to: query, aliases: aliases()) else { return }
|
||||
var framed = Data([UInt8(response.count >> 8), UInt8(response.count & 255)]); framed.append(response)
|
||||
framed.withUnsafeBytes { bytes in
|
||||
guard let base = bytes.baseAddress else { return }
|
||||
var offset = 0
|
||||
while offset < framed.count {
|
||||
let count = Darwin.write(fd, base.advanced(by: offset), framed.count - offset)
|
||||
if count <= 0 { return }
|
||||
offset += count
|
||||
}
|
||||
}
|
||||
}
|
||||
private func readExact(_ count: Int, _ fd: Int32) -> Data? {
|
||||
var output = Data(); var buffer = [UInt8](repeating: 0, count: count)
|
||||
while output.count < count {
|
||||
let n = Darwin.read(fd, &buffer, count - output.count); if n <= 0 { return nil }
|
||||
output.append(buffer, count: n)
|
||||
}
|
||||
return output
|
||||
}
|
||||
}
|
||||
@@ -24,6 +24,7 @@ final class MeshDaemon {
|
||||
private let identity: NodeIdentity
|
||||
private let router: LinkStateRouter
|
||||
private let firewall: MeshFirewall
|
||||
private let nativeFirewall: NativePacketFirewall
|
||||
private let configURL: URL
|
||||
private var config: DaemonConfig
|
||||
private let ipc: IPCServer
|
||||
@@ -35,12 +36,19 @@ final class MeshDaemon {
|
||||
private var lsaSequence: UInt64 = 0
|
||||
private var seenPackets = Set<UUID>()
|
||||
private var seenOrder: [UUID] = []
|
||||
private var replayPackets = Set<UUID>()
|
||||
private var replayOrder: [UUID] = []
|
||||
private var textBindings: [UInt16: IPCClient] = [:]
|
||||
private var streamBindings: [UInt16: IPCClient] = [:]
|
||||
private var streams: [UInt64: StreamState] = [:]
|
||||
private var pendingPings: [UUID: (IPCClient, Date)] = [:]
|
||||
private var pingWindows: [MeshAddress: (Date, Int)] = [:]
|
||||
private var timer: DispatchSourceTimer?
|
||||
private var nativeInterface: NativeInterfaceClient?
|
||||
private var synchronizedRoutes = Set<MeshAddress>()
|
||||
private var nextInterfaceAttempt = Date.distantPast
|
||||
private var dns: MeshDNSServer?
|
||||
private var dnsState = "stopped"
|
||||
|
||||
init(baseURL: URL = FileManager.default.homeDirectoryForCurrentUser
|
||||
.appendingPathComponent("Library/Application Support/UltraMesh"), socketPath: String? = nil) throws {
|
||||
@@ -50,6 +58,7 @@ final class MeshDaemon {
|
||||
configURL = baseURL.appendingPathComponent("config.json")
|
||||
config = DaemonConfig.load(from: configURL)
|
||||
firewall = MeshFirewall(rules: config.firewall)
|
||||
nativeFirewall = NativePacketFirewall(rules: firewall)
|
||||
ipc = IPCServer(path: socketPath ?? baseURL.appendingPathComponent("umnd.sock").path)
|
||||
let wallClockBase = UInt64(Date().timeIntervalSince1970 * 1000)
|
||||
lsaSequence = max((config.lsaSequence ?? 0) &+ 1_000_000, wallClockBase)
|
||||
@@ -63,6 +72,8 @@ final class MeshDaemon {
|
||||
try startNetwork()
|
||||
ipc.onMessage = { [weak self] client, message in self?.queue.async { self?.handleIPC(client, message) } }
|
||||
try ipc.start()
|
||||
startDNS()
|
||||
attemptInterface()
|
||||
publishLocalState()
|
||||
let timer = DispatchSource.makeTimerSource(queue: queue)
|
||||
timer.schedule(deadline: .now() + 1, repeating: 1)
|
||||
@@ -74,6 +85,7 @@ final class MeshDaemon {
|
||||
func stop() {
|
||||
queue.sync {
|
||||
timer?.cancel(); browser?.cancel(); listener?.cancel()
|
||||
nativeInterface?.disconnect(); dns?.stop()
|
||||
allSessions.values.forEach { $0.stop() }; ipc.stop()
|
||||
}
|
||||
}
|
||||
@@ -185,7 +197,8 @@ final class MeshDaemon {
|
||||
@discardableResult private func send(_ inner: InnerFrame, to remote: NodeRecord) -> Bool {
|
||||
guard remote.address != address,
|
||||
let sealed = try? identity.seal(inner, to: remote) else { return false }
|
||||
let packet = RoutedPacket(source: address, destination: remote.address, sealed: sealed)
|
||||
let packet = RoutedPacket(source: address, destination: remote.address, sealed: sealed,
|
||||
trafficClass: inner.kind == .ipv6Packet ? .nativeIPv6 : .overlay)
|
||||
guard router.routes()[remote.address] != nil || peers[remote.address] != nil else { return false }
|
||||
forward(packet); return true
|
||||
}
|
||||
@@ -214,6 +227,12 @@ final class MeshDaemon {
|
||||
case .streamData: receiveStreamData(inner)
|
||||
case .streamAck: receiveStreamAck(inner)
|
||||
case .streamClose, .streamReset: receiveStreamClose(inner)
|
||||
case .ipv6Packet:
|
||||
guard let replayID = inner.replayID, markReplay(replayID),
|
||||
let info = try? IPv6PacketParser.parse(inner.payload),
|
||||
info.source == packet.source, info.destination == address,
|
||||
nativeFirewall.allowInbound(info, packet: inner.payload) else { return }
|
||||
_ = nativeInterface?.inject(inner.payload)
|
||||
case .error: receiveStreamClose(inner)
|
||||
}
|
||||
}
|
||||
@@ -304,9 +323,23 @@ final class MeshDaemon {
|
||||
response.ok = ok; response.message = text; response.values = values; response.streamID = streamID
|
||||
client.send(response)
|
||||
}
|
||||
guard message.version == 2 else { reply(false, "incompatible CLI protocol version"); return }
|
||||
switch message.operation {
|
||||
case .status:
|
||||
reply(true, "umnd running as \(address); \(peers.count) direct peer(s), \(router.routes().count) route(s)")
|
||||
let state = nativeInterface?.status
|
||||
reply(true, "umnd running as \(address); \(peers.count) direct peer(s), \(router.routes().count) route(s); " +
|
||||
"interface \(state?.name ?? "unavailable") MTU 1280, helper \(state?.helperState ?? "disconnected"), " +
|
||||
"\(state?.routeCount ?? 0) route(s) synchronized, DNS \(dnsState)")
|
||||
case .interfaceStatus:
|
||||
let state = nativeInterface?.status
|
||||
let name = state?.name ?? "unavailable"
|
||||
var response = IPCMessage(operation: message.operation, requestID: message.requestID)
|
||||
response.ok = true; response.interfaceName = state?.name; response.interfaceAddress = address.description
|
||||
response.interfaceMTU = 1280; response.helperState = state?.helperState ?? "disconnected"
|
||||
response.synchronizedRouteCount = state?.routeCount ?? 0; response.dnsState = dnsState
|
||||
response.message = "interface: \(name)\naddress: \(address)\nMTU: 1280\nhelper: \(response.helperState!)\n" +
|
||||
"routes synchronized: \(response.synchronizedRouteCount!)\nDNS: \(dnsState)"
|
||||
client.send(response)
|
||||
case .address: reply(true, address.description)
|
||||
case .peers: reply(true, values: peers.keys.sorted().map(\.description))
|
||||
case .routes:
|
||||
@@ -323,14 +356,16 @@ final class MeshDaemon {
|
||||
guard let name = message.name?.lowercased() else { reply(false, "missing alias"); return }
|
||||
config.aliases.removeValue(forKey: name); saveConfig(); reply(true)
|
||||
case .firewallList:
|
||||
let values = firewall.allRules().map { "allow \($0.port) from \($0.source?.description ?? "any")" }
|
||||
let values = firewall.allRules().map { "allow \($0.protocolKind.rawValue) \($0.port) from \($0.source?.description ?? "any")" }
|
||||
reply(true, values: values)
|
||||
case .firewallAllow, .firewallRevoke:
|
||||
guard let port = message.port, let sourceText = message.source else { reply(false, "missing port or source"); return }
|
||||
guard let protocolKind = message.firewallProtocol, let port = message.port, let sourceText = message.source else {
|
||||
reply(false, "new firewall rules require overlay, tcp, or udp"); return
|
||||
}
|
||||
let source: MeshAddress?
|
||||
if sourceText == "any" { source = nil } else if let parsed = resolve(sourceText) { source = parsed } else { reply(false, "invalid source"); return }
|
||||
if message.operation == .firewallAllow { firewall.allow(port: port, source: source) }
|
||||
else { firewall.revoke(port: port, source: source) }
|
||||
if message.operation == .firewallAllow { firewall.allow(protocol: protocolKind, port: port, source: source) }
|
||||
else { firewall.revoke(protocol: protocolKind, port: port, source: source) }
|
||||
config.firewall = Set(firewall.allRules()); saveConfig(); reply(true)
|
||||
case .ping:
|
||||
guard let target = message.target, let remote = record(for: target) else { reply(false, "no route to destination"); return }
|
||||
@@ -410,6 +445,51 @@ final class MeshDaemon {
|
||||
return true
|
||||
}
|
||||
|
||||
private func markReplay(_ id: UUID) -> Bool {
|
||||
guard !replayPackets.contains(id) else { return false }
|
||||
replayPackets.insert(id); replayOrder.append(id)
|
||||
if replayOrder.count > 4096 { replayPackets.remove(replayOrder.removeFirst()) }
|
||||
return true
|
||||
}
|
||||
|
||||
private func startDNS() {
|
||||
let server = MeshDNSServer { [weak self] in
|
||||
guard let self else { return [:] }
|
||||
return self.queue.sync { self.config.aliases }
|
||||
}
|
||||
do { try server.start(); dns = server; dnsState = "listening on 127.0.0.1:53535" }
|
||||
catch { dnsState = "unavailable (\(error.localizedDescription))"; log("DNS unavailable: \(error)") }
|
||||
}
|
||||
|
||||
private func attemptInterface() {
|
||||
guard nativeInterface == nil, Date() >= nextInterfaceAttempt else { return }
|
||||
nextInterfaceAttempt = Date().addingTimeInterval(5)
|
||||
let socket = ProcessInfo.processInfo.environment["UMN_HELPER_SOCKET"] ?? "/var/run/ultramesh-helper.sock"
|
||||
let client = NativeInterfaceClient(local: address, socketPath: socket, queue: queue)
|
||||
client.onPacket = { [weak self] packet in self?.receiveOutboundIPv6(packet) }
|
||||
client.onDisconnect = { [weak self] in self?.nativeInterface = nil; self?.synchronizedRoutes.removeAll() }
|
||||
do {
|
||||
try client.connect(); nativeInterface = client; synchronizeNativeRoutes(force: true)
|
||||
if nativeInterface != nil { log("native interface connected") }
|
||||
}
|
||||
catch { client.disconnect(); log("native interface unavailable: \(error.localizedDescription)") }
|
||||
}
|
||||
|
||||
private func synchronizeNativeRoutes(force: Bool = false) {
|
||||
guard let nativeInterface else { return }
|
||||
let desired = Set(router.routes().keys.filter { router.record(for: $0) != nil })
|
||||
guard force || desired != synchronizedRoutes else { return }
|
||||
do { try nativeInterface.setRoutes(desired); synchronizedRoutes = desired }
|
||||
catch { log("route synchronization failed: \(error.localizedDescription)"); nativeInterface.disconnect() }
|
||||
}
|
||||
|
||||
private func receiveOutboundIPv6(_ packet: Data) {
|
||||
guard let info = try? IPv6PacketParser.parse(packet), info.source == address,
|
||||
let remote = router.record(for: info.destination), nativeFirewall.allowOutbound(info) else { return }
|
||||
let frame = InnerFrame(kind: .ipv6Packet, payload: packet, sourceRecord: identity.record, replayID: UUID())
|
||||
_ = send(frame, to: remote)
|
||||
}
|
||||
|
||||
private func allowPing(from source: MeshAddress) -> Bool {
|
||||
let now = Date()
|
||||
if let current = pingWindows[source], now.timeIntervalSince(current.0) < 1 {
|
||||
@@ -427,6 +507,7 @@ final class MeshDaemon {
|
||||
if Int(now.timeIntervalSince1970) % 5 == 0 { broadcast(.keepalive, excluding: nil) }
|
||||
if Int(now.timeIntervalSince1970) % 10 == 0 { publishLocalState() }
|
||||
router.expire(now: now)
|
||||
attemptInterface(); synchronizeNativeRoutes()
|
||||
pingWindows = pingWindows.filter { now.timeIntervalSince($0.value.0) < 60 }
|
||||
let expiredPings = pendingPings.filter { now.timeIntervalSince($0.value.1) > 10 }
|
||||
for (id, pending) in expiredPings {
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
import Foundation
|
||||
import Darwin
|
||||
import UltraMeshCore
|
||||
import UMNSystem
|
||||
|
||||
struct NativeInterfaceStatus {
|
||||
var name: String?
|
||||
var address: MeshAddress
|
||||
var mtu = 1280
|
||||
var helperState: String
|
||||
var routeCount = 0
|
||||
}
|
||||
|
||||
final class NativeInterfaceClient {
|
||||
private let local: MeshAddress
|
||||
private let socketPath: String
|
||||
private let callbackQueue: DispatchQueue
|
||||
private var helperFD: Int32 = -1
|
||||
private var deviceFD: Int32 = -1
|
||||
private var readSource: DispatchSourceRead?
|
||||
private var helperSource: DispatchSourceRead?
|
||||
private let writeLock = NSLock()
|
||||
private(set) var status: NativeInterfaceStatus
|
||||
var onPacket: ((Data) -> Void)?
|
||||
var onDisconnect: (() -> Void)?
|
||||
|
||||
init(local: MeshAddress, socketPath: String = "/var/run/ultramesh-helper.sock", queue: DispatchQueue) {
|
||||
self.local = local; self.socketPath = socketPath; callbackQueue = queue
|
||||
status = .init(address: local, helperState: "disconnected")
|
||||
}
|
||||
|
||||
func connect() throws {
|
||||
guard helperFD < 0 else { return }
|
||||
let socket = try UnixIPC.connect(path: socketPath)
|
||||
do {
|
||||
try writeAll(Data("UMN2 \(local)\n".utf8), fd: socket)
|
||||
var receivedFD: Int32 = -1; var buffer = [UInt8](repeating: 0, count: 256)
|
||||
let count = umn_recv_fd(socket, &receivedFD, &buffer, buffer.count)
|
||||
guard count > 0, receivedFD >= 0 else { throw UMNError.invalidFrame }
|
||||
var adopted = false
|
||||
defer { if !adopted { Darwin.close(receivedFD) } }
|
||||
var readyData = Data(buffer.prefix(count))
|
||||
while !readyData.contains(10), readyData.count < buffer.count {
|
||||
var byte: UInt8 = 0
|
||||
let amount = Darwin.read(socket, &byte, 1)
|
||||
guard amount == 1 else { throw UMNError.invalidFrame }
|
||||
readyData.append(byte)
|
||||
}
|
||||
guard readyData.last == 10, let ready = String(data: readyData, encoding: .utf8) else { throw UMNError.invalidFrame }
|
||||
let fields = ready.trimmingCharacters(in: .whitespacesAndNewlines).split(separator: " ")
|
||||
guard fields.count == 3, fields[0] == "READY", fields[1].hasPrefix("utun"), fields[2] == "1280" else {
|
||||
throw UMNError.invalidFrame
|
||||
}
|
||||
helperFD = socket; deviceFD = receivedFD
|
||||
adopted = true
|
||||
status.name = String(fields[1]); status.helperState = "connected"
|
||||
let source = DispatchSource.makeReadSource(fileDescriptor: receivedFD, queue: callbackQueue)
|
||||
source.setEventHandler { [weak self] in self?.readPacket() }
|
||||
source.setCancelHandler { Darwin.close(receivedFD) }
|
||||
source.resume(); readSource = source
|
||||
let helperMonitor = DispatchSource.makeReadSource(fileDescriptor: socket, queue: callbackQueue)
|
||||
helperMonitor.setEventHandler { [weak self] in self?.helperBecameReadable() }
|
||||
helperMonitor.resume(); helperSource = helperMonitor
|
||||
} catch { UnixIPC.close(socket); throw error }
|
||||
}
|
||||
|
||||
func setRoutes(_ addresses: Set<MeshAddress>) throws {
|
||||
guard helperFD >= 0 else { throw UMNError.message("helper disconnected") }
|
||||
let line = "ROUTES" + addresses.sorted().map { " \($0)" }.joined() + "\n"
|
||||
do { try writeAll(Data(line.utf8), fd: helperFD); status.routeCount = addresses.count }
|
||||
catch { disconnect(); throw error }
|
||||
}
|
||||
|
||||
func inject(_ packet: Data) -> Bool {
|
||||
guard deviceFD >= 0 else { return false }
|
||||
var family = UInt32(AF_INET6).bigEndian
|
||||
let framed = withUnsafeBytes(of: &family) { Data($0) } + packet
|
||||
writeLock.lock(); defer { writeLock.unlock() }
|
||||
return (try? writeAll(framed, fd: deviceFD)) != nil
|
||||
}
|
||||
|
||||
func disconnect() {
|
||||
let wasConnected = helperFD >= 0
|
||||
readSource?.cancel(); readSource = nil; deviceFD = -1
|
||||
helperSource?.cancel(); helperSource = nil
|
||||
if helperFD >= 0 { UnixIPC.close(helperFD); helperFD = -1 }
|
||||
status.name = nil; status.helperState = "disconnected"; status.routeCount = 0
|
||||
if wasConnected { onDisconnect?() }
|
||||
}
|
||||
|
||||
private func readPacket() {
|
||||
guard deviceFD >= 0 else { return }
|
||||
var buffer = [UInt8](repeating: 0, count: IPv6PacketParser.maximumPacketSize + 4)
|
||||
let count = Darwin.read(deviceFD, &buffer, buffer.count)
|
||||
guard count > 4 else { if count <= 0 { disconnect() }; return }
|
||||
let family = buffer.prefix(4).reduce(UInt32(0)) { ($0 << 8) | UInt32($1) }
|
||||
guard family == UInt32(AF_INET6) else { return }
|
||||
onPacket?(Data(buffer[4..<count]))
|
||||
}
|
||||
|
||||
private func helperBecameReadable() {
|
||||
guard helperFD >= 0 else { return }
|
||||
var byte: UInt8 = 0
|
||||
_ = recv(helperFD, &byte, 1, MSG_PEEK)
|
||||
// The protocol has no unsolicited helper messages, so data or EOF means
|
||||
// the session is no longer trustworthy.
|
||||
disconnect()
|
||||
}
|
||||
|
||||
private func writeAll(_ data: Data, fd: Int32) throws {
|
||||
try data.withUnsafeBytes { bytes in
|
||||
guard let base = bytes.baseAddress else { return }
|
||||
var offset = 0
|
||||
while offset < data.count {
|
||||
let count = Darwin.write(fd, base.advanced(by: offset), data.count - offset)
|
||||
if count < 0 { if errno == EINTR { continue }; throw POSIXError(.init(rawValue: errno) ?? .EIO) }
|
||||
offset += count
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -10,6 +10,12 @@ final class PeerSession {
|
||||
var onStop: ((PeerSession) -> Void)?
|
||||
private let queue: DispatchQueue
|
||||
private var stopped = false
|
||||
private struct Pending { let data: Data; let bulk: Bool }
|
||||
private var priorityQueue: [Pending] = []
|
||||
private var bulkQueue: [Pending] = []
|
||||
private var bulkBytes = 0
|
||||
private var sending = false
|
||||
private static let maximumPendingBulkBytes = 512 * 1024
|
||||
|
||||
init(connection: NWConnection, outbound: Bool, queue: DispatchQueue) {
|
||||
self.connection = connection; self.outbound = outbound; self.queue = queue
|
||||
@@ -30,9 +36,13 @@ final class PeerSession {
|
||||
|
||||
func send(_ message: WireMessage) {
|
||||
guard !stopped, let data = try? FrameCodec.encode(WireEnvelope(message: message)) else { return }
|
||||
connection.send(content: data, completion: .contentProcessed { [weak self] error in
|
||||
if error != nil { self?.stop() }
|
||||
})
|
||||
let bulk: Bool
|
||||
if case .packet(let packet) = message { bulk = packet.trafficClass == .nativeIPv6 } else { bulk = false }
|
||||
if bulk {
|
||||
guard bulkBytes + data.count <= Self.maximumPendingBulkBytes else { return }
|
||||
bulkBytes += data.count; bulkQueue.append(.init(data: data, bulk: true))
|
||||
} else { priorityQueue.append(.init(data: data, bulk: false)) }
|
||||
pumpWrites()
|
||||
}
|
||||
|
||||
func stop() {
|
||||
@@ -54,6 +64,22 @@ final class PeerSession {
|
||||
}
|
||||
}
|
||||
|
||||
private func pumpWrites() {
|
||||
guard !stopped, !sending else { return }
|
||||
let pending: Pending
|
||||
if !priorityQueue.isEmpty { pending = priorityQueue.removeFirst() }
|
||||
else if !bulkQueue.isEmpty { pending = bulkQueue.removeFirst(); bulkBytes -= pending.data.count }
|
||||
else { return }
|
||||
sending = true
|
||||
connection.send(content: pending.data, completion: .contentProcessed { [weak self] error in
|
||||
guard let self else { return }
|
||||
self.queue.async {
|
||||
self.sending = false
|
||||
if error != nil { self.stop() } else { self.pumpWrites() }
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
private func receiveExactly(_ count: Int, accumulated: Data, completion: @escaping (Data?) -> Void) {
|
||||
if accumulated.count == count { completion(accumulated); return }
|
||||
connection.receive(minimumIncompleteLength: 1, maximumLength: count - accumulated.count) { [weak self] data, _, complete, error in
|
||||
|
||||
@@ -2,6 +2,7 @@ import Foundation
|
||||
import UltraMeshCore
|
||||
|
||||
do {
|
||||
signal(SIGPIPE, SIG_IGN)
|
||||
let arguments = Array(CommandLine.arguments.dropFirst())
|
||||
var stateURL = FileManager.default.homeDirectoryForCurrentUser
|
||||
.appendingPathComponent("Library/Application Support/UltraMesh")
|
||||
|
||||
Reference in New Issue
Block a user