punchnet-macos/Tun/Punchnet/UDPHole/SDLUDPHole.swift

185 lines
5.5 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() async 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() async {
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
//
init() throws {
let (stream, continuation) = AsyncThrowingStream.makeStream(of: (SocketAddress, SDLHoleMessage).self, bufferingPolicy: .bufferingNewest(2048))
self.messageStream = stream
self.messageContinuation = continuation
}
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
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
SDLLogger.log("[SDLUDPHole] read data: \(bytesCount), from: \(remoteAddress.description)", for: .debug)
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
}
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)
let channel = self.channel
self.channel = nil
try? channel?.close().wait()
try? self.group.syncShutdownGracefully()
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)
}
}