diff --git a/ApplicationLibrary/Service/NWSocket.swift b/ApplicationLibrary/Service/NWSocket.swift index 292c674..4a98838 100644 --- a/ApplicationLibrary/Service/NWSocket.swift +++ b/ApplicationLibrary/Service/NWSocket.swift @@ -2,52 +2,83 @@ import Foundation import Libbox import Network -public class NWSocket { +public enum NWSocketError: Error { + case connectionClosed + case invalidLength(Int) + case messageTooLarge(Int) + case timeout(String) +} + +extension NWSocketError: LocalizedError { + public var errorDescription: String? { + switch self { + case .connectionClosed: + "Connection closed" + case let .invalidLength(length): + "Invalid message length: \(length)" + case let .messageTooLarge(length): + "Message too large: \(length)" + case let .timeout(phase): + "Timed out: \(phase)" + } + } +} + +private final class OneShot: @unchecked Sendable { + private let lock = NSLock() + private var continuation: CheckedContinuation? + + init(_ continuation: CheckedContinuation) { + self.continuation = continuation + } + + @discardableResult + func resume(_ result: Result) -> Bool { + lock.lock() + guard let continuation else { + lock.unlock() + return false + } + self.continuation = nil + lock.unlock() + continuation.resume(with: result) + return true + } +} + +public final class NWSocket { private let connection: NWConnection public init(_ connection: NWConnection) { self.connection = connection } - public func read() throws -> Data { - let semaphore = DispatchSemaphore(value: 0) - var result: Result! - connection.receive(minimumIncompleteLength: 2, maximumLength: 2) { content, _, _, error in - if let error { - result = .failure(error) - } else { - result = .success(content!) - } - semaphore.signal() - } - semaphore.wait() - let lengthChunk = try result.get() + public func read( + headerTimeout: TimeInterval = 0, + bodyTimeout: TimeInterval = 60, + maxMessageSize: Int = 32 * 1024 * 1024 + ) async throws -> Data { + let lengthChunk = try await receiveExactly(count: 2, timeout: headerTimeout, phase: "read header") let length = Int(LibboxDecodeLengthChunk(lengthChunk)) - connection.receive(minimumIncompleteLength: length, maximumLength: length) { content, _, _, error in - if let error { - result = .failure(error) - } else { - result = .success(content!) - } - semaphore.signal() + guard length >= 0 else { + connection.cancel() + throw NWSocketError.invalidLength(length) } - semaphore.wait() - return try result.get() + guard length <= maxMessageSize else { + connection.cancel() + throw NWSocketError.messageTooLarge(length) + } + guard length > 0 else { + return Data() + } + return try await receiveExactly(count: length, timeout: bodyTimeout, phase: "read body") } - public func write(_ data: Data?) throws { + public func write(_ data: Data?, timeout: TimeInterval = 30) async throws { guard let data else { return } - let semaphore = DispatchSemaphore(value: 0) - var result: Error? - connection.send(content: LibboxEncodeChunkedMessage(data), isComplete: false, completion: .contentProcessed { error in - result = error - semaphore.wait() - }) - if let result { - throw result - } + try await sendAndAwait(content: LibboxEncodeChunkedMessage(data), timeout: timeout, phase: "write") } public func send(_ data: Data?) { @@ -60,4 +91,71 @@ public class NWSocket { public func cancel() { connection.cancel() } + + private func receiveExactly(count: Int, timeout: TimeInterval, phase: String) async throws -> Data { + guard count > 0 else { + return Data() + } + return try await withCheckedThrowingContinuation { continuation in + let oneShot = OneShot(continuation) + let timeoutItem: DispatchWorkItem? + if timeout > 0 { + timeoutItem = DispatchWorkItem { [connection] in + if oneShot.resume(.failure(NWSocketError.timeout(phase))) { + connection.cancel() + } + } + DispatchQueue.global().asyncAfter(deadline: .now() + timeout, execute: timeoutItem!) + } else { + timeoutItem = nil + } + + connection.receive(minimumIncompleteLength: count, maximumLength: count) { content, _, isComplete, error in + timeoutItem?.cancel() + if let error { + oneShot.resume(.failure(error)) + return + } + guard let content else { + oneShot.resume(.failure(NWSocketError.connectionClosed)) + return + } + guard content.count == count else { + if isComplete { + oneShot.resume(.failure(NWSocketError.connectionClosed)) + } else { + oneShot.resume(.failure(NWSocketError.invalidLength(content.count))) + } + return + } + oneShot.resume(.success(content)) + } + } + } + + private func sendAndAwait(content: Data?, timeout: TimeInterval, phase: String) async throws { + try await withCheckedThrowingContinuation { continuation in + let oneShot = OneShot(continuation) + let timeoutItem: DispatchWorkItem? + if timeout > 0 { + timeoutItem = DispatchWorkItem { [connection] in + if oneShot.resume(.failure(NWSocketError.timeout(phase))) { + connection.cancel() + } + } + DispatchQueue.global().asyncAfter(deadline: .now() + timeout, execute: timeoutItem!) + } else { + timeoutItem = nil + } + + connection.send(content: content, isComplete: false, completion: .contentProcessed { error in + timeoutItem?.cancel() + if let error { + oneShot.resume(.failure(error)) + } else { + oneShot.resume(.success(())) + } + }) + } + } } diff --git a/ApplicationLibrary/Service/ProfileServer.swift b/ApplicationLibrary/Service/ProfileServer.swift index 6d0242b..abc5fbd 100644 --- a/ApplicationLibrary/Service/ProfileServer.swift +++ b/ApplicationLibrary/Service/ProfileServer.swift @@ -55,17 +55,17 @@ public class ProfileServer { try await writeProfilePreviewList() } catch { NSLog("profile server: write profile list: \(error.localizedDescription)") - writeError(error.localizedDescription) + await writeError(error.localizedDescription) return } do { while true { - let message = try connection.read() - try processMessage(message) + let message = try await connection.read() + try await processMessage(message) } } catch { NSLog("profile server: process connection: \(error.localizedDescription)") - writeError(error.localizedDescription) + await writeError(error.localizedDescription) } } @@ -89,16 +89,14 @@ public class ProfileServer { } #endif - private func processMessage(_ data: Data) throws { + private func processMessage(_ data: Data) async throws { if data.count == 0 { return } let messageType = Int64(data[0]) switch messageType { case LibboxMessageTypeProfileContentRequest: - Task { - try await processProfileContentRequest(data) - } + try await processProfileContentRequest(data) default: throw NSError(domain: "ProfileServer", code: 0, userInfo: [NSLocalizedDescriptionKey: String(localized: "Unexpected message type \(messageType)")]) } @@ -136,7 +134,7 @@ public class ProfileServer { content.lastUpdated = Int64(lastUpdated.timeIntervalSince1970) } } - try connection.write(content.encode()) + try await connection.write(content.encode()) } private func writeProfilePreviewList() async throws { @@ -156,13 +154,13 @@ public class ProfileServer { } encoder.append(preview) } - try connection.write(encoder.encode()) + try await connection.write(encoder.encode()) } - private func writeError(_ message: String) { + private func writeError(_ message: String) async { let errorMessage = LibboxErrorMessage() errorMessage.message = message - try? connection.write(errorMessage.encode()) + try? await connection.write(errorMessage.encode()) } } } diff --git a/ApplicationLibrary/Views/Profile/ImportProfileViewModel.swift b/ApplicationLibrary/Views/Profile/ImportProfileViewModel.swift index 5978530..83f2c3a 100644 --- a/ApplicationLibrary/Views/Profile/ImportProfileViewModel.swift +++ b/ApplicationLibrary/Views/Profile/ImportProfileViewModel.swift @@ -59,10 +59,13 @@ var message: Data while true { do { - message = try socket.read() + message = try await socket.read() } catch { throw NSError(domain: "ImportProfileViewModel", code: 0, userInfo: [NSLocalizedDescriptionKey: String(localized: "Read from connection: \(error.localizedDescription)")]) } + if message.isEmpty { + continue + } var error: NSError? switch Int64(message[0]) { case LibboxMessageTypeError: @@ -112,12 +115,15 @@ connection.stateUpdateHandler = nil let request = LibboxProfileContentRequest() request.profileID = profileID - do { - try socket.write(request.encode()) - isImporting = true - } catch { - alert = AlertState(error: error) - reset() + isImporting = true + Task { + do { + try await socket.write(request.encode()) + } catch { + isImporting = false + alert = AlertState(error: error) + reset() + } } }