DomainDig/SSLCheckService.swift

ede15c586af71b8d6f2c190ffdad0324414cac06
domain-dig/DomainDig/SSLCheckService.swift history · blame · raw

322 lines · 12270 bytes

  1import Foundation
  2import Security
  3
  4struct SSLCheckService {
  5
  6    static func check(domain: String) async throws -> SSLCertificateInfo {
  7        let delegate = SSLSessionDelegate()
  8        let session = URLSession(
  9            configuration: .ephemeral,
 10            delegate: delegate,
 11            delegateQueue: nil
 12        )
 13        defer { session.invalidateAndCancel() }
 14
 15        let url = URL(string: "https://\(domain)")!
 16        let request = URLRequest(url: url, timeoutInterval: 10)
 17
 18        // We only need to establish the connection to grab the cert
 19        _ = try await session.data(for: request)
 20
 21        guard let trust = delegate.serverTrust else {
 22            throw SSLError.noCertificate
 23        }
 24
 25        return try extractCertificateInfo(from: trust)
 26    }
 27
 28    private static func extractCertificateInfo(from trust: SecTrust) throws -> SSLCertificateInfo {
 29        let chainCount = SecTrustGetCertificateCount(trust)
 30        guard chainCount > 0,
 31              let certChain = SecTrustCopyCertificateChain(trust) as? [SecCertificate],
 32              let leaf = certChain.first else {
 33            throw SSLError.noCertificate
 34        }
 35
 36        // Common Name — use subject summary (available on iOS)
 37        let commonName = SecCertificateCopySubjectSummary(leaf) as String? ?? "Unknown"
 38
 39        // Validity dates
 40        let validFrom: Date
 41        let validUntil: Date
 42
 43        if let notBefore = SecCertificateCopyNotValidBeforeDate(leaf) as Date? {
 44            validFrom = notBefore
 45        } else {
 46            validFrom = Date.distantPast
 47        }
 48
 49        if let notAfter = SecCertificateCopyNotValidAfterDate(leaf) as Date? {
 50            validUntil = notAfter
 51        } else {
 52            validUntil = Date.distantFuture
 53        }
 54
 55        let daysUntilExpiry = Calendar.current.dateComponents([.day], from: Date(), to: validUntil).day ?? 0
 56
 57        // Parse the DER-encoded certificate to extract SANs and Issuer
 58        let derData = SecCertificateCopyData(leaf) as Data
 59        let parsed = DERCertificateParser.parse(derData)
 60
 61        let sans = parsed.subjectAltNames.isEmpty ? [commonName] : parsed.subjectAltNames
 62
 63        // Issuer: prefer parsed issuer, fall back to chain's next cert summary
 64        var issuer = parsed.issuerCommonName ?? "Unknown"
 65        if issuer == "Unknown" && certChain.count > 1 {
 66            let issuerCert = certChain[1]
 67            if let issuerSummary = SecCertificateCopySubjectSummary(issuerCert) as String? {
 68                issuer = issuerSummary
 69            }
 70        }
 71
 72        return SSLCertificateInfo(
 73            commonName: commonName,
 74            subjectAltNames: sans,
 75            issuer: issuer,
 76            validFrom: validFrom,
 77            validUntil: validUntil,
 78            daysUntilExpiry: daysUntilExpiry,
 79            chainDepth: Int(chainCount)
 80        )
 81    }
 82}
 83
 84// MARK: - Minimal DER/ASN.1 parser for X.509 certificate fields
 85
 86private enum DERCertificateParser {
 87    struct Result {
 88        var issuerCommonName: String?
 89        var subjectAltNames: [String] = []
 90    }
 91
 92    static func parse(_ data: Data) -> Result {
 93        var result = Result()
 94        let bytes = [UInt8](data)
 95
 96        // X.509 structure: SEQUENCE { tbsCertificate, signatureAlgorithm, signatureValue }
 97        // tbsCertificate: SEQUENCE { version, serialNumber, signature, issuer, validity, subject, ... extensions }
 98        guard let tbsRange = readSequence(bytes, offset: 0),
 99              let tbsContent = readSequence(bytes, offset: tbsRange.contentStart) else {
100            return result
101        }
102
103        var offset = tbsContent.contentStart
104
105        // Skip version (explicit tag [0]) if present
106        if offset < bytes.count && (bytes[offset] & 0xE0) == 0xA0 {
107            if let tagLen = readTagAndLength(bytes, offset: offset) {
108                offset = tagLen.contentStart + tagLen.length
109            }
110        }
111
112        // Skip serialNumber
113        if let serial = readTagAndLength(bytes, offset: offset) {
114            offset = serial.contentStart + serial.length
115        }
116
117        // Skip signature algorithm
118        if let sigAlg = readTagAndLength(bytes, offset: offset) {
119            offset = sigAlg.contentStart + sigAlg.length
120        }
121
122        // Issuer — a SEQUENCE of SETs of attribute type-value pairs
123        if let issuerSeq = readTagAndLength(bytes, offset: offset) {
124            result.issuerCommonName = extractCommonName(bytes, sequenceStart: issuerSeq.contentStart, length: issuerSeq.length)
125            offset = issuerSeq.contentStart + issuerSeq.length
126        }
127
128        // Skip validity
129        if let validity = readTagAndLength(bytes, offset: offset) {
130            offset = validity.contentStart + validity.length
131        }
132
133        // Skip subject
134        if let subject = readTagAndLength(bytes, offset: offset) {
135            offset = subject.contentStart + subject.length
136        }
137
138        // Skip subjectPublicKeyInfo
139        if let spki = readTagAndLength(bytes, offset: offset) {
140            offset = spki.contentStart + spki.length
141        }
142
143        // Extensions are in an explicit tag [3]
144        while offset < tbsContent.contentStart + tbsContent.length {
145            if bytes[offset] == 0xA3 {
146                if let extWrapper = readTagAndLength(bytes, offset: offset) {
147                    // Inside is a SEQUENCE of SEQUENCE extensions
148                    if let extsSeq = readTagAndLength(bytes, offset: extWrapper.contentStart) {
149                        result.subjectAltNames = extractSANs(bytes, sequenceStart: extsSeq.contentStart, length: extsSeq.length)
150                    }
151                }
152                break
153            }
154            // Skip optional implicit tags (issuerUniqueID [1], subjectUniqueID [2])
155            if let tl = readTagAndLength(bytes, offset: offset) {
156                offset = tl.contentStart + tl.length
157            } else {
158                break
159            }
160        }
161
162        return result
163    }
164
165    // OID for commonName: 2.5.4.3 = 55 04 03
166    private static let cnOID: [UInt8] = [0x55, 0x04, 0x03]
167
168    // OID for subjectAltName: 2.5.29.17 = 55 1D 11
169    private static let sanOID: [UInt8] = [0x55, 0x1D, 0x11]
170
171    private static func extractCommonName(_ bytes: [UInt8], sequenceStart: Int, length: Int) -> String? {
172        let end = sequenceStart + length
173        var pos = sequenceStart
174        while pos < end {
175            // Each SET in the issuer
176            guard let setTL = readTagAndLength(bytes, offset: pos) else { break }
177            let setEnd = setTL.contentStart + setTL.length
178
179            // Inside the SET is a SEQUENCE with OID + value
180            if let seqTL = readTagAndLength(bytes, offset: setTL.contentStart) {
181                let seqEnd = seqTL.contentStart + seqTL.length
182                if let oidTL = readTagAndLength(bytes, offset: seqTL.contentStart) {
183                    let oidBytes = Array(bytes[oidTL.contentStart..<oidTL.contentStart + oidTL.length])
184                    if oidBytes == cnOID {
185                        let valueStart = oidTL.contentStart + oidTL.length
186                        if let valueTL = readTagAndLength(bytes, offset: valueStart) {
187                            let strBytes = bytes[valueTL.contentStart..<valueTL.contentStart + valueTL.length]
188                            return String(bytes: strBytes, encoding: .utf8)
189                        }
190                    }
191                    _ = seqEnd // suppress unused warning
192                }
193            }
194            pos = setEnd
195        }
196        return nil
197    }
198
199    private static func extractSANs(_ bytes: [UInt8], sequenceStart: Int, length: Int) -> [String] {
200        let end = sequenceStart + length
201        var pos = sequenceStart
202        var sans: [String] = []
203
204        while pos < end {
205            guard let extSeq = readTagAndLength(bytes, offset: pos) else { break }
206            let extEnd = extSeq.contentStart + extSeq.length
207
208            // Each extension is SEQUENCE { OID, [critical], value }
209            if let oidTL = readTagAndLength(bytes, offset: extSeq.contentStart) {
210                let oidBytes = Array(bytes[oidTL.contentStart..<oidTL.contentStart + oidTL.length])
211                if oidBytes == sanOID {
212                    var valuePos = oidTL.contentStart + oidTL.length
213                    // Skip optional critical BOOLEAN
214                    if valuePos < extEnd && bytes[valuePos] == 0x01 {
215                        if let boolTL = readTagAndLength(bytes, offset: valuePos) {
216                            valuePos = boolTL.contentStart + boolTL.length
217                        }
218                    }
219                    // The value is an OCTET STRING wrapping a SEQUENCE of GeneralNames
220                    if let octetTL = readTagAndLength(bytes, offset: valuePos) {
221                        if let sanSeq = readTagAndLength(bytes, offset: octetTL.contentStart) {
222                            let sanEnd = sanSeq.contentStart + sanSeq.length
223                            var sanPos = sanSeq.contentStart
224                            while sanPos < sanEnd {
225                                guard let nameTL = readTagAndLength(bytes, offset: sanPos) else { break }
226                                // Context tag [2] = dNSName (IA5String)
227                                if (bytes[sanPos] & 0x1F) == 2 {
228                                    let nameBytes = bytes[nameTL.contentStart..<nameTL.contentStart + nameTL.length]
229                                    if let name = String(bytes: nameBytes, encoding: .ascii) {
230                                        sans.append(name)
231                                    }
232                                }
233                                sanPos = nameTL.contentStart + nameTL.length
234                            }
235                        }
236                    }
237                }
238            }
239            pos = extEnd
240        }
241        return sans
242    }
243
244    private struct TLV {
245        let contentStart: Int
246        let length: Int
247    }
248
249    private static func readSequence(_ bytes: [UInt8], offset: Int) -> TLV? {
250        guard offset < bytes.count, bytes[offset] == 0x30 else { return nil }
251        return readTagAndLength(bytes, offset: offset)
252    }
253
254    private static func readTagAndLength(_ bytes: [UInt8], offset: Int) -> TLV? {
255        guard offset < bytes.count else { return nil }
256        var pos = offset + 1 // skip tag byte
257        guard pos < bytes.count else { return nil }
258
259        let firstLen = bytes[pos]
260        pos += 1
261
262        let length: Int
263        if firstLen < 0x80 {
264            length = Int(firstLen)
265        } else {
266            let numBytes = Int(firstLen & 0x7F)
267            guard numBytes > 0, numBytes <= 4, pos + numBytes <= bytes.count else { return nil }
268            var len = 0
269            for i in 0..<numBytes {
270                len = (len << 8) | Int(bytes[pos + i])
271            }
272            pos += numBytes
273            length = len
274        }
275
276        return TLV(contentStart: pos, length: length)
277    }
278}
279
280enum SSLError: LocalizedError {
281    case noCertificate
282    case connectionFailed
283
284    var errorDescription: String? {
285        switch self {
286        case .noCertificate:
287            return "No certificate found"
288        case .connectionFailed:
289            return "Failed to connect to server"
290        }
291    }
292}
293
294final class SSLSessionDelegate: NSObject, URLSessionDelegate, @unchecked Sendable {
295    private let lock = NSLock()
296    private var _serverTrust: SecTrust?
297
298    var serverTrust: SecTrust? {
299        lock.lock()
300        defer { lock.unlock() }
301        return _serverTrust
302    }
303
304    func urlSession(
305        _ session: URLSession,
306        didReceive challenge: URLAuthenticationChallenge,
307        completionHandler: @escaping (URLSession.AuthChallengeDisposition, URLCredential?) -> Void
308    ) {
309        guard challenge.protectionSpace.authenticationMethod == NSURLAuthenticationMethodServerTrust,
310              let trust = challenge.protectionSpace.serverTrust else {
311            completionHandler(.performDefaultHandling, nil)
312            return
313        }
314
315        lock.lock()
316        _serverTrust = trust
317        lock.unlock()
318
319        let credential = URLCredential(trust: trust)
320        completionHandler(.useCredential, credential)
321    }
322}