punchnet-macos/Tun/Punchnet/UDPHole/SDLUDPHole.swift
2026-05-06 16:05:03 +08:00

206 lines
6.1 KiB
Swift
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

//
// SDLanServer.swift
// Tun
//
// Created by on 2024/1/31.
//
import Foundation
import NIOCore
import NIOPosix
import SwiftProtobuf
enum SDLUDPHoleError: Error {
case invalidLocalAddress
case closed
case errorCaught
case sendFaied(Error)
}
actor SDLUDPHole {
enum State {
case idle
case running
case stopped
}
private var state: State = .idle
private let udpHoleHandler: SDLUDPHoleHandler
init() throws {
self.udpHoleHandler = try SDLUDPHoleHandler()
}
func start() throws -> SocketAddress {
let localAddress = try self.udpHoleHandler.start()
self.state = .running
return localAddress
}
func messageStream() -> AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error> {
return self.udpHoleHandler.messageStream
}
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) {
guard self.state == .running else {
return
}
self.udpHoleHandler.send(type: type, data: data, remoteAddress: remoteAddress)
}
func stop() {
guard self.state != .stopped else {
return
}
self.state = .stopped
self.udpHoleHandler.stop()
}
deinit {
SDLLogger.log("[SDLUDPHoleActor] deinit", for: .debug)
}
}
// sn-server
private final class SDLUDPHoleHandler: ChannelInboundHandler {
typealias InboundIn = AddressedEnvelope<ByteBuffer>
private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1)
private var channel: Channel?
private let locker = NSLock()
public let messageStream: AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error>
private let messageContinuation: AsyncThrowingStream<(SocketAddress, SDLHoleMessage), Error>.Continuation
private var isMessageContinuationFinished: Bool = false
private let counterActor: SDLUDPCounter
//
init() throws {
let (stream, continuation) = AsyncThrowingStream.makeStream(of: (SocketAddress, SDLHoleMessage).self, bufferingPolicy: .bufferingNewest(2048))
self.messageStream = stream
self.messageContinuation = continuation
self.counterActor = SDLUDPCounter()
}
func start() throws -> SocketAddress {
let bootstrap = DatagramBootstrap(group: group)
.channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1)
.channelInitializer { channel in
channel.pipeline.addHandler(self)
}
// IPv4IPv4
let channel = try bootstrap.bind(host: "0.0.0.0", port: 0).wait()
guard let localAddress = channel.localAddress else {
throw SDLUDPHoleError.invalidLocalAddress
}
self.channel = channel
let counterActor = self.counterActor
Task {
await counterActor.start()
}
return localAddress
}
// --MARK: ChannelInboundHandler delegate
func channelRead(context: ChannelHandlerContext, data: NIOAny) {
let envelope = unwrapInboundIn(data)
var buffer = envelope.data
let remoteAddress = envelope.remoteAddress
let bytesCount = buffer.readableBytes
let counterActor = self.counterActor
Task {
await counterActor.increment(direction: .inbound, from: remoteAddress.description, bytes: bytesCount)
}
do {
if let message = try SDLHoleMessage.decode(buffer: &buffer) {
self.messageContinuation.yield((remoteAddress, message))
} else {
SDLLogger.log("[SDLUDPHole] decode message, get null", for: .debug)
}
} catch let err {
SDLLogger.log("[SDLUDPHole] decode message, get error: \(err)", for: .debug)
}
}
func channelInactive(context: ChannelHandlerContext) {
self.finishMessageContinuationIfNeed(throwing: .closed)
}
func errorCaught(context: ChannelHandlerContext, error: any Error) {
context.close(promise: nil)
self.finishMessageContinuationIfNeed(throwing: .errorCaught)
}
// MARK:
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) {
guard let channel = self.channel else {
return
}
let counterActor = self.counterActor
Task {
await counterActor.increment(direction: .outbound, from: remoteAddress.description, bytes: data.count)
}
var buffer = channel.allocator.buffer(capacity: data.count + 1)
buffer.writeBytes([type.rawValue])
buffer.writeBytes(data)
let envelope = AddressedEnvelope<ByteBuffer>(remoteAddress: remoteAddress, data: buffer)
let promise = channel.eventLoop.makePromise(of: Void.self)
channel.eventLoop.execute {
channel.writeAndFlush(envelope, promise: promise)
}
promise.futureResult.whenFailure { [weak self] err in
self?.finishMessageContinuationIfNeed(throwing: .sendFaied(err))
}
}
func stop() {
self.finishMessageContinuationIfNeed(throwing: nil)
try? self.channel?.close().wait()
self.channel = nil
try? self.group.syncShutdownGracefully()
let counterActor = self.counterActor
Task {
await counterActor.stop()
}
SDLLogger.log("[SDLUDPHole] stopped", for: .debug)
}
private func finishMessageContinuationIfNeed(throwing error: SDLUDPHoleError?) {
locker.lock()
defer {
locker.unlock()
}
guard !self.isMessageContinuationFinished else {
return
}
self.isMessageContinuationFinished = true
self.messageContinuation.finish(throwing: error)
}
deinit {
SDLLogger.log("[SDLUDPHole] deinit", for: .debug)
}
}