fix udpHole

This commit is contained in:
anlicheng 2026-05-04 20:15:08 +08:00
parent 783c1329e1
commit 55701085e3
3 changed files with 106 additions and 124 deletions

View File

@ -53,8 +53,6 @@ actor SDLContextActor {
private var udpHoleLocalAddress: SocketAddress?
private var udpHoleV6: SDLUDPHoleV6?
private var udpHoleV6Workers: [Task<Void, Never>]?
private var udpHoleV6LocalAddress: SocketAddress?
// dnsclient
private var dnsClient: DNSCloudClient?
@ -149,9 +147,8 @@ actor SDLContextActor {
}
// await self.supervisor.addWorker(name: "udpHoleV6") {
// let udpHoleV6 = try await self.startUDPHoleV6()
// SDLLogger.log("[SDLContext] udp v6 running!!!!")
// try await udpHoleV6.waitClose()
// try await self.startUDPHoleV6()
// SDLLogger.log("[SDLContext] udp v6 closed!!!!")
// }
}
@ -357,6 +354,7 @@ actor SDLContextActor {
self.udpHoleLocalAddress = localAddress
defer {
self.udpHole?.stop()
self.udpHole = nil
self.udpHoleLocalAddress = nil
}
@ -371,8 +369,17 @@ actor SDLContextActor {
try Task.checkCancellation()
switch message.inboundMessage {
case .control(let controlMessage):
await self.handleHoleControlMessage(controlMessage, localAddress: localAddress, remoteAddress: remoteAddress, source: .v4)
case .control(let message):
switch message {
case .stunReply(_):
SDLLogger.log("[SDLContext] get a stunReply", for: .debug)
case .stunProbeReply(let probeReply):
await self.proberActor.handleProbeReply(localAddress: localAddress, reply: probeReply)
case .register(let register):
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v4)
case .registerAck(let registerAck):
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6)
}
case .data(let data):
try? await self.handleHoleData(data: data)
}
@ -399,26 +406,64 @@ actor SDLContextActor {
}
}
private func startUDPHoleV6() async throws -> SDLUDPHoleV6 {
self.udpHoleV6Workers?.forEach {$0.cancel()}
self.udpHoleV6Workers = nil
private func startUDPHoleV6() async throws {
// udp
let udpHoleV6 = try SDLUDPHoleV6()
let localAddress = try udpHoleV6.start()
SDLLogger.log("[SDLContext] udpHoleV6 started, on address: \(localAddress)")
self.udpHoleV6 = udpHoleV6
//
let messageStream = udpHoleV6.messageStream
let messageTask = Task.detached {
await self.consumeUDPHoleMessages(stream: messageStream, localAddress: localAddress, source: .v6)
if let localAddress {
SDLLogger.log("[SDLContext] udpHoleV6 started, on address: \(localAddress)")
} else {
SDLLogger.log("[SDLContext] udpHoleV6 started, no local address")
}
self.udpHoleV6 = udpHoleV6
self.udpHoleV6LocalAddress = localAddress
self.udpHoleV6Workers = [messageTask]
defer {
self.udpHoleV6?.stop()
self.udpHoleV6 = nil
}
return udpHoleV6
try await withThrowingTaskGroup { group in
defer {
group.cancelAll()
}
//
group.addTask {
for await (remoteAddress, message) in udpHoleV6.messageStream {
try Task.checkCancellation()
switch message.inboundMessage {
case .control(let message):
switch message {
case .register(let register):
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: .v6)
case .registerAck(let registerAck):
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: .v6)
default:
()
}
case .data(let data):
try? await self.handleHoleData(data: data)
}
}
}
group.addTask {
for await event in udpHoleV6.eventStream {
try Task.checkCancellation()
switch event {
case .ready:
SDLLogger.log("[SDLContext] udpHoleV6 ready")
case .closed, .errorCaught:
throw SDLContextError.udpHoleClosed
}
}
}
try await group.next()
}
}
// context
@ -435,14 +480,12 @@ actor SDLContextActor {
self.flowSessionManager.clear()
self.udpHole?.stop()
self.udpHole = nil
self.udpHoleLocalAddress = nil
self.udpHoleV6Workers?.forEach { $0.cancel() }
self.udpHoleV6Workers = nil
self.udpHoleV6?.stop()
self.udpHoleV6 = nil
self.udpHoleV6LocalAddress = nil
self.quicClient?.stop()
self.quicClient = nil
@ -682,7 +725,6 @@ actor SDLContextActor {
self.udpHole = nil
self.udpHoleLocalAddress = nil
self.udpHoleV6 = nil
self.udpHoleV6LocalAddress = nil
self.dnsClient = nil
}
}
@ -800,40 +842,6 @@ extension SDLContextActor {
// Hole
extension SDLContextActor {
private func consumeUDPHoleMessages(stream: AsyncStream<(SocketAddress, SDLHoleMessage)>, localAddress: SocketAddress, source: UDPHoleKind) async {
for await (remoteAddress, message) in stream {
if Task.isCancelled {
break
}
switch message.inboundMessage {
case .control(let controlMessage):
await self.handleHoleControlMessage(controlMessage, localAddress: localAddress, remoteAddress: remoteAddress, source: source)
case .data(let data):
try? await self.handleHoleData(data: data)
}
}
}
private func handleHoleControlMessage(_ message: SDLHoleControlMessage, localAddress: SocketAddress, remoteAddress: SocketAddress, source: UDPHoleKind) async {
switch message {
case .stunReply(_):
guard source == .v4 else {
return
}
SDLLogger.log("[SDLContext] get a stunReply", for: .debug)
case .stunProbeReply(let probeReply):
guard source == .v4 else {
return
}
await self.proberActor.handleProbeReply(localAddress: localAddress, reply: probeReply)
case .register(let register):
try? await self.handleRegister(remoteAddress: remoteAddress, register: register, source: source)
case .registerAck(let registerAck):
await self.handleRegisterAck(remoteAddress: remoteAddress, registerAck: registerAck, source: source)
}
}
private func handleRegister(remoteAddress: SocketAddress, register: SDLRegister, source: UDPHoleKind) async throws {
let networkAddr = config.networkAddress
SDLLogger.log("[SDLContext] register packet: \(register), network_address: \(networkAddr)")

View File

@ -24,6 +24,8 @@ final class SDLUDPHole: ChannelInboundHandler {
case errorCaught
}
private var isStopped: Bool = false
private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1)
private var channel: Channel?
@ -113,8 +115,14 @@ final class SDLUDPHole: ChannelInboundHandler {
}
}
deinit {
SDLLogger.log("[SDLUDPHole] deinit", for: .debug)
func stop() {
guard !self.isStopped else {
return
}
self.isStopped = true
SDLLogger.log("[SDLUDPHole] stop", for: .debug)
self.messageContinuation.finish()
self.eventContinuation.finish()
self.channel = nil

View File

@ -14,30 +14,37 @@ import SwiftProtobuf
final class SDLUDPHoleV6: ChannelInboundHandler {
typealias InboundIn = AddressedEnvelope<ByteBuffer>
private enum State: Equatable {
case idle
//
enum HoleEvent {
case ready
case stopping
case stopped
case closed
case errorCaught
}
private var isStopped: Bool = false
private let group = MultiThreadedEventLoopGroup(numberOfThreads: 1)
private var channel: Channel?
private var closeFuture: EventLoopFuture<Void>?
private var state: State = .idle
private var didFinishMessageStream: Bool = false
public let messageStream: AsyncStream<(SocketAddress, SDLHoleMessage)>
private let messageContinuation: AsyncStream<(SocketAddress, SDLHoleMessage)>.Continuation
//
public let eventStream: AsyncStream<HoleEvent>
private let eventContinuation: AsyncStream<HoleEvent>.Continuation
//
init() throws {
let (stream, continuation) = AsyncStream.makeStream(of: (SocketAddress, SDLHoleMessage).self, bufferingPolicy: .bufferingNewest(2048))
self.messageStream = stream
self.messageContinuation = continuation
let eventPair = AsyncStream.makeStream(of: HoleEvent.self)
self.eventStream = eventPair.stream
self.eventContinuation = eventPair.continuation
}
func start() throws -> SocketAddress {
func start() throws -> SocketAddress? {
let bootstrap = DatagramBootstrap(group: group)
.channelOption(ChannelOptions.socketOption(.so_reuseaddr), value: 1)
.channelInitializer { channel in
@ -47,51 +54,13 @@ final class SDLUDPHoleV6: ChannelInboundHandler {
// IPv6IPv6
let channel = try bootstrap.bind(host: "::", port: 0).wait()
self.channel = channel
self.closeFuture = channel.closeFuture
self.state = .ready
precondition(channel.localAddress != nil, "UDP v6 channel has no localAddress after bind")
return channel.localAddress!
}
func waitClose() async throws {
switch self.state {
case .idle:
SDLLogger.log("[SDLUDPHoleV6] waitClose11", for: .debug)
return
case .ready, .stopping, .stopped:
guard let closeFuture = self.closeFuture else {
SDLLogger.log("[SDLUDPHoleV6] waitClose22", for: .debug)
return
}
try await closeFuture.get()
SDLLogger.log("[SDLUDPHoleV6] waitClose33", for: .debug)
}
}
func stop() {
switch self.state {
case .stopping, .stopped:
return
case .idle:
self.state = .stopped
self.finishMessageStream()
return
case .ready:
self.state = .stopping
}
self.finishMessageStream()
self.channel?.close(promise: nil)
return channel.localAddress
}
// --MARK: ChannelInboundHandler delegate
func channelRead(context: ChannelHandlerContext, data: NIOAny) {
guard case .ready = self.state else {
return
}
let envelope = unwrapInboundIn(data)
var buffer = envelope.data
@ -113,23 +82,19 @@ final class SDLUDPHoleV6: ChannelInboundHandler {
}
func channelInactive(context: ChannelHandlerContext) {
self.finishMessageStream()
self.channel = nil
self.state = .stopped
SDLLogger.log("[SDLUDPHoleV6] channelInactive", for: .debug)
self.eventContinuation.yield(.closed)
}
func errorCaught(context: ChannelHandlerContext, error: any Error) {
SDLLogger.log("[SDLUDPHoleV6] channel error: \(error)", for: .debug)
self.finishMessageStream()
if self.state != .stopped {
self.state = .stopping
}
context.close(promise: nil)
self.eventContinuation.yield(.errorCaught)
}
// MARK:
func send(type: SDLPacketType, data: Data, remoteAddress: SocketAddress) {
guard case .ready = self.state, let channel = self.channel else {
guard let channel = self.channel else {
return
}
@ -143,17 +108,18 @@ final class SDLUDPHoleV6: ChannelInboundHandler {
}
}
private func finishMessageStream() {
guard !self.didFinishMessageStream else {
func stop() {
guard !self.isStopped else {
return
}
self.didFinishMessageStream = true
self.messageContinuation.finish()
}
self.isStopped = true
SDLLogger.log("[SDLUDPHoleV6] stop", for: .debug)
self.messageContinuation.finish()
self.eventContinuation.finish()
self.channel = nil
deinit {
self.stop()
try? self.group.syncShutdownGracefully()
}