diff --git a/Tun/DNS/DNSHelper.swift b/Tun/DNS/DNSHelper.swift index b336a29..efd0f56 100644 --- a/Tun/DNS/DNSHelper.swift +++ b/Tun/DNS/DNSHelper.swift @@ -12,7 +12,7 @@ struct DNSHelper { static let dnsDestIpAddr: UInt32 = 1684300900 // 判断是否是dns请求的数据包 - static func isDnsRequestPacket(ipPacket: IPPacket) -> Bool { + static func isDnsRequestPacket(ipPacket: IPPacketView) -> Bool { return ipPacket.header.destination == dnsDestIpAddr } diff --git a/Tun/DNS/DNSParser.swift b/Tun/DNS/DNSParser.swift index 9c7d9fa..d467a38 100644 --- a/Tun/DNS/DNSParser.swift +++ b/Tun/DNS/DNSParser.swift @@ -34,6 +34,13 @@ struct DNSMessage { } } +struct DNSQuerySummary { + let transactionID: UInt16 + let name: String + let type: UInt16 + let qclass: UInt16 +} + // MARK: - DNS 完整解析器 final class DNSParser { private let data: Data @@ -135,3 +142,105 @@ final class DNSParser { return val } } + +extension DNSParser { + static func parseFirstQuestion(data: Data, offset: Int) -> DNSQuerySummary? { + guard offset >= 0, data.count >= offset + 12 else { + return nil + } + + return data.withUnsafeBytes { rawBuffer -> DNSQuerySummary? in + let bytes = rawBuffer.bindMemory(to: UInt8.self) + guard let baseAddress = bytes.baseAddress else { + return nil + } + + func readUInt16(at index: Int) -> UInt16 { + UInt16(baseAddress[index]) << 8 | UInt16(baseAddress[index + 1]) + } + + let transactionID = readUInt16(at: offset) + let questionCount = readUInt16(at: offset + 4) + guard questionCount > 0 else { + return nil + } + + var cursor = offset + 12 + guard let name = parseName( + baseAddress: baseAddress, + count: data.count, + messageStart: offset, + cursor: &cursor + ), cursor + 4 <= data.count else { + return nil + } + + return DNSQuerySummary( + transactionID: transactionID, + name: name, + type: readUInt16(at: cursor), + qclass: readUInt16(at: cursor + 2) + ) + } + } + + private static func parseName( + baseAddress: UnsafePointer, + count: Int, + messageStart: Int, + cursor: inout Int + ) -> String? { + var currentOffset = cursor + var resumeOffset: Int? + var jumpCount = 0 + var name = "" + + while currentOffset < count { + let length = Int(baseAddress[currentOffset]) + + if length == 0 { + currentOffset += 1 + cursor = resumeOffset ?? currentOffset + return name + } + + if (length & 0xC0) == 0xC0 { + guard currentOffset + 1 < count else { + return nil + } + + let pointer = Int(UInt16(baseAddress[currentOffset] & 0x3F) << 8 | UInt16(baseAddress[currentOffset + 1])) + let targetOffset = messageStart + pointer + guard targetOffset < count, jumpCount < 8 else { + return nil + } + + if resumeOffset == nil { + resumeOffset = currentOffset + 2 + } + currentOffset = targetOffset + jumpCount += 1 + continue + } + + guard (length & 0xC0) == 0, length <= 63 else { + return nil + } + + let labelStart = currentOffset + 1 + guard labelStart + length <= count else { + return nil + } + + if !name.isEmpty { + name.append(".") + } + + let labelBuffer = UnsafeBufferPointer(start: baseAddress.advanced(by: labelStart), count: length) + name.append(String(decoding: labelBuffer, as: UTF8.self)) + currentOffset = labelStart + length + } + + return nil + } +} diff --git a/Tun/Inbound/PacketInboundActor.swift b/Tun/Inbound/PacketInboundActor.swift index af46777..2423ac3 100644 --- a/Tun/Inbound/PacketInboundActor.swift +++ b/Tun/Inbound/PacketInboundActor.swift @@ -162,7 +162,7 @@ actor PacketInboundActor { } private func makeIPv4Plan(layerData: Data, identityID: UInt32, inboundBytes: Int, policyRuntime: PolicyRuntime) -> ProcessingPlan { - guard let ipPacket = IPPacket(layerData) else { + guard let ipPacket = IPPacketView(layerData) else { return .init(inboundBytes: inboundBytes, action: .none) } @@ -185,13 +185,13 @@ actor PacketInboundActor { } } - private func makePolicyPacketContext(identityID: UInt32, ipPacket: IPPacket) -> PolicyPacketContext { + private func makePolicyPacketContext(identityID: UInt32, ipPacket: IPPacketView) -> PolicyPacketContext { let ports: (UInt16?, UInt16?) switch ipPacket.transportPacket { - case .tcp(let tcpPacket): - ports = (tcpPacket.header.srcPort, tcpPacket.header.dstPort) - case .udp(let udpPacket): - ports = (udpPacket.srcPort, udpPacket.dstPort) + case .tcp(let srcPort, let dstPort): + ports = (srcPort, dstPort) + case .udp(let srcPort, let dstPort, _): + ports = (srcPort, dstPort) default: ports = (nil, nil) } diff --git a/Tun/NetworkStack/IPPacket.swift b/Tun/NetworkStack/IPPacket.swift index cb3988a..c4a97be 100644 --- a/Tun/NetworkStack/IPPacket.swift +++ b/Tun/NetworkStack/IPPacket.swift @@ -111,6 +111,91 @@ struct IPPacket { } +// MARK: - Lightweight IP Packet View + +struct IPPacketView { + let header: IPHeader + let data: Data + let transportPacket: TransportPacket + + enum TransportPacket { + case tcp(srcPort: UInt16, dstPort: UInt16) + case udp(srcPort: UInt16, dstPort: UInt16, payloadOffset: Int) + case icmp + case unsupported(UInt8) + case malformed + } + + init?(_ data: Data) { + guard data.count >= 20 else { + return nil + } + + let firstByte = Self.byte(at: 0, in: data) + let version = firstByte >> 4 + let headerLen = (firstByte & 0x0F) * 4 + + guard headerLen >= 20, data.count >= headerLen else { + return nil + } + + let proto = Self.byte(at: 9, in: data) + self.header = IPHeader( + version: version, + headerLength: headerLen, + typeOfService: Self.byte(at: 1, in: data), + totalLength: UInt16(bytes: (Self.byte(at: 2, in: data), Self.byte(at: 3, in: data))), + id: UInt16(bytes: (Self.byte(at: 4, in: data), Self.byte(at: 5, in: data))), + offset: UInt16(bytes: (Self.byte(at: 6, in: data), Self.byte(at: 7, in: data))), + ttl: Self.byte(at: 8, in: data), + proto: proto, + checksum: UInt16(bytes: (Self.byte(at: 10, in: data), Self.byte(at: 11, in: data))), + source: UInt32(bytes: (Self.byte(at: 12, in: data), Self.byte(at: 13, in: data), Self.byte(at: 14, in: data), Self.byte(at: 15, in: data))), + destination: UInt32(bytes: (Self.byte(at: 16, in: data), Self.byte(at: 17, in: data), Self.byte(at: 18, in: data), Self.byte(at: 19, in: data))) + ) + + self.data = data + self.transportPacket = Self.parseTransportPacket(proto: proto, offset: Int(headerLen), data: data) + } + + private static func parseTransportPacket(proto: UInt8, offset: Int, data: Data) -> TransportPacket { + guard let proto = TransportProtocol(rawValue: proto) else { + return .unsupported(proto) + } + + switch proto { + case .tcp: + guard data.count >= offset + 20 else { + return .malformed + } + return .tcp( + srcPort: UInt16(bytes: (Self.byte(at: offset, in: data), Self.byte(at: offset + 1, in: data))), + dstPort: UInt16(bytes: (Self.byte(at: offset + 2, in: data), Self.byte(at: offset + 3, in: data))) + ) + + case .udp: + guard data.count >= offset + 8 else { + return .malformed + } + return .udp( + srcPort: UInt16(bytes: (Self.byte(at: offset, in: data), Self.byte(at: offset + 1, in: data))), + dstPort: UInt16(bytes: (Self.byte(at: offset + 2, in: data), Self.byte(at: offset + 3, in: data))), + payloadOffset: offset + 8 + ) + + case .icmp: + guard data.count >= offset + 4 else { + return .malformed + } + return .icmp + } + } + + private static func byte(at offset: Int, in data: Data) -> UInt8 { + data[data.index(data.startIndex, offsetBy: offset)] + } +} + // MARK: - TCP Flags struct TCPFlags: OptionSet { diff --git a/Tun/Outbound/PacketOutboundActor.swift b/Tun/Outbound/PacketOutboundActor.swift index 72d012b..c50ecff 100644 --- a/Tun/Outbound/PacketOutboundActor.swift +++ b/Tun/Outbound/PacketOutboundActor.swift @@ -98,7 +98,7 @@ actor PacketOutboundActor { break } - if let packet = IPPacket(data) { + if let packet = IPPacketView(data) { await self?.handleTunPacket(packet) } } @@ -117,7 +117,7 @@ actor PacketOutboundActor { packetReaderTask?.cancel() } - func handleTunPacket(_ packet: IPPacket) async { + func handleTunPacket(_ packet: IPPacketView) async { let router = PacketOutboundRouter(networkAddress: self.networkAddress, exitNode: self.exitNode) let decision = router.route(packet: packet) diff --git a/Tun/Outbound/PacketOutboundRouter.swift b/Tun/Outbound/PacketOutboundRouter.swift index 5743575..14ac4fc 100644 --- a/Tun/Outbound/PacketOutboundRouter.swift +++ b/Tun/Outbound/PacketOutboundRouter.swift @@ -30,7 +30,7 @@ struct PacketOutboundRouter { let networkAddress: SDLConfiguration.NetworkAddress let exitNode: SDLConfiguration.ExitNode? - func route(packet: IPPacket, now: Date = Date()) -> RouteDecision { + func route(packet: IPPacketView, now: Date = Date()) -> RouteDecision { let dstIp = packet.header.destination // 本地通讯, 目标地址是本地服务器的ip地址 @@ -57,25 +57,22 @@ struct PacketOutboundRouter { return .drop(reason: .noRoute) } - private func routeDNS(packet: IPPacket, now: Date) -> RouteDecision? { + private func routeDNS(packet: IPPacketView, now: Date) -> RouteDecision? { guard DNSHelper.isDnsRequestPacket(ipPacket: packet) else { return nil } - guard case .udp(let udpPacket) = packet.transportPacket else { + guard case .udp(let srcPort, _, let payloadOffset) = packet.transportPacket else { return .drop(reason: .invalidDNSRequest) } - // 数据是通过offset解析的, dns查询必然是udp包 - let payloadOffset = udpPacket.payloadOffset - let dnsParser = DNSParser(data: packet.data, offset: payloadOffset) - guard let dnsMessage = dnsParser.parse(), let name = dnsMessage.questions.first?.name else { + guard let query = DNSParser.parseFirstQuestion(data: packet.data, offset: payloadOffset) else { return .drop(reason: .invalidDNSRequest) } // 如果是内部域名,则转发整个ip包的内容到云端服务器 - if name.contains(self.networkAddress.networkDomain) { - return .cloudDNS(name: name, ipPacketData: packet.data) + if query.name.contains(self.networkAddress.networkDomain) { + return .cloudDNS(name: query.name, ipPacketData: packet.data) } // 如果开启了出口节点,则转发给出口节点 @@ -86,12 +83,12 @@ struct PacketOutboundRouter { // 通过本地的dns解析,发送的是udp的payload部分 let dnsPayload = Data(packet.data[payloadOffset.. FlowSession? { + switch self.transportPacket { + case .tcp(let srcPort, let dstPort): + return FlowSession(srcIP: header.source, dstIP: header.destination, srcPort: srcPort, dstPort: dstPort, proto: header.proto) + case .udp(let srcPort, let dstPort, _): + return FlowSession(srcIP: header.source, dstIP: header.destination, srcPort: srcPort, dstPort: dstPort, proto: header.proto) + default: + return nil + } + } +} diff --git a/Tun/Policy/PolicyRuntime.swift b/Tun/Policy/PolicyRuntime.swift index 06a162a..4ccbf40 100644 --- a/Tun/Policy/PolicyRuntime.swift +++ b/Tun/Policy/PolicyRuntime.swift @@ -22,7 +22,7 @@ struct PolicyRuntime: @unchecked Sendable { self.flowSessionTable = flowSessionTable } - func evaluateInbound(srcIdentityID: UInt32, ipPacket: IPPacket) -> InboundDecision { + func evaluateInbound(srcIdentityID: UInt32, ipPacket: IPPacketView) -> InboundDecision { if let reverseFlowSession = ipPacket.flowSession()?.reverse(), self.flowSessionTable.hasSession(reverseFlowSession) { self.flowSessionTable.updateSession(reverseFlowSession) @@ -40,14 +40,14 @@ struct PolicyRuntime: @unchecked Sendable { return self.isAllowedByRule(ruleMap: ruleMap, ipPacket: ipPacket) ? .allow : .deny } - private func isAllowedByRule(ruleMap: PolicyRuleMap, ipPacket: IPPacket) -> Bool { + private func isAllowedByRule(ruleMap: PolicyRuleMap, ipPacket: IPPacketView) -> Bool { let proto = ipPacket.header.proto switch ipPacket.transportPacket { - case .tcp(let tcpPacket): - return ruleMap.isAllow(proto: proto, port: tcpPacket.header.dstPort) - case .udp(let udpPacket): - return ruleMap.isAllow(proto: proto, port: udpPacket.dstPort) + case .tcp(_, let dstPort): + return ruleMap.isAllow(proto: proto, port: dstPort) + case .udp(_, let dstPort, _): + return ruleMap.isAllow(proto: proto, port: dstPort) default: return false } diff --git a/Tun/Policy/PolicyService.swift b/Tun/Policy/PolicyService.swift index 2a72c72..e936172 100644 --- a/Tun/Policy/PolicyService.swift +++ b/Tun/Policy/PolicyService.swift @@ -30,7 +30,7 @@ actor PolicyService { return PolicyRuntime(policyRuleSnapshot: self.snapshotPublisher.current(), flowSessionTable: self.flowSessionTable) } - nonisolated func recordOutboundFlow(ipPacket: IPPacket) { + nonisolated func recordOutboundFlow(ipPacket: IPPacketView) { guard let flowSession = ipPacket.flowSession() else { return } diff --git a/Tun/Snapshot/SnapshotPublisher.swift b/Tun/Snapshot/SnapshotPublisher.swift index fdaee29..913a550 100644 --- a/Tun/Snapshot/SnapshotPublisher.swift +++ b/Tun/Snapshot/SnapshotPublisher.swift @@ -4,29 +4,27 @@ // // Created by 安礼成 on 2026/2/5. // -import Atomics +import Foundation final class SnapshotPublisher: @unchecked Sendable { - private let atomic: ManagedAtomic> + private let lock = NSLock() + private var snapshot: S init(initial snapshot: S) { - self.atomic = ManagedAtomic(.passRetained(snapshot)) + self.snapshot = snapshot } func publish(_ snapshot: S) { - let newRef = Unmanaged.passRetained(snapshot) - let oldRef = atomic.exchange(newRef, ordering: .releasing) - oldRef.release() + self.lock.lock() + self.snapshot = snapshot + self.lock.unlock() } @inline(__always) func current() -> S { - atomic.load(ordering: .acquiring).takeUnretainedValue() + self.lock.lock() + let snapshot = self.snapshot + self.lock.unlock() + return snapshot } - - deinit { - let ref = atomic.load(ordering: .acquiring) - ref.release() - } - }