完善代码逻辑
This commit is contained in:
parent
2cf0bcec73
commit
b3526cac4d
@ -12,7 +12,7 @@ struct DNSHelper {
|
|||||||
static let dnsDestIpAddr: UInt32 = 1684300900
|
static let dnsDestIpAddr: UInt32 = 1684300900
|
||||||
|
|
||||||
// 判断是否是dns请求的数据包
|
// 判断是否是dns请求的数据包
|
||||||
static func isDnsRequestPacket(ipPacket: IPPacket) -> Bool {
|
static func isDnsRequestPacket(ipPacket: IPPacketView) -> Bool {
|
||||||
return ipPacket.header.destination == dnsDestIpAddr
|
return ipPacket.header.destination == dnsDestIpAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -34,6 +34,13 @@ struct DNSMessage {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct DNSQuerySummary {
|
||||||
|
let transactionID: UInt16
|
||||||
|
let name: String
|
||||||
|
let type: UInt16
|
||||||
|
let qclass: UInt16
|
||||||
|
}
|
||||||
|
|
||||||
// MARK: - DNS 完整解析器
|
// MARK: - DNS 完整解析器
|
||||||
final class DNSParser {
|
final class DNSParser {
|
||||||
private let data: Data
|
private let data: Data
|
||||||
@ -135,3 +142,105 @@ final class DNSParser {
|
|||||||
return val
|
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<UInt8>,
|
||||||
|
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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@ -162,7 +162,7 @@ actor PacketInboundActor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private func makeIPv4Plan(layerData: Data, identityID: UInt32, inboundBytes: Int, policyRuntime: PolicyRuntime) -> ProcessingPlan {
|
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)
|
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?)
|
let ports: (UInt16?, UInt16?)
|
||||||
switch ipPacket.transportPacket {
|
switch ipPacket.transportPacket {
|
||||||
case .tcp(let tcpPacket):
|
case .tcp(let srcPort, let dstPort):
|
||||||
ports = (tcpPacket.header.srcPort, tcpPacket.header.dstPort)
|
ports = (srcPort, dstPort)
|
||||||
case .udp(let udpPacket):
|
case .udp(let srcPort, let dstPort, _):
|
||||||
ports = (udpPacket.srcPort, udpPacket.dstPort)
|
ports = (srcPort, dstPort)
|
||||||
default:
|
default:
|
||||||
ports = (nil, nil)
|
ports = (nil, nil)
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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
|
// MARK: - TCP Flags
|
||||||
|
|
||||||
struct TCPFlags: OptionSet {
|
struct TCPFlags: OptionSet {
|
||||||
|
|||||||
@ -98,7 +98,7 @@ actor PacketOutboundActor {
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
if let packet = IPPacket(data) {
|
if let packet = IPPacketView(data) {
|
||||||
await self?.handleTunPacket(packet)
|
await self?.handleTunPacket(packet)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@ -117,7 +117,7 @@ actor PacketOutboundActor {
|
|||||||
packetReaderTask?.cancel()
|
packetReaderTask?.cancel()
|
||||||
}
|
}
|
||||||
|
|
||||||
func handleTunPacket(_ packet: IPPacket) async {
|
func handleTunPacket(_ packet: IPPacketView) async {
|
||||||
let router = PacketOutboundRouter(networkAddress: self.networkAddress, exitNode: self.exitNode)
|
let router = PacketOutboundRouter(networkAddress: self.networkAddress, exitNode: self.exitNode)
|
||||||
let decision = router.route(packet: packet)
|
let decision = router.route(packet: packet)
|
||||||
|
|
||||||
|
|||||||
@ -30,7 +30,7 @@ struct PacketOutboundRouter {
|
|||||||
let networkAddress: SDLConfiguration.NetworkAddress
|
let networkAddress: SDLConfiguration.NetworkAddress
|
||||||
let exitNode: SDLConfiguration.ExitNode?
|
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
|
let dstIp = packet.header.destination
|
||||||
|
|
||||||
// 本地通讯, 目标地址是本地服务器的ip地址
|
// 本地通讯, 目标地址是本地服务器的ip地址
|
||||||
@ -57,25 +57,22 @@ struct PacketOutboundRouter {
|
|||||||
return .drop(reason: .noRoute)
|
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 {
|
guard DNSHelper.isDnsRequestPacket(ipPacket: packet) else {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
guard case .udp(let udpPacket) = packet.transportPacket else {
|
guard case .udp(let srcPort, _, let payloadOffset) = packet.transportPacket else {
|
||||||
return .drop(reason: .invalidDNSRequest)
|
return .drop(reason: .invalidDNSRequest)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 数据是通过offset解析的, dns查询必然是udp包
|
guard let query = DNSParser.parseFirstQuestion(data: packet.data, offset: payloadOffset) else {
|
||||||
let payloadOffset = udpPacket.payloadOffset
|
|
||||||
let dnsParser = DNSParser(data: packet.data, offset: payloadOffset)
|
|
||||||
guard let dnsMessage = dnsParser.parse(), let name = dnsMessage.questions.first?.name else {
|
|
||||||
return .drop(reason: .invalidDNSRequest)
|
return .drop(reason: .invalidDNSRequest)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 如果是内部域名,则转发整个ip包的内容到云端服务器
|
// 如果是内部域名,则转发整个ip包的内容到云端服务器
|
||||||
if name.contains(self.networkAddress.networkDomain) {
|
if query.name.contains(self.networkAddress.networkDomain) {
|
||||||
return .cloudDNS(name: name, ipPacketData: packet.data)
|
return .cloudDNS(name: query.name, ipPacketData: packet.data)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 如果开启了出口节点,则转发给出口节点
|
// 如果开启了出口节点,则转发给出口节点
|
||||||
@ -86,12 +83,12 @@ struct PacketOutboundRouter {
|
|||||||
// 通过本地的dns解析,发送的是udp的payload部分
|
// 通过本地的dns解析,发送的是udp的payload部分
|
||||||
let dnsPayload = Data(packet.data[payloadOffset..<packet.data.count])
|
let dnsPayload = Data(packet.data[payloadOffset..<packet.data.count])
|
||||||
let tracker = DNSLocalClient.DNSTracker(
|
let tracker = DNSLocalClient.DNSTracker(
|
||||||
transactionID: dnsMessage.transactionID,
|
transactionID: query.transactionID,
|
||||||
clientIP: packet.header.source,
|
clientIP: packet.header.source,
|
||||||
clientPort: udpPacket.srcPort,
|
clientPort: srcPort,
|
||||||
createdAt: now
|
createdAt: now
|
||||||
)
|
)
|
||||||
return .localDNS(name: name, payload: dnsPayload, tracker: tracker)
|
return .localDNS(name: query.name, payload: dnsPayload, tracker: tracker)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -133,3 +133,17 @@ extension IPPacket {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
extension IPPacketView {
|
||||||
|
|
||||||
|
func flowSession() -> 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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@ -22,7 +22,7 @@ struct PolicyRuntime: @unchecked Sendable {
|
|||||||
self.flowSessionTable = flowSessionTable
|
self.flowSessionTable = flowSessionTable
|
||||||
}
|
}
|
||||||
|
|
||||||
func evaluateInbound(srcIdentityID: UInt32, ipPacket: IPPacket) -> InboundDecision {
|
func evaluateInbound(srcIdentityID: UInt32, ipPacket: IPPacketView) -> InboundDecision {
|
||||||
if let reverseFlowSession = ipPacket.flowSession()?.reverse(),
|
if let reverseFlowSession = ipPacket.flowSession()?.reverse(),
|
||||||
self.flowSessionTable.hasSession(reverseFlowSession) {
|
self.flowSessionTable.hasSession(reverseFlowSession) {
|
||||||
self.flowSessionTable.updateSession(reverseFlowSession)
|
self.flowSessionTable.updateSession(reverseFlowSession)
|
||||||
@ -40,14 +40,14 @@ struct PolicyRuntime: @unchecked Sendable {
|
|||||||
return self.isAllowedByRule(ruleMap: ruleMap, ipPacket: ipPacket) ? .allow : .deny
|
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
|
let proto = ipPacket.header.proto
|
||||||
|
|
||||||
switch ipPacket.transportPacket {
|
switch ipPacket.transportPacket {
|
||||||
case .tcp(let tcpPacket):
|
case .tcp(_, let dstPort):
|
||||||
return ruleMap.isAllow(proto: proto, port: tcpPacket.header.dstPort)
|
return ruleMap.isAllow(proto: proto, port: dstPort)
|
||||||
case .udp(let udpPacket):
|
case .udp(_, let dstPort, _):
|
||||||
return ruleMap.isAllow(proto: proto, port: udpPacket.dstPort)
|
return ruleMap.isAllow(proto: proto, port: dstPort)
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|||||||
@ -30,7 +30,7 @@ actor PolicyService {
|
|||||||
return PolicyRuntime(policyRuleSnapshot: self.snapshotPublisher.current(), flowSessionTable: self.flowSessionTable)
|
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 {
|
guard let flowSession = ipPacket.flowSession() else {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@ -4,29 +4,27 @@
|
|||||||
//
|
//
|
||||||
// Created by 安礼成 on 2026/2/5.
|
// Created by 安礼成 on 2026/2/5.
|
||||||
//
|
//
|
||||||
import Atomics
|
import Foundation
|
||||||
|
|
||||||
final class SnapshotPublisher<S: Snapshot>: @unchecked Sendable {
|
final class SnapshotPublisher<S: Snapshot>: @unchecked Sendable {
|
||||||
private let atomic: ManagedAtomic<Unmanaged<S>>
|
private let lock = NSLock()
|
||||||
|
private var snapshot: S
|
||||||
|
|
||||||
init(initial snapshot: S) {
|
init(initial snapshot: S) {
|
||||||
self.atomic = ManagedAtomic(.passRetained(snapshot))
|
self.snapshot = snapshot
|
||||||
}
|
}
|
||||||
|
|
||||||
func publish(_ snapshot: S) {
|
func publish(_ snapshot: S) {
|
||||||
let newRef = Unmanaged.passRetained(snapshot)
|
self.lock.lock()
|
||||||
let oldRef = atomic.exchange(newRef, ordering: .releasing)
|
self.snapshot = snapshot
|
||||||
oldRef.release()
|
self.lock.unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
@inline(__always)
|
@inline(__always)
|
||||||
func current() -> S {
|
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()
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user