Files
umn/Sources/umnd/MeshDaemon.swift
T
2026-08-31 15:55:09 -05:00

453 lines
24 KiB
Swift

import Foundation
import Network
import UltraMeshCore
private final class StreamState {
struct Pending { var frame: InnerFrame; var lastSent: Date; var attempts: Int }
let id: UInt64
let remote: NodeRecord
weak var client: IPCClient?
var nextSequence: UInt64 = 1
var expectedSequence: UInt64 = 1
var pending: [UInt64: Pending] = [:]
var queued: [Data] = []
var closing = false
var transferred = 0
var lastProgress = Date()
init(id: UInt64, remote: NodeRecord, client: IPCClient) {
self.id = id; self.remote = remote; self.client = client
}
}
final class MeshDaemon {
private let queue = DispatchQueue(label: "com.ultramesh.daemon")
private let identity: NodeIdentity
private let router: LinkStateRouter
private let firewall: MeshFirewall
private let configURL: URL
private var config: DaemonConfig
private let ipc: IPCServer
private var listener: NWListener?
private var browser: NWBrowser?
private var pendingEndpoints = Set<String>()
private var peers: [MeshAddress: PeerSession] = [:]
private var allSessions: [ObjectIdentifier: PeerSession] = [:]
private var lsaSequence: UInt64 = 0
private var seenPackets = Set<UUID>()
private var seenOrder: [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?
init(baseURL: URL = FileManager.default.homeDirectoryForCurrentUser
.appendingPathComponent("Library/Application Support/UltraMesh"), socketPath: String? = nil) throws {
try FileManager.default.createDirectory(at: baseURL, withIntermediateDirectories: true)
identity = try NodeIdentity.loadOrCreate(at: baseURL.appendingPathComponent("identity.plist"))
router = LinkStateRouter(local: identity.record.address)
configURL = baseURL.appendingPathComponent("config.json")
config = DaemonConfig.load(from: configURL)
firewall = MeshFirewall(rules: config.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)
config.lsaSequence = lsaSequence
try config.save(to: configURL)
}
var address: MeshAddress { identity.record.address }
func start() throws {
try startNetwork()
ipc.onMessage = { [weak self] client, message in self?.queue.async { self?.handleIPC(client, message) } }
try ipc.start()
publishLocalState()
let timer = DispatchSource.makeTimerSource(queue: queue)
timer.schedule(deadline: .now() + 1, repeating: 1)
timer.setEventHandler { [weak self] in self?.tick() }
timer.resume(); self.timer = timer
log("node \(address) started")
}
func stop() {
queue.sync {
timer?.cancel(); browser?.cancel(); listener?.cancel()
allSessions.values.forEach { $0.stop() }; ipc.stop()
}
}
private func startNetwork() throws {
let parameters = NWParameters.tcp
parameters.includePeerToPeer = true
let listener = try NWListener(using: parameters)
listener.service = .init(name: "umn-\(address.description.suffix(8))", type: "_umn._tcp")
listener.newConnectionHandler = { [weak self] connection in self?.queue.async { self?.addSession(connection, outbound: false) } }
listener.stateUpdateHandler = { [weak self] state in
if case .failed(let error) = state { self?.log("listener failed: \(error)") }
}
listener.start(queue: queue); self.listener = listener
let browseParameters = NWParameters.tcp
browseParameters.includePeerToPeer = true
let browser = NWBrowser(for: .bonjour(type: "_umn._tcp", domain: nil), using: browseParameters)
browser.browseResultsChangedHandler = { [weak self] results, _ in
guard let self else { return }
self.queue.async {
for result in results {
let key = String(describing: result.endpoint)
guard !self.pendingEndpoints.contains(key) else { continue }
self.pendingEndpoints.insert(key)
let connection = NWConnection(to: result.endpoint, using: browseParameters)
self.addSession(connection, outbound: true)
}
}
}
browser.stateUpdateHandler = { [weak self] state in
if case .failed(let error) = state { self?.log("browser failed: \(error)") }
}
browser.start(queue: queue); self.browser = browser
}
private func addSession(_ connection: NWConnection, outbound: Bool) {
let session = PeerSession(connection: connection, outbound: outbound, queue: queue)
allSessions[ObjectIdentifier(session)] = session
session.onMessage = { [weak self] peer, message in self?.queue.async { self?.handle(peer, message) } }
session.onStop = { [weak self] peer in self?.queue.async { self?.remove(peer) } }
session.start(local: identity.record)
}
private func remove(_ session: PeerSession) {
allSessions.removeValue(forKey: ObjectIdentifier(session))
pendingEndpoints.remove(String(describing: session.connection.endpoint))
if let remote = session.remote, peers[remote.address] === session {
peers.removeValue(forKey: remote.address); log("peer \(remote.address) disconnected"); publishLocalState()
let endpoint = session.connection.endpoint
queue.asyncAfter(deadline: .now() + 1) { [weak self] in
guard let self, self.peers[remote.address] == nil else { return }
let parameters = NWParameters.tcp; parameters.includePeerToPeer = true
self.addSession(NWConnection(to: endpoint, using: parameters), outbound: true)
}
}
}
private func handle(_ session: PeerSession, _ message: WireMessage) {
switch message {
case .hello(let record):
guard record.validate(), record.address != address else { session.stop(); return }
guard peers[record.address] != nil || peers.count < 31 else { session.stop(); return }
session.remote = record
let shouldBeOutbound = address < record.address
if session.outbound != shouldBeOutbound, let existing = peers[record.address], existing.outbound == shouldBeOutbound {
session.stop(); return
}
if let existing = peers[record.address], existing !== session {
if existing.outbound == shouldBeOutbound { session.stop(); return }
existing.stop()
}
peers[record.address] = session
log("peer \(record.address) connected over \(session.outbound ? "outbound" : "inbound") path")
for state in currentLinkStates() { session.send(.linkState(state)) }
publishLocalState()
case .linkState(let state):
guard session.remote != nil, router.ingest(state) else { return }
broadcast(.linkState(state), excluding: session)
case .packet(var packet):
guard session.remote != nil, markSeen(packet.id) else { return }
if packet.destination == address { receiveLocal(packet) }
else if packet.hopLimit > 1 {
packet.hopLimit -= 1; forward(packet)
}
case .keepalive: break
}
}
private func publishLocalState() {
lsaSequence &+= 1
guard let state = try? identity.makeLinkState(sequence: lsaSequence, neighbors: Array(peers.keys)) else { return }
_ = router.ingest(state); broadcast(.linkState(state), excluding: nil)
}
private func currentLinkStates() -> [LinkState] {
router.allStates()
}
private func broadcast(_ message: WireMessage, excluding: PeerSession?) {
for peer in peers.values where peer !== excluding { peer.send(message) }
}
private func forward(_ packet: RoutedPacket) {
guard let route = router.routes()[packet.destination], let peer = peers[route.nextHop] else { return }
peer.send(.packet(packet))
}
@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)
guard router.routes()[remote.address] != nil || peers[remote.address] != nil else { return false }
forward(packet); return true
}
private func receiveLocal(_ packet: RoutedPacket) {
guard let inner = try? identity.open(packet.sealed, expectedSource: packet.source) else { return }
switch inner.kind {
case .pingRequest:
guard allowPing(from: packet.source) else { return }
let reply = InnerFrame(kind: .pingReply, streamID: inner.streamID, payload: inner.payload, sourceRecord: identity.record)
_ = send(reply, to: inner.sourceRecord)
case .pingReply:
guard let idString = String(data: inner.payload, encoding: .utf8), let id = UUID(uuidString: idString),
let pending = pendingPings.removeValue(forKey: id) else { return }
var response = IPCMessage(operation: .ping, requestID: id); response.ok = true
response.message = String(format: "reply from %@: %.1f ms", packet.source.description,
Date().timeIntervalSince(pending.1) * 1000)
pending.0.send(response)
case .text:
guard let port = inner.port, firewall.allows(port: port, source: packet.source),
let client = textBindings[port] else { return }
var event = IPCMessage(operation: .textSend); event.ok = true
event.source = packet.source.description; event.port = port; event.data = inner.payload
client.send(event)
case .streamOpen: receiveStreamOpen(inner, from: packet.source)
case .streamData: receiveStreamData(inner)
case .streamAck: receiveStreamAck(inner)
case .streamClose, .streamReset: receiveStreamClose(inner)
case .error: receiveStreamClose(inner)
}
}
private func receiveStreamOpen(_ inner: InnerFrame, from source: MeshAddress) {
if let existing = streams[inner.streamID] {
sendAck(for: existing, sequence: 0); return
}
guard streams.count < 32,
streams.values.filter({ $0.remote.address == source }).count < 8,
let port = inner.port,
firewall.allows(port: port, source: source), let client = streamBindings[port] else {
let reset = InnerFrame(kind: .streamReset, streamID: inner.streamID,
payload: Data("denied or unbound".utf8), sourceRecord: identity.record)
_ = send(reset, to: inner.sourceRecord); return
}
let state = StreamState(id: inner.streamID, remote: inner.sourceRecord, client: client)
streams[inner.streamID] = state
var event = IPCMessage(operation: .openStream); event.ok = true; event.streamID = inner.streamID
event.source = source.description; event.port = port; client.send(event)
sendAck(for: state, sequence: 0)
}
private func receiveStreamData(_ inner: InnerFrame) {
guard let state = streams[inner.streamID] else { return }
if inner.sequence == state.expectedSequence {
guard state.transferred + inner.payload.count <= 1_048_576 else { reset(state, reason: "1 MiB limit exceeded"); return }
state.transferred += inner.payload.count; state.expectedSequence &+= 1; state.lastProgress = Date()
var event = IPCMessage(operation: .streamData); event.streamID = state.id; event.data = inner.payload; event.ok = true
state.client?.send(event)
}
sendAck(for: state, sequence: state.expectedSequence - 1)
}
private func receiveStreamAck(_ inner: InnerFrame) {
guard let state = streams[inner.streamID] else { return }
state.pending = state.pending.filter { $0.key > inner.acknowledgment }
state.lastProgress = Date()
pump(state)
finishIfDrained(state)
}
private func receiveStreamClose(_ inner: InnerFrame) {
guard let state = streams.removeValue(forKey: inner.streamID) else { return }
var event = IPCMessage(operation: .streamClose); event.streamID = state.id; event.data = inner.payload
event.ok = inner.kind == .streamClose; state.client?.send(event)
}
private func sendAck(for state: StreamState, sequence: UInt64) {
let ack = InnerFrame(kind: .streamAck, streamID: state.id, acknowledgment: sequence, sourceRecord: identity.record)
_ = send(ack, to: state.remote)
}
private func reset(_ state: StreamState, reason: String) {
let frame = InnerFrame(kind: .streamReset, streamID: state.id, payload: Data(reason.utf8), sourceRecord: identity.record)
_ = send(frame, to: state.remote); streams.removeValue(forKey: state.id)
var event = IPCMessage(operation: .streamClose); event.streamID = state.id; event.ok = false; event.message = reason
state.client?.send(event)
}
private func sendReliable(_ frame: InnerFrame, state: StreamState) -> Bool {
guard send(frame, to: state.remote) else { return false }
state.pending[frame.sequence] = .init(frame: frame, lastSent: Date(), attempts: 0)
return true
}
private func pump(_ state: StreamState) {
var inFlightBytes = state.pending.values.reduce(0) { $0 + $1.frame.payload.count }
while let data = state.queued.first, inFlightBytes + data.count <= 65_536 {
state.queued.removeFirst()
let frame = InnerFrame(kind: .streamData, streamID: state.id, sequence: state.nextSequence,
payload: data, sourceRecord: identity.record)
state.nextSequence &+= 1
guard sendReliable(frame, state: state) else { state.queued.insert(data, at: 0); return }
inFlightBytes += data.count
}
}
private func finishIfDrained(_ state: StreamState) {
guard state.closing, state.queued.isEmpty, state.pending.isEmpty else { return }
let frame = InnerFrame(kind: .streamClose, streamID: state.id, sequence: state.nextSequence, sourceRecord: identity.record)
_ = send(frame, to: state.remote); streams.removeValue(forKey: state.id)
}
private func handleIPC(_ client: IPCClient, _ message: IPCMessage) {
func reply(_ ok: Bool, _ text: String? = nil, values: [String]? = nil, streamID: UInt64? = nil) {
var response = IPCMessage(operation: message.operation, requestID: message.requestID)
response.ok = ok; response.message = text; response.values = values; response.streamID = streamID
client.send(response)
}
switch message.operation {
case .status:
reply(true, "umnd running as \(address); \(peers.count) direct peer(s), \(router.routes().count) route(s)")
case .address: reply(true, address.description)
case .peers: reply(true, values: peers.keys.sorted().map(\.description))
case .routes:
let values = router.routes().values.sorted { $0.destination < $1.destination }
.map { "\($0.destination) via \($0.nextHop) (\($0.hopCount) hop\($0.hopCount == 1 ? "" : "s"))" }
reply(true, values: values)
case .aliasList:
reply(true, values: config.aliases.sorted { $0.key < $1.key }.map { "\($0.key).mesh \($0.value)" })
case .aliasSet:
guard let name = message.name?.lowercased(), validAlias(name), let target = message.target,
let value = try? MeshAddress(target) else { reply(false, "invalid alias or address"); return }
config.aliases[name] = value; saveConfig(); reply(true, "\(name).mesh -> \(value)")
case .aliasRemove:
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")" }
reply(true, values: values)
case .firewallAllow, .firewallRevoke:
guard let port = message.port, let sourceText = message.source else { reply(false, "missing port or source"); 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) }
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 }
let id = message.requestID
let inner = InnerFrame(kind: .pingRequest, streamID: UInt64.random(in: 1...UInt64.max),
payload: Data(id.uuidString.utf8), sourceRecord: identity.record)
guard send(inner, to: remote) else { reply(false, "no route to destination"); return }
pendingPings[id] = (client, Date())
case .textSend:
guard let target = message.target, let remote = record(for: target), let port = message.port,
let data = message.data, data.count <= 4096, String(data: data, encoding: .utf8) != nil else {
reply(false, "invalid destination, port, or UTF-8 message (maximum 4 KiB)"); return
}
let inner = InnerFrame(kind: .text, port: port, payload: data, sourceRecord: identity.record)
let sent = send(inner, to: remote)
reply(sent, sent ? "sent" : "no route to destination")
case .bindText:
guard let port = message.port, textBindings[port] == nil else { reply(false, "port is already bound"); return }
textBindings[port] = client; bindCleanup(client: client); reply(true, "listening on mesh port \(port)")
case .bindStream:
guard let port = message.port, streamBindings[port] == nil else { reply(false, "port is already bound"); return }
streamBindings[port] = client; bindCleanup(client: client); reply(true, "bound mesh port \(port)")
case .openStream:
guard streams.count < 32, let target = message.target, let remote = record(for: target), let port = message.port else {
reply(false, "no route to destination"); return
}
var id = UInt64.random(in: 1...UInt64.max); while streams[id] != nil { id = UInt64.random(in: 1...UInt64.max) }
let state = StreamState(id: id, remote: remote, client: client); streams[id] = state
let open = InnerFrame(kind: .streamOpen, streamID: id, port: port, sequence: 0, sourceRecord: identity.record)
guard sendReliable(open, state: state) else { streams.removeValue(forKey: id); reply(false, "no route"); return }
bindCleanup(client: client); reply(true, "opening", streamID: id)
case .streamData:
guard let id = message.streamID, let state = streams[id], state.client === client, let data = message.data,
state.transferred + data.count <= 1_048_576 else { reply(false, "unknown stream or 1 MiB limit exceeded"); return }
state.transferred += data.count
state.queued.append(data); pump(state)
case .streamClose:
guard let id = message.streamID, let state = streams[id], state.client === client else { return }
state.closing = true; finishIfDrained(state)
}
}
private func bindCleanup(client: IPCClient) {
client.onClose = { [weak self, weak client] _ in
guard let self, let client else { return }
self.queue.async {
self.textBindings = self.textBindings.filter { $0.value !== client }
self.streamBindings = self.streamBindings.filter { $0.value !== client }
let ids = self.streams.filter { $0.value.client === client }.map(\.key)
for id in ids { if let state = self.streams[id] { self.reset(state, reason: "local service disconnected") } }
}
}
}
private func record(for text: String) -> NodeRecord? {
guard let target = resolve(text) else { return nil }
if let peer = peers[target]?.remote { return peer }
return router.record(for: target)
}
private func resolve(_ text: String) -> MeshAddress? {
if let value = try? MeshAddress(text) { return value }
let name = text.lowercased().hasSuffix(".mesh") ? String(text.lowercased().dropLast(5)) : text.lowercased()
return config.aliases[name]
}
private func validAlias(_ value: String) -> Bool {
!value.isEmpty && value.count <= 63 && value.allSatisfy { $0.isLetter || $0.isNumber || $0 == "-" }
}
private func saveConfig() { try? config.save(to: configURL) }
private func markSeen(_ id: UUID) -> Bool {
guard !seenPackets.contains(id) else { return false }
seenPackets.insert(id); seenOrder.append(id)
if seenOrder.count > 4096 { seenPackets.remove(seenOrder.removeFirst()) }
return true
}
private func allowPing(from source: MeshAddress) -> Bool {
let now = Date()
if let current = pingWindows[source], now.timeIntervalSince(current.0) < 1 {
guard current.1 < 10 else { return false }
pingWindows[source] = (current.0, current.1 + 1)
} else {
guard pingWindows[source] != nil || pingWindows.count < 128 else { return false }
pingWindows[source] = (now, 1)
}
return true
}
private func tick() {
let now = Date()
if Int(now.timeIntervalSince1970) % 5 == 0 { broadcast(.keepalive, excluding: nil) }
if Int(now.timeIntervalSince1970) % 10 == 0 { publishLocalState() }
router.expire(now: now)
pingWindows = pingWindows.filter { now.timeIntervalSince($0.value.0) < 60 }
let expiredPings = pendingPings.filter { now.timeIntervalSince($0.value.1) > 10 }
for (id, pending) in expiredPings {
var response = IPCMessage(operation: .ping, requestID: id); response.ok = false; response.message = "ping timed out"
pending.0.send(response); pendingPings.removeValue(forKey: id)
}
for state in Array(streams.values) {
if now.timeIntervalSince(state.lastProgress) > 60 { reset(state, reason: "stream timed out"); continue }
for (sequence, pending) in Array(state.pending) {
let delay = min(5.0, 0.5 * pow(2.0, Double(min(pending.attempts, 4))))
if now.timeIntervalSince(pending.lastSent) >= delay {
_ = send(pending.frame, to: state.remote)
state.pending[sequence] = .init(frame: pending.frame, lastSent: now, attempts: pending.attempts + 1)
}
}
}
}
private func log(_ message: String) {
let line = "[\(ISO8601DateFormatter().string(from: Date()))] \(message)\n"
FileHandle.standardError.write(Data(line.utf8))
}
}