krz/domain-dig

an ios app for DNS & SSL analysis

clone: git clone https://gitbay.org/krz/domain-dig.git

v1.5.0: DomainDig/DNSLookupService.swift · raw

  1import Foundation
  2
  3enum DNSResolverOption: String, CaseIterable, Identifiable {
  4    case cloudflare
  5    case google
  6    case quad9
  7    case custom
  8
  9    static let userDefaultsKey = "dnsResolverURL"
 10    static let defaultURLString = "https://cloudflare-dns.com/dns-query"
 11
 12    var id: String { rawValue }
 13
 14    var title: String {
 15        switch self {
 16        case .cloudflare: return "Cloudflare"
 17        case .google: return "Google"
 18        case .quad9: return "Quad9"
 19        case .custom: return "Custom"
 20        }
 21    }
 22
 23    var urlString: String? {
 24        switch self {
 25        case .cloudflare: return Self.defaultURLString
 26        case .google: return "https://dns.google/dns-query"
 27        case .quad9: return "https://dns.quad9.net/dns-query"
 28        case .custom: return nil
 29        }
 30    }
 31
 32    static func option(for urlString: String) -> DNSResolverOption {
 33        let trimmedURL = urlString.trimmingCharacters(in: .whitespacesAndNewlines)
 34        return Self.allCases.first(where: { $0.urlString == trimmedURL }) ?? .custom
 35    }
 36
 37    static func isValidCustomURL(_ urlString: String) -> Bool {
 38        let trimmedURL = urlString.trimmingCharacters(in: .whitespacesAndNewlines)
 39        guard trimmedURL.hasPrefix("https://") else {
 40            return false
 41        }
 42        return URL(string: trimmedURL) != nil
 43    }
 44
 45    static func resolvedURLString(from storedValue: String?) -> String {
 46        guard let storedValue else {
 47            return defaultURLString
 48        }
 49
 50        let trimmedURL = storedValue.trimmingCharacters(in: .whitespacesAndNewlines)
 51        guard !trimmedURL.isEmpty else {
 52            return defaultURLString
 53        }
 54
 55        return isValidCustomURL(trimmedURL) ? trimmedURL : defaultURLString
 56    }
 57}
 58
 59struct DNSLookupService {
 60    private static let rrsigQueryType = 46
 61    private static let internetClass = 1
 62
 63    static func lookup(domain: String, recordType: DNSRecordType) async throws -> [DNSRecord] {
 64        try await lookup(
 65            domain: domain,
 66            recordType: recordType,
 67            resolverURLString: currentResolverURLString()
 68        )
 69    }
 70
 71    static func lookup(
 72        domain: String,
 73        recordType: DNSRecordType,
 74        resolverURLString: String
 75    ) async throws -> [DNSRecord] {
 76        let answers = try await lookupAnswers(
 77            domain: domain,
 78            queryType: recordType.queryType,
 79            resolverURLString: resolverURLString
 80        )
 81
 82        return answers
 83            .filter { $0.type == recordType.queryType }
 84            .map { answer in
 85                let value: String
 86                if recordType.usesRawDataValue {
 87                    value = answer.data
 88                } else {
 89                    value = answer.data.trimmingCharacters(in: CharacterSet(charactersIn: "\""))
 90                }
 91                return DNSRecord(value: value, ttl: answer.TTL)
 92            }
 93    }
 94
 95    static func lookupAll(domain: String) async -> [DNSSection] {
 96        typealias Result = (
 97            type: DNSRecordType,
 98            records: [DNSRecord],
 99            wildcard: [DNSRecord],
100            dnssecSigned: Bool?,
101            error: String?
102        )
103
104        let wildcardTypes: Set<DNSRecordType> = [.A, .AAAA, .MX, .TXT, .SRV, .CAA]
105        let resolverURLString = currentResolverURLString()
106
107        return await withTaskGroup(of: Result.self, returning: [DNSSection].self) { group in
108            for recordType in DNSRecordType.allCases {
109                let shouldQueryWildcard = wildcardTypes.contains(recordType)
110                group.addTask {
111                    var apexRecords: [DNSRecord] = []
112                    var wildcardRecords: [DNSRecord] = []
113                    var dnssecSigned: Bool?
114                    var lookupError: String?
115
116                    do {
117                        apexRecords = try await lookup(
118                            domain: domain,
119                            recordType: recordType,
120                            resolverURLString: resolverURLString
121                        )
122                        dnssecSigned = try await lookupDNSSECStatus(
123                            domain: domain,
124                            resolverURLString: resolverURLString
125                        )
126                    } catch {
127                        lookupError = error.localizedDescription
128                    }
129
130                    if shouldQueryWildcard && lookupError == nil {
131                        do {
132                            wildcardRecords = try await lookup(
133                                domain: "*.\(domain)",
134                                recordType: recordType,
135                                resolverURLString: resolverURLString
136                            )
137                        } catch {
138                            // Wildcard failure is non-fatal; just leave empty.
139                        }
140                    }
141
142                    return (recordType, apexRecords, wildcardRecords, dnssecSigned, lookupError)
143                }
144            }
145
146            var sections: [DNSSection] = []
147            for await result in group {
148                sections.append(DNSSection(
149                    recordType: result.type,
150                    records: result.records,
151                    wildcardRecords: result.wildcard,
152                    dnssecSigned: result.dnssecSigned,
153                    error: result.error
154                ))
155            }
156
157            let order = DNSRecordType.allCases
158            return sections.sorted { a, b in
159                (order.firstIndex(of: a.recordType) ?? 0) < (order.firstIndex(of: b.recordType) ?? 0)
160            }
161        }
162    }
163
164    private static func lookupAnswers(
165        domain: String,
166        queryType: Int,
167        resolverURLString: String
168    ) async throws -> [CloudflareDNSResponse.CloudflareDNSAnswer] {
169        let resolverURL = try validatedResolverURL(from: resolverURLString)
170        var components = URLComponents(url: resolverURL, resolvingAgainstBaseURL: false)!
171        components.queryItems = [
172            URLQueryItem(name: "name", value: domain),
173            URLQueryItem(name: "type", value: String(queryType))
174        ]
175
176        var request = URLRequest(url: components.url!)
177        request.setValue("application/dns-json", forHTTPHeaderField: "Accept")
178
179        let (data, response) = try await URLSession.shared.data(for: request)
180
181        guard let httpResponse = response as? HTTPURLResponse,
182              httpResponse.statusCode == 200 else {
183            return try await lookupAnswersViaRFC8484(
184                domain: domain,
185                queryType: queryType,
186                resolverURL: resolverURL
187            )
188        }
189
190        let dnsResponse = try JSONDecoder().decode(CloudflareDNSResponse.self, from: data)
191
192        return dnsResponse.Answer ?? []
193    }
194
195    private static func lookupDNSSECStatus(
196        domain: String,
197        resolverURLString: String
198    ) async throws -> Bool {
199        let answers = try await lookupAnswers(
200            domain: domain,
201            queryType: rrsigQueryType,
202            resolverURLString: resolverURLString
203        )
204        return answers.contains(where: { $0.type == rrsigQueryType })
205    }
206
207    private static func currentResolverURLString() -> String {
208        let storedValue = UserDefaults.standard.string(forKey: DNSResolverOption.userDefaultsKey)
209        return DNSResolverOption.resolvedURLString(from: storedValue)
210    }
211
212    private static func validatedResolverURL(from urlString: String) throws -> URL {
213        guard let url = URL(string: urlString) else {
214            throw URLError(.badURL)
215        }
216        return url
217    }
218
219    private static func lookupAnswersViaRFC8484(
220        domain: String,
221        queryType: Int,
222        resolverURL: URL
223    ) async throws -> [CloudflareDNSResponse.CloudflareDNSAnswer] {
224        let queryData = try buildDNSQueryMessage(domain: domain, queryType: queryType)
225        let encodedQuery = base64URLEncodedString(for: queryData)
226
227        var components = URLComponents(url: resolverURL, resolvingAgainstBaseURL: false)!
228        components.queryItems = [URLQueryItem(name: "dns", value: encodedQuery)]
229
230        var request = URLRequest(url: components.url!)
231        request.setValue("application/dns-message", forHTTPHeaderField: "Accept")
232
233        let (data, response) = try await URLSession.shared.data(for: request)
234
235        guard let httpResponse = response as? HTTPURLResponse,
236              httpResponse.statusCode == 200 else {
237            throw URLError(.badServerResponse)
238        }
239
240        return try parseDNSMessage(data)
241    }
242
243    private static func buildDNSQueryMessage(domain: String, queryType: Int) throws -> Data {
244        let normalizedName = domain.trimmingCharacters(in: .whitespacesAndNewlines)
245        let labels = normalizedName.split(separator: ".")
246
247        var data = Data()
248        data.appendUInt16(UInt16.random(in: UInt16.min ... UInt16.max))
249        data.appendUInt16(0x0100)
250        data.appendUInt16(1)
251        data.appendUInt16(0)
252        data.appendUInt16(0)
253        data.appendUInt16(0)
254
255        for label in labels {
256            guard let labelData = label.data(using: .utf8),
257                  labelData.count <= 63 else {
258                throw URLError(.badURL)
259            }
260            data.append(UInt8(labelData.count))
261            data.append(labelData)
262        }
263
264        data.append(0)
265        data.appendUInt16(UInt16(queryType))
266        data.appendUInt16(UInt16(internetClass))
267
268        return data
269    }
270
271    private static func base64URLEncodedString(for data: Data) -> String {
272        data.base64EncodedString()
273            .replacingOccurrences(of: "+", with: "-")
274            .replacingOccurrences(of: "/", with: "_")
275            .replacingOccurrences(of: "=", with: "")
276    }
277
278    private static func parseDNSMessage(_ data: Data) throws -> [CloudflareDNSResponse.CloudflareDNSAnswer] {
279        guard data.count >= 12 else {
280            throw URLError(.cannotParseResponse)
281        }
282
283        let answerCount = Int(readUInt16(in: data, at: 6))
284        let questionCount = Int(readUInt16(in: data, at: 4))
285        var offset = 12
286
287        for _ in 0 ..< questionCount {
288            _ = try readDomainName(in: data, offset: &offset)
289            offset += 4
290        }
291
292        var answers: [CloudflareDNSResponse.CloudflareDNSAnswer] = []
293        for _ in 0 ..< answerCount {
294            let name = try readDomainName(in: data, offset: &offset)
295            let type = Int(readUInt16(in: data, at: offset))
296            offset += 2
297            _ = readUInt16(in: data, at: offset)
298            offset += 2
299            let ttl = Int(readUInt32(in: data, at: offset))
300            offset += 4
301            let dataLength = Int(readUInt16(in: data, at: offset))
302            offset += 2
303
304            guard offset + dataLength <= data.count else {
305                throw URLError(.cannotParseResponse)
306            }
307
308            let recordDataOffset = offset
309            let recordData = data.subdata(in: recordDataOffset ..< (recordDataOffset + dataLength))
310            offset += dataLength
311
312            let parsedValue = try parseRecordData(
313                from: data,
314                recordType: type,
315                recordDataOffset: recordDataOffset,
316                recordData: recordData
317            )
318
319            answers.append(.init(
320                name: name,
321                type: type,
322                TTL: ttl,
323                data: parsedValue
324            ))
325        }
326
327        return answers
328    }
329
330    private static func parseRecordData(
331        from message: Data,
332        recordType: Int,
333        recordDataOffset: Int,
334        recordData: Data
335    ) throws -> String {
336        switch recordType {
337        case 1:
338            guard recordData.count == 4 else { throw URLError(.cannotParseResponse) }
339            return recordData.map(String.init).joined(separator: ".")
340        case 2, 5:
341            var offset = recordDataOffset
342            return try readDomainName(in: message, offset: &offset)
343        case 15:
344            guard recordData.count >= 3 else { throw URLError(.cannotParseResponse) }
345            let preference = readUInt16(in: recordData, at: 0)
346            var exchangeOffset = recordDataOffset + 2
347            let exchange = try readDomainName(in: message, offset: &exchangeOffset)
348            return "\(preference) \(exchange)"
349        case 16:
350            return try parseTXTData(recordData)
351        case 28:
352            guard recordData.count == 16 else { throw URLError(.cannotParseResponse) }
353            return stride(from: 0, to: 16, by: 2)
354                .map { index in
355                    String(format: "%x", readUInt16(in: recordData, at: index))
356                }
357                .joined(separator: ":")
358        case 6:
359            var offset = recordDataOffset
360            let mname = try readDomainName(in: message, offset: &offset)
361            let rname = try readDomainName(in: message, offset: &offset)
362            let serial = readUInt32(in: message, at: offset)
363            let refresh = readUInt32(in: message, at: offset + 4)
364            let retry = readUInt32(in: message, at: offset + 8)
365            let expire = readUInt32(in: message, at: offset + 12)
366            let minimum = readUInt32(in: message, at: offset + 16)
367            return "\(mname) \(rname) \(serial) \(refresh) \(retry) \(expire) \(minimum)"
368        case 33:
369            guard recordData.count >= 7 else { throw URLError(.cannotParseResponse) }
370            let priority = readUInt16(in: recordData, at: 0)
371            let weight = readUInt16(in: recordData, at: 2)
372            let port = readUInt16(in: recordData, at: 4)
373            var targetOffset = recordDataOffset + 6
374            let target = try readDomainName(in: message, offset: &targetOffset)
375            return "\(priority) \(weight) \(port) \(target)"
376        case 43:
377            guard recordData.count >= 4 else { throw URLError(.cannotParseResponse) }
378            let keyTag = readUInt16(in: recordData, at: 0)
379            let algorithm = recordData[2]
380            let digestType = recordData[3]
381            let digest = recordData.dropFirst(4).map { String(format: "%02X", $0) }.joined()
382            return "\(keyTag) \(algorithm) \(digestType) \(digest)"
383        case 46:
384            return "RRSIG"
385        case 257:
386            guard recordData.count >= 2 else { throw URLError(.cannotParseResponse) }
387            let flags = recordData[0]
388            let tagLength = Int(recordData[1])
389            guard recordData.count >= 2 + tagLength else {
390                throw URLError(.cannotParseResponse)
391            }
392            let tagData = recordData.subdata(in: 2 ..< (2 + tagLength))
393            let valueData = recordData.dropFirst(2 + tagLength)
394            let tag = String(decoding: tagData, as: UTF8.self)
395            let value = String(decoding: valueData, as: UTF8.self)
396            return "\(flags) \(tag) \"\(value)\""
397        default:
398            return recordData.base64EncodedString()
399        }
400    }
401
402    private static func parseTXTData(_ data: Data) throws -> String {
403        var offset = 0
404        var strings: [String] = []
405
406        while offset < data.count {
407            let count = Int(data[offset])
408            offset += 1
409            guard offset + count <= data.count else {
410                throw URLError(.cannotParseResponse)
411            }
412            let stringData = data.subdata(in: offset ..< (offset + count))
413            strings.append(String(decoding: stringData, as: UTF8.self))
414            offset += count
415        }
416
417        return strings.joined()
418    }
419
420    private static func readDomainName(in data: Data, offset: inout Int) throws -> String {
421        var labels: [String] = []
422        var currentOffset = offset
423        var jumped = false
424        var seenOffsets = Set<Int>()
425
426        while true {
427            guard currentOffset < data.count else {
428                throw URLError(.cannotParseResponse)
429            }
430
431            let length = Int(data[currentOffset])
432
433            if length == 0 {
434                if !jumped {
435                    offset = currentOffset + 1
436                }
437                break
438            }
439
440            if length & 0xC0 == 0xC0 {
441                guard currentOffset + 1 < data.count else {
442                    throw URLError(.cannotParseResponse)
443                }
444
445                let pointer = ((length & 0x3F) << 8) | Int(data[currentOffset + 1])
446                guard seenOffsets.insert(pointer).inserted else {
447                    throw URLError(.cannotParseResponse)
448                }
449
450                if !jumped {
451                    offset = currentOffset + 2
452                }
453                currentOffset = pointer
454                jumped = true
455                continue
456            }
457
458            let labelStart = currentOffset + 1
459            let labelEnd = labelStart + length
460            guard labelEnd <= data.count else {
461                throw URLError(.cannotParseResponse)
462            }
463
464            let labelData = data.subdata(in: labelStart ..< labelEnd)
465            labels.append(String(decoding: labelData, as: UTF8.self))
466            currentOffset = labelEnd
467        }
468
469        return labels.joined(separator: ".")
470    }
471
472    private static func readUInt16(in data: Data, at offset: Int) -> UInt16 {
473        let upper = UInt16(data[offset]) << 8
474        let lower = UInt16(data[offset + 1])
475        return upper | lower
476    }
477
478    private static func readUInt32(in data: Data, at offset: Int) -> UInt32 {
479        let first = UInt32(data[offset]) << 24
480        let second = UInt32(data[offset + 1]) << 16
481        let third = UInt32(data[offset + 2]) << 8
482        let fourth = UInt32(data[offset + 3])
483        return first | second | third | fourth
484    }
485}
486
487private extension Data {
488    mutating func appendUInt16(_ value: UInt16) {
489        append(UInt8((value >> 8) & 0xFF))
490        append(UInt8(value & 0xFF))
491    }
492}