diff --git a/CHANGELOG.md b/CHANGELOG.md index 04579fbf00..3ec105d569 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -108,6 +108,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Vim Replace mode writing an invisible character for Backspace and keypad Enter. (#2717) - Option and Control chords editing text in Vim Normal and Visual mode, and Vim's Ctrl commands never running. (#2717) - Line and paragraph separators (U+2028, U+2029) shown as line breaks the database does not see. (#2717) +- Stop on Cloudflare D1, libSQL and Trino cancelling a sidebar read instead of the running query. - Numeric-looking filter values sent unquoted to text columns when a table first opens or after a foreign key jump. ## [0.73.0] - 2026-09-09 diff --git a/Packages/TableProCore/Sources/TableProTrinoCore/TrinoStatementClient.swift b/Packages/TableProCore/Sources/TableProTrinoCore/TrinoStatementClient.swift index f0358fe8fd..b4499bdd83 100644 --- a/Packages/TableProCore/Sources/TableProTrinoCore/TrinoStatementClient.swift +++ b/Packages/TableProCore/Sources/TableProTrinoCore/TrinoStatementClient.swift @@ -11,8 +11,7 @@ public final class TrinoStatementClient: @unchecked Sendable { private let config: TrinoClientConfig private let session: TrinoSessionState private let lock = NSLock() - private var _cancelled = false - private var _currentNextUri: String? + private var running: [ObjectIdentifier: TrinoRunningStatement] = [:] private static let maxTransientRetries = 5 private static let logger = Logger(subsystem: "com.TablePro", category: "TrinoStatementClient") @@ -62,14 +61,14 @@ public final class TrinoStatementClient: @unchecked Sendable { } } + /// Stops every statement running on this client, not only the one that started last: the app + /// runs sidebar and autocomplete reads on the same client while a query runs. Each statement + /// is told to stop, its request in flight is cancelled, and Trino gets one DELETE for it. public func cancel() { - let uri = lock.withLock { () -> String? in - _cancelled = true - return _currentNextUri - } - if let uri { - fireDelete(uri) - } + let statements = lock.withLock { Array(running.values) } + statements.forEach { $0.markCancelled() } + transport.cancelAll() + statements.compactMap { $0.claimRelease() }.forEach(fireDelete) } private struct StatementOutcome { @@ -83,16 +82,31 @@ public final class TrinoStatementClient: @unchecked Sendable { onColumns: ([TrinoColumn]) -> Void, onPage: ([[TrinoValue]]) -> Void ) async throws -> StatementOutcome { - lock.withLock { - _cancelled = false - _currentNextUri = nil - } guard let statementURL = config.statementURL else { throw TrinoError.invalidConfiguration("Invalid Trino server URL") } + let statement = TrinoRunningStatement() + lock.withLock { running[ObjectIdentifier(statement)] = statement } + defer { lock.withLock { running[ObjectIdentifier(statement)] = nil } } + + do { + return try await drive(statement, url: statementURL, sql: sql, onColumns: onColumns, onPage: onPage) + } catch { + releaseIfCancelled(statement) + throw error + } + } + private func drive( + _ statement: TrinoRunningStatement, + url statementURL: URL, + sql: String, + onColumns: ([TrinoColumn]) -> Void, + onPage: ([[TrinoValue]]) -> Void + ) async throws -> StatementOutcome { var httpResponse = try await sendWithRetry( - makeRequest(method: .post, url: statementURL, headers: initialHeaders(), body: Data(sql.utf8)) + makeRequest(method: .post, url: statementURL, headers: initialHeaders(), body: Data(sql.utf8)), + for: statement ) var results = try decode(httpResponse) session.apply(responseHeaders: httpResponse.headers, protocolHeaders: config.protocolHeaders) @@ -113,12 +127,14 @@ public final class TrinoStatementClient: @unchecked Sendable { var nextUri = results.nextUri while let uri = nextUri { - try abortIfCancelled(currentUri: uri) - lock.withLock { _currentNextUri = uri } + statement.advance(to: uri) guard let nextURL = URL(string: uri) else { throw TrinoError.invalidResponse("Trino returned an invalid nextUri") } - httpResponse = try await sendWithRetry(makeRequest(method: .get, url: nextURL, headers: followHeaders())) + httpResponse = try await sendWithRetry( + makeRequest(method: .get, url: nextURL, headers: followHeaders()), + for: statement + ) results = try decode(httpResponse) session.apply(responseHeaders: httpResponse.headers, protocolHeaders: config.protocolHeaders) if let error = results.error { @@ -141,7 +157,6 @@ public final class TrinoStatementClient: @unchecked Sendable { nextUri = results.nextUri } - lock.withLock { _currentNextUri = nil } return StatementOutcome(updateType: updateType, updateCount: updateCount, queryId: queryId) } @@ -167,16 +182,23 @@ public final class TrinoStatementClient: @unchecked Sendable { } } - private func abortIfCancelled(currentUri: String) throws { - let cancelled = lock.withLock { _cancelled } || Task.isCancelled - guard cancelled else { return } - fireDelete(currentUri) + private func abortIfCancelled(_ statement: TrinoRunningStatement) throws { + guard statement.isCancelled || Task.isCancelled else { return } throw TrinoError.cancelled } - private func sendWithRetry(_ request: TrinoHTTPRequest) async throws -> TrinoHTTPResponse { + private func releaseIfCancelled(_ statement: TrinoRunningStatement) { + guard statement.isCancelled || Task.isCancelled, let uri = statement.claimRelease() else { return } + fireDelete(uri) + } + + private func sendWithRetry( + _ request: TrinoHTTPRequest, + for statement: TrinoRunningStatement + ) async throws -> TrinoHTTPResponse { var attempt = 0 while true { + try abortIfCancelled(statement) let response = try await transport.send(request) switch response.statusCode { case 200...299: @@ -299,3 +321,30 @@ public final class TrinoStatementClient: @unchecked Sendable { String(data: response.body, encoding: .utf8)?.trimmingCharacters(in: .whitespacesAndNewlines) ?? "" } } + +private final class TrinoRunningStatement: @unchecked Sendable { + private let lock = NSLock() + private var cancelled = false + private var released = false + private var nextUri: String? + + var isCancelled: Bool { + lock.withLock { cancelled } + } + + func markCancelled() { + lock.withLock { cancelled = true } + } + + func advance(to uri: String) { + lock.withLock { nextUri = uri } + } + + func claimRelease() -> String? { + lock.withLock { + guard !released, let nextUri else { return nil } + released = true + return nextUri + } + } +} diff --git a/Packages/TableProCore/Sources/TableProTrinoCore/TrinoTransport.swift b/Packages/TableProCore/Sources/TableProTrinoCore/TrinoTransport.swift index 0efb260cca..0986725543 100644 --- a/Packages/TableProCore/Sources/TableProTrinoCore/TrinoTransport.swift +++ b/Packages/TableProCore/Sources/TableProTrinoCore/TrinoTransport.swift @@ -75,19 +75,37 @@ public struct TrinoHTTPResponse: Sendable { public protocol TrinoTransport: Sendable { func send(_ request: TrinoHTTPRequest) async throws -> TrinoHTTPResponse + func cancelAll() } +/// Sends with `URLSession.data(for:delegate:)`, so cancelling the Swift task that awaits a request +/// cancels its URL task, and keeps every request in flight so `cancelAll` stops each one. A DELETE +/// is never tracked: it is how a statement tells Trino to stop, and a cancel must not cancel it. public final class URLSessionTrinoTransport: NSObject, TrinoTransport, @unchecked Sendable { private let session: URLSession + private let lock = NSLock() + private var inFlight: [ObjectIdentifier: URLSessionTask] = [:] - public init(tls: TrinoTLSOptions) { - let configuration = URLSessionConfiguration.ephemeral + public convenience init(tls: TrinoTLSOptions) { + self.init(tls: tls, configuration: .ephemeral) + } + + init(tls: TrinoTLSOptions, configuration: URLSessionConfiguration) { configuration.requestCachePolicy = .reloadIgnoringLocalCacheData let delegateProxy = TrinoTLSDelegate(tls: tls) self.session = URLSession(configuration: configuration, delegate: delegateProxy, delegateQueue: nil) super.init() } + deinit { + session.invalidateAndCancel() + } + + public func cancelAll() { + let tasks = lock.withLock { Array(inFlight.values) } + tasks.forEach { $0.cancel() } + } + public func send(_ request: TrinoHTTPRequest) async throws -> TrinoHTTPResponse { var urlRequest = URLRequest(url: request.url) urlRequest.httpMethod = request.method.rawValue @@ -97,24 +115,18 @@ public final class URLSessionTrinoTransport: NSObject, TrinoTransport, @unchecke urlRequest.setValue(value, forHTTPHeaderField: name) } - let (data, response) = try await withCheckedThrowingContinuation { - (continuation: CheckedContinuation<(Data, URLResponse), Error>) in - let task = session.dataTask(with: urlRequest) { data, response, error in - if let error { - if (error as? URLError)?.code == .cancelled { - continuation.resume(throwing: TrinoError.cancelled) - } else { - continuation.resume(throwing: TrinoError.transport(error.localizedDescription)) - } - return - } - guard let data, let response else { - continuation.resume(throwing: TrinoError.invalidResponse("Empty response from Trino")) - return - } - continuation.resume(returning: (data, response)) - } - task.resume() + let tracker = request.method == .delete ? nil : TrinoTaskTracker(transport: self) + defer { tracker?.finish() } + let data: Data + let response: URLResponse + do { + (data, response) = try await session.data(for: urlRequest, delegate: tracker) + } catch let error as URLError where error.code == .cancelled { + throw TrinoError.cancelled + } catch is CancellationError { + throw TrinoError.cancelled + } catch { + throw TrinoError.transport(error.localizedDescription) } guard let httpResponse = response as? HTTPURLResponse else { @@ -126,6 +138,38 @@ public final class URLSessionTrinoTransport: NSObject, TrinoTransport, @unchecke body: data ) } + + var inFlightCount: Int { + lock.withLock { inFlight.count } + } + + fileprivate func register(_ task: URLSessionTask) { + lock.withLock { inFlight[ObjectIdentifier(task)] = task } + } + + fileprivate func unregister(_ task: URLSessionTask) { + lock.withLock { inFlight[ObjectIdentifier(task)] = nil } + } +} + +private final class TrinoTaskTracker: NSObject, URLSessionTaskDelegate, @unchecked Sendable { + private weak var transport: URLSessionTrinoTransport? + private let lock = NSLock() + private var task: URLSessionTask? + + init(transport: URLSessionTrinoTransport) { + self.transport = transport + } + + func urlSession(_ session: URLSession, didCreateTask task: URLSessionTask) { + lock.withLock { self.task = task } + transport?.register(task) + } + + func finish() { + guard let task = lock.withLock({ task }) else { return } + transport?.unregister(task) + } } private final class TrinoTLSDelegate: NSObject, URLSessionDelegate { diff --git a/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoStatementClientTests.swift b/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoStatementClientTests.swift index b4d6dd626d..ddc40e129c 100644 --- a/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoStatementClientTests.swift +++ b/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoStatementClientTests.swift @@ -1,5 +1,5 @@ -import XCTest @testable import TableProTrinoCore +import XCTest final class TrinoStatementClientTests: XCTestCase { private func makeClient(_ transport: StubTransport, session: TrinoSessionState = TrinoSessionState(catalog: "c", schema: "s")) -> TrinoStatementClient { @@ -136,5 +136,6 @@ final class TrinoStatementClientTests: XCTestCase { } catch let error as TrinoError { XCTAssertEqual(error, .cancelled) } + XCTAssertEqual(transport.cancelAllCount, 1) } } diff --git a/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoTestSupport.swift b/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoTestSupport.swift index a9a2965a34..7721d0c134 100644 --- a/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoTestSupport.swift +++ b/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoTestSupport.swift @@ -11,6 +11,7 @@ final class StubTransport: TrinoTransport, @unchecked Sendable { private let lock = NSLock() private var queue: [Canned] private var recorded: [TrinoHTTPRequest] = [] + private var cancelAllCalls = 0 var onSend: ((TrinoHTTPRequest, Int) -> Void)? init(_ responses: [Canned]) { @@ -21,6 +22,14 @@ final class StubTransport: TrinoTransport, @unchecked Sendable { lock.withLock { recorded } } + var cancelAllCount: Int { + lock.withLock { cancelAllCalls } + } + + func cancelAll() { + lock.withLock { cancelAllCalls += 1 } + } + func send(_ request: TrinoHTTPRequest) async throws -> TrinoHTTPResponse { let index = lock.withLock { () -> Int in recorded.append(request) diff --git a/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoURLSessionTransportTests.swift b/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoURLSessionTransportTests.swift new file mode 100644 index 0000000000..29171c8c32 --- /dev/null +++ b/Packages/TableProCore/Tests/TableProTrinoCoreTests/TrinoURLSessionTransportTests.swift @@ -0,0 +1,229 @@ +import Foundation +@testable import TableProTrinoCore +import Testing + +/// Plays a Trino coordinator: a statement POST answers at once, `SELECT hang` hands back a nextUri +/// whose GET never answers until cancelled, and every DELETE is recorded, so a test can hold +/// several statements in flight and watch what a cancel does to each. +private final class TrinoStubProtocol: URLProtocol, @unchecked Sendable { + static let origin = "https://h:8443" + + private static let lock = NSLock() + nonisolated(unsafe) private static var startedRequests: [String] = [] + nonisolated(unsafe) private static var queryCount = 0 + nonisolated(unsafe) private static var timeout: TimeInterval? + + static func reset() { + lock.withLock { + startedRequests = [] + queryCount = 0 + timeout = nil + } + } + + static var started: [String] { + lock.withLock { startedRequests } + } + + static var lastTimeout: TimeInterval? { + lock.withLock { timeout } + } + + static var parkedPolls: [String] { + started.filter { $0.hasPrefix("GET /v1/statement/executing/") }.map { String($0.dropFirst("GET ".count)) } + } + + static var deletes: [String] { + started.filter { $0.hasPrefix("DELETE ") }.map { String($0.dropFirst("DELETE ".count)) } + } + + override class func canInit(with request: URLRequest) -> Bool { true } + override class func canonicalRequest(for request: URLRequest) -> URLRequest { request } + + override func startLoading() { + let method = request.httpMethod ?? "GET" + let path = request.url?.path ?? "" + Self.lock.withLock { + Self.startedRequests.append("\(method) \(path)") + Self.timeout = request.timeoutInterval + } + + switch (method, path) { + case ("POST", "/v1/statement"): + respond(body: statementReply(sql: Self.bodyText(of: request))) + case ("DELETE", "/slow"): + DispatchQueue.global().asyncAfter(deadline: .now() + 0.3) { self.respond(status: 204, body: "") } + case ("DELETE", _): + respond(status: 204, body: "") + case (_, "/ok"): + respond(body: #"{"id":"ok"}"#, headers: ["X-Trino-Set-Schema": "analytics"]) + default: + return + } + } + + override func stopLoading() {} + + private func statementReply(sql: String) -> String { + let queryId = Self.lock.withLock { () -> String in + Self.queryCount += 1 + return "q\(Self.queryCount)" + } + guard sql == "SELECT hang" else { + let column = #"{"name":"n","type":"bigint","typeSignature":{"rawType":"bigint"}}"# + return #"{"id":"\#(queryId)","columns":[\#(column)],"data":[[1]]}"# + } + return #"{"id":"\#(queryId)","nextUri":"\#(Self.origin)/v1/statement/executing/\#(queryId)/1"}"# + } + + private func respond(status: Int = 200, body: String, headers: [String: String] = [:]) { + guard let url = request.url, + let response = HTTPURLResponse(url: url, statusCode: status, httpVersion: nil, headerFields: headers) + else { return } + client?.urlProtocol(self, didReceive: response, cacheStoragePolicy: .notAllowed) + client?.urlProtocol(self, didLoad: Data(body.utf8)) + client?.urlProtocolDidFinishLoading(self) + } + + private static func bodyText(of request: URLRequest) -> String { + if let body = request.httpBody { + return String(bytes: body, encoding: .utf8) ?? "" + } + guard let stream = request.httpBodyStream else { return "" } + stream.open() + defer { stream.close() } + var data = Data() + var buffer = [UInt8](repeating: 0, count: 1_024) + while stream.hasBytesAvailable { + let read = stream.read(&buffer, maxLength: buffer.count) + guard read > 0 else { break } + data.append(buffer, count: read) + } + return String(bytes: data, encoding: .utf8) ?? "" + } +} + +@Suite("Trino URLSession transport", .serialized) +struct TrinoURLSessionTransportTests { + init() { + TrinoStubProtocol.reset() + } + + private func transport() -> URLSessionTrinoTransport { + let configuration = URLSessionConfiguration.ephemeral + configuration.protocolClasses = [TrinoStubProtocol.self] + return URLSessionTrinoTransport(tls: .systemDefault, configuration: configuration) + } + + private func client(_ transport: URLSessionTrinoTransport) -> TrinoStatementClient { + TrinoStatementClient( + transport: transport, + config: TrinoClientConfig(host: "h", port: 8_443, useTLS: true, user: "u"), + session: TrinoSessionState(catalog: "c", schema: "s") + ) + } + + private func request(_ method: TrinoHTTPRequest.Method, _ path: String, timeout: Int = 60) throws -> TrinoHTTPRequest { + TrinoHTTPRequest( + method: method, + url: try #require(URL(string: TrinoStubProtocol.origin + path)), + headers: [:], + timeoutSeconds: timeout + ) + } + + private func waitUntil(_ condition: () -> Bool) async { + for _ in 0 ..< 300 where !condition() { + try? await Task.sleep(for: .milliseconds(10)) + } + } + + @Test("A request carries its own timeout and returns the status, headers and body") + func roundTrip() async throws { + let response = try await transport().send(try request(.get, "/ok", timeout: 330)) + + #expect(response.statusCode == 200) + #expect(response.headers.first("x-trino-set-schema") == "analytics") + #expect(String(bytes: response.body, encoding: .utf8) == #"{"id":"ok"}"#) + #expect(TrinoStubProtocol.lastTimeout == 330) + } + + @Test("Cancelling everything stops every request in flight, not just the latest") + func cancelAllStopsEveryRequest() async throws { + let transport = transport() + let firstRequest = try request(.get, "/hang/1") + let secondRequest = try request(.get, "/hang/2") + let first = Task { try await transport.send(firstRequest) } + let second = Task { try await transport.send(secondRequest) } + await waitUntil { transport.inFlightCount == 2 } + + transport.cancelAll() + + for task in [first, second] { + await #expect(throws: TrinoError.cancelled) { try await task.value } + } + #expect(transport.inFlightCount == 0) + } + + @Test("Cancelling the awaiting task cancels its request") + func taskCancellation() async throws { + let transport = transport() + let hanging = try request(.get, "/hang") + let task = Task { try await transport.send(hanging) } + await waitUntil { transport.inFlightCount == 1 } + + task.cancel() + + await #expect(throws: TrinoError.cancelled) { try await task.value } + #expect(transport.inFlightCount == 0) + } + + @Test("Cancelling everything leaves a DELETE to finish telling Trino to stop") + func cancelAllSparesDelete() async throws { + let transport = transport() + let deleteRequest = try request(.delete, "/slow") + let pollRequest = try request(.get, "/hang") + let delete = Task { try await transport.send(deleteRequest) } + let poll = Task { try await transport.send(pollRequest) } + await waitUntil { transport.inFlightCount == 1 && TrinoStubProtocol.started.contains("DELETE /slow") } + + transport.cancelAll() + + await #expect(throws: TrinoError.cancelled) { try await poll.value } + #expect(try await delete.value.statusCode == 204) + } + + @Test("Cancel stops every running statement and deletes each on the server") + func cancelStopsEveryStatement() async throws { + let transport = transport() + let client = client(transport) + let first = Task { try await client.execute("SELECT hang") } + let second = Task { try await client.execute("SELECT hang") } + await waitUntil { TrinoStubProtocol.parkedPolls.count == 2 } + let parked = TrinoStubProtocol.parkedPolls + + let finished = try await client.execute("SELECT 1") + client.cancel() + + #expect(finished.rows == [[.text("1")]]) + for task in [first, second] { + await #expect(throws: TrinoError.cancelled) { try await task.value } + } + await waitUntil { TrinoStubProtocol.deletes.count == 2 } + #expect(Set(TrinoStubProtocol.deletes) == Set(parked)) + } + + @Test("Cancelling the task running a statement cancels its poll and deletes it on the server") + func taskCancellationReleasesStatement() async throws { + let client = client(transport()) + let task = Task { try await client.execute("SELECT hang") } + await waitUntil { TrinoStubProtocol.parkedPolls.count == 1 } + let parked = TrinoStubProtocol.parkedPolls + + task.cancel() + + await #expect(throws: TrinoError.cancelled) { try await task.value } + await waitUntil { TrinoStubProtocol.deletes.count == 1 } + #expect(TrinoStubProtocol.deletes == parked) + } +} diff --git a/Plugins/CloudflareD1DriverPlugin/CloudflareD1PluginDriver.swift b/Plugins/CloudflareD1DriverPlugin/CloudflareD1PluginDriver.swift index 2f56c8c9e0..a5f92e67d6 100644 --- a/Plugins/CloudflareD1DriverPlugin/CloudflareD1PluginDriver.swift +++ b/Plugins/CloudflareD1DriverPlugin/CloudflareD1PluginDriver.swift @@ -169,14 +169,14 @@ final class CloudflareD1PluginDriver: PluginDatabaseDriver, @unchecked Sendable let payloads = try await client.executeBatchRaw(statements: statements) let elapsed = Date().timeIntervalSince(startTime) - return payloads.enumerated().map { _, payload in + return payloads.map { payload in mapRawResult(payload, executionTime: payload.meta?.duration ?? (elapsed / Double(payloads.count))) } } func cancelQuery() throws { lock.lock() - httpClient?.cancelCurrentTask() + httpClient?.cancelAll() lock.unlock() } @@ -188,7 +188,7 @@ final class CloudflareD1PluginDriver: PluginDatabaseDriver, @unchecked Sendable // MARK: - Streaming func streamRows(query: String) -> AsyncThrowingStream { - return AsyncThrowingStream(bufferingPolicy: .unbounded) { continuation in + AsyncThrowingStream(bufferingPolicy: .unbounded) { continuation in let streamTask = Task { do { try await self.performStreamRows(query: query, continuation: continuation) diff --git a/Plugins/CloudflareD1DriverPlugin/D1HttpClient.swift b/Plugins/CloudflareD1DriverPlugin/D1HttpClient.swift index b1ea4d1938..fa82b5c313 100644 --- a/Plugins/CloudflareD1DriverPlugin/D1HttpClient.swift +++ b/Plugins/CloudflareD1DriverPlugin/D1HttpClient.swift @@ -128,6 +128,10 @@ enum D1Value: Decodable { // MARK: - HTTP Client +/// Sends with `URLSession.data(for:delegate:)`, so cancelling the Swift task that awaits a request +/// cancels its URL task, and keeps every request in flight so `cancelAll` stops each one. The app +/// runs sidebar and autocomplete reads on this client while a query runs, so a single stored task +/// handle cancelled whichever request started last instead of the query the user stopped. final class D1HttpClient: @unchecked Sendable { private static let logger = Logger(subsystem: "com.TablePro", category: "D1HttpClient") @@ -136,7 +140,7 @@ final class D1HttpClient: @unchecked Sendable { private let lock = NSLock() private var _databaseId: String private var session: URLSession? - private var currentTask: URLSessionDataTask? + private var inFlight: [ObjectIdentifier: URLSessionTask] = [:] private let queryTimeout = HttpQueryTimeoutBox() var databaseId: String { @@ -162,30 +166,37 @@ final class D1HttpClient: @unchecked Sendable { queryTimeout.set(serverTimeoutSeconds: seconds) } - func createSession() { - let config = URLSessionConfiguration.default - config.timeoutIntervalForRequest = HttpQueryTimeout.sessionBootstrapRequestTimeout - config.timeoutIntervalForResource = HttpQueryTimeout.sessionResourceTimeout + func createSession(configuration: URLSessionConfiguration = .default) { + configuration.timeoutIntervalForRequest = HttpQueryTimeout.sessionBootstrapRequestTimeout + configuration.timeoutIntervalForResource = HttpQueryTimeout.sessionResourceTimeout lock.lock() - session = URLSession(configuration: config) + session = URLSession(configuration: configuration) lock.unlock() } func invalidateSession() { lock.lock() - currentTask?.cancel() - currentTask = nil session?.invalidateAndCancel() session = nil lock.unlock() } - func cancelCurrentTask() { - lock.lock() - currentTask?.cancel() - currentTask = nil - lock.unlock() + func cancelAll() { + let tasks = lock.withLock { Array(inFlight.values) } + tasks.forEach { $0.cancel() } + } + + var inFlightCount: Int { + lock.withLock { inFlight.count } + } + + fileprivate func register(_ task: URLSessionTask) { + lock.withLock { inFlight[ObjectIdentifier(task)] = task } + } + + fileprivate func unregister(_ task: URLSessionTask) { + lock.withLock { inFlight[ObjectIdentifier(task)] = nil } } // MARK: - API Methods @@ -315,38 +326,16 @@ final class D1HttpClient: @unchecked Sendable { request.setValue("application/json", forHTTPHeaderField: "Content-Type") request.httpBody = body - let (data, response) = try await withTaskCancellationHandler { - try await withCheckedThrowingContinuation { - (continuation: CheckedContinuation<(Data, URLResponse), Error>) in - let task = session.dataTask(with: request) { data, response, error in - if let error { - continuation.resume(throwing: error) - return - } - guard let data, let response else { - continuation.resume( - throwing: D1HttpError(message: "Empty response from server") - ) - return - } - continuation.resume(returning: (data, response)) - } - - self.lock.lock() - self.currentTask = task - self.lock.unlock() - - task.resume() - } - } onCancel: { - self.lock.lock() - self.currentTask?.cancel() - self.currentTask = nil - self.lock.unlock() + let tracker = D1TaskTracker(client: self) + defer { tracker.finish() } + let data: Data + let response: URLResponse + do { + (data, response) = try await session.data(for: request, delegate: tracker) + } catch let error as URLError where error.code == .cancelled { + throw CancellationError() } - lock.withLock { currentTask = nil } - guard let httpResponse = response as? HTTPURLResponse else { throw D1HttpError(message: "Invalid response from server") } @@ -401,6 +390,26 @@ final class D1HttpClient: @unchecked Sendable { } } +private final class D1TaskTracker: NSObject, URLSessionTaskDelegate, @unchecked Sendable { + private weak var client: D1HttpClient? + private let lock = NSLock() + private var task: URLSessionTask? + + init(client: D1HttpClient) { + self.client = client + } + + func urlSession(_ session: URLSession, didCreateTask task: URLSessionTask) { + lock.withLock { self.task = task } + client?.register(task) + } + + func finish() { + guard let task = lock.withLock({ task }) else { return } + client?.unregister(task) + } +} + // MARK: - Error struct D1HttpError: Error, LocalizedError { diff --git a/Plugins/LibSQLDriverPlugin/HranaHttpClient.swift b/Plugins/LibSQLDriverPlugin/HranaHttpClient.swift index ffe2a6b7cd..d4f84bc1b5 100644 --- a/Plugins/LibSQLDriverPlugin/HranaHttpClient.swift +++ b/Plugins/LibSQLDriverPlugin/HranaHttpClient.swift @@ -107,6 +107,10 @@ struct HranaErrorDetail: Decodable { // MARK: - HTTP Client +/// Sends with `URLSession.data(for:delegate:)`, so cancelling the Swift task that awaits a request +/// cancels its URL task, and keeps every request in flight so `cancelAll` stops each one. The app +/// runs sidebar and autocomplete reads on this client while a query runs, so a single stored task +/// handle cancelled whichever request started last instead of the query the user stopped. final class HranaHttpClient: @unchecked Sendable { private static let logger = Logger(subsystem: "com.TablePro", category: "HranaHttpClient") @@ -114,7 +118,7 @@ final class HranaHttpClient: @unchecked Sendable { private let authToken: String? private let lock = NSLock() private var session: URLSession? - private var currentTask: URLSessionDataTask? + private var inFlight: [ObjectIdentifier: URLSessionTask] = [:] private let queryTimeout = HttpQueryTimeoutBox() init(baseUrl: URL, authToken: String?) { @@ -126,30 +130,37 @@ final class HranaHttpClient: @unchecked Sendable { queryTimeout.set(serverTimeoutSeconds: seconds) } - func createSession() { - let config = URLSessionConfiguration.default - config.timeoutIntervalForRequest = HttpQueryTimeout.sessionBootstrapRequestTimeout - config.timeoutIntervalForResource = HttpQueryTimeout.sessionResourceTimeout + func createSession(configuration: URLSessionConfiguration = .default) { + configuration.timeoutIntervalForRequest = HttpQueryTimeout.sessionBootstrapRequestTimeout + configuration.timeoutIntervalForResource = HttpQueryTimeout.sessionResourceTimeout lock.lock() - session = URLSession(configuration: config) + session = URLSession(configuration: configuration) lock.unlock() } func invalidateSession() { lock.lock() - currentTask?.cancel() - currentTask = nil session?.invalidateAndCancel() session = nil lock.unlock() } - func cancelCurrentTask() { - lock.lock() - currentTask?.cancel() - currentTask = nil - lock.unlock() + func cancelAll() { + let tasks = lock.withLock { Array(inFlight.values) } + tasks.forEach { $0.cancel() } + } + + var inFlightCount: Int { + lock.withLock { inFlight.count } + } + + fileprivate func register(_ task: URLSessionTask) { + lock.withLock { inFlight[ObjectIdentifier(task)] = task } + } + + fileprivate func unregister(_ task: URLSessionTask) { + lock.withLock { inFlight[ObjectIdentifier(task)] = nil } } // MARK: - API Methods @@ -222,38 +233,16 @@ final class HranaHttpClient: @unchecked Sendable { } request.httpBody = body - let (data, response) = try await withTaskCancellationHandler { - try await withCheckedThrowingContinuation { - (continuation: CheckedContinuation<(Data, URLResponse), Error>) in - let task = session.dataTask(with: request) { data, response, error in - if let error { - continuation.resume(throwing: error) - return - } - guard let data, let response else { - continuation.resume( - throwing: HranaHttpError(message: "Empty response from server") - ) - return - } - continuation.resume(returning: (data, response)) - } - - self.lock.lock() - self.currentTask = task - self.lock.unlock() - - task.resume() - } - } onCancel: { - self.lock.lock() - self.currentTask?.cancel() - self.currentTask = nil - self.lock.unlock() + let tracker = HranaTaskTracker(client: self) + defer { tracker.finish() } + let data: Data + let response: URLResponse + do { + (data, response) = try await session.data(for: request, delegate: tracker) + } catch let error as URLError where error.code == .cancelled { + throw CancellationError() } - lock.withLock { currentTask = nil } - guard let httpResponse = response as? HTTPURLResponse else { throw HranaHttpError(message: "Invalid response from server") } @@ -317,6 +306,26 @@ final class HranaHttpClient: @unchecked Sendable { } } +private final class HranaTaskTracker: NSObject, URLSessionTaskDelegate, @unchecked Sendable { + private weak var client: HranaHttpClient? + private let lock = NSLock() + private var task: URLSessionTask? + + init(client: HranaHttpClient) { + self.client = client + } + + func urlSession(_ session: URLSession, didCreateTask task: URLSessionTask) { + lock.withLock { self.task = task } + client?.register(task) + } + + func finish() { + guard let task = lock.withLock({ task }) else { return } + client?.unregister(task) + } +} + // MARK: - Error struct HranaHttpError: Error, LocalizedError { diff --git a/Plugins/LibSQLDriverPlugin/LibSQLPluginDriver.swift b/Plugins/LibSQLDriverPlugin/LibSQLPluginDriver.swift index b164cfed80..106032aaaf 100644 --- a/Plugins/LibSQLDriverPlugin/LibSQLPluginDriver.swift +++ b/Plugins/LibSQLDriverPlugin/LibSQLPluginDriver.swift @@ -240,7 +240,7 @@ final class LibSQLPluginDriver: PluginDatabaseDriver, @unchecked Sendable { lock.unlock() if case .remote(let client) = current { - client.cancelCurrentTask() + client.cancelAll() } } @@ -259,7 +259,7 @@ final class LibSQLPluginDriver: PluginDatabaseDriver, @unchecked Sendable { // MARK: - Streaming func streamRows(query: String) -> AsyncThrowingStream { - return AsyncThrowingStream(bufferingPolicy: .unbounded) { continuation in + AsyncThrowingStream(bufferingPolicy: .unbounded) { continuation in let streamTask = Task { do { try await self.performStreamRows(query: query, continuation: continuation) diff --git a/TableProTests/Core/CloudflareD1/D1HttpClientCancellationTests.swift b/TableProTests/Core/CloudflareD1/D1HttpClientCancellationTests.swift new file mode 100644 index 0000000000..1bb46a6921 --- /dev/null +++ b/TableProTests/Core/CloudflareD1/D1HttpClientCancellationTests.swift @@ -0,0 +1,129 @@ +import Foundation +import TableProPluginKit +import Testing + +/// Plays the D1 REST API: a statement whose SQL is `SELECT hang` never answers until cancelled, +/// any other statement answers at once, so a test can hold several requests in flight and watch +/// what a cancel does to each. +private final class D1StubProtocol: URLProtocol, @unchecked Sendable { + private static let lock = NSLock() + nonisolated(unsafe) private static var lastRequest: URLRequest? + + static var recorded: URLRequest? { + lock.withLock { lastRequest } + } + + override class func canInit(with request: URLRequest) -> Bool { true } + override class func canonicalRequest(for request: URLRequest) -> URLRequest { request } + + override func startLoading() { + Self.lock.withLock { Self.lastRequest = request } + guard !Self.bodyText(of: request).contains("SELECT hang"), + let url = request.url, + let response = HTTPURLResponse(url: url, statusCode: 200, httpVersion: nil, headerFields: nil) + else { return } + let body = #"{"result":[{"results":{"columns":["n"],"rows":[[1]]},"success":true}],"success":true}"# + client?.urlProtocol(self, didReceive: response, cacheStoragePolicy: .notAllowed) + client?.urlProtocol(self, didLoad: Data(body.utf8)) + client?.urlProtocolDidFinishLoading(self) + } + + override func stopLoading() {} + + private static func bodyText(of request: URLRequest) -> String { + guard let stream = request.httpBodyStream else { + return String(bytes: request.httpBody ?? Data(), encoding: .utf8) ?? "" + } + stream.open() + defer { stream.close() } + var data = Data() + var buffer = [UInt8](repeating: 0, count: 1_024) + while stream.hasBytesAvailable { + let read = stream.read(&buffer, maxLength: buffer.count) + guard read > 0 else { break } + data.append(buffer, count: read) + } + return String(bytes: data, encoding: .utf8) ?? "" + } +} + +@Suite("Cloudflare D1 HTTP client cancellation", .serialized) +struct D1HttpClientCancellationTests { + private func connectedClient() -> D1HttpClient { + let client = D1HttpClient(accountId: "account", apiToken: "token", databaseId: "database") + let configuration = URLSessionConfiguration.ephemeral + configuration.protocolClasses = [D1StubProtocol.self] + client.createSession(configuration: configuration) + return client + } + + private func waitUntil(_ condition: () -> Bool) async { + for _ in 0 ..< 300 where !condition() { + try? await Task.sleep(for: .milliseconds(10)) + } + } + + @Test("A statement carries the token and the query timeout and returns its rows") + func roundTrip() async throws { + let client = connectedClient() + client.setQueryTimeout(300) + + let payload = try await client.executeRaw(sql: "SELECT 1") + + #expect(payload.results.columns == ["n"]) + #expect(payload.results.rows?.first?.first?.stringValue == "1") + let request = try #require(D1StubProtocol.recorded) + #expect(request.value(forHTTPHeaderField: "Authorization") == "Bearer token") + #expect(request.timeoutInterval == HttpQueryTimeout(serverTimeoutSeconds: 300).requestTimeoutInterval) + #expect(client.inFlightCount == 0) + } + + @Test("Cancelling everything stops every request in flight, not just the latest") + func cancelAllStopsEveryRequest() async throws { + let client = connectedClient() + let query = Task { try await client.executeRaw(sql: "SELECT hang") } + let sidebarRead = Task { try await client.executeRaw(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 2 } + + client.cancelAll() + + await #expect(throws: CancellationError.self) { try await query.value } + await #expect(throws: CancellationError.self) { try await sidebarRead.value } + #expect(client.inFlightCount == 0) + } + + @Test("A request that finishes leaves the others cancellable") + func finishedRequestKeepsOthersTracked() async throws { + let client = connectedClient() + let query = Task { try await client.executeRaw(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 1 } + + _ = try await client.executeRaw(sql: "SELECT 1") + client.cancelAll() + + await #expect(throws: CancellationError.self) { try await query.value } + } + + @Test("Cancelling the awaiting task cancels its request") + func taskCancellation() async throws { + let client = connectedClient() + let query = Task { try await client.executeRaw(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 1 } + + query.cancel() + + await #expect(throws: CancellationError.self) { try await query.value } + #expect(client.inFlightCount == 0) + } + + @Test("Disconnecting cancels every request in flight") + func invalidateSessionCancels() async throws { + let client = connectedClient() + let query = Task { try await client.executeRaw(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 1 } + + client.invalidateSession() + + await #expect(throws: CancellationError.self) { try await query.value } + } +} diff --git a/TableProTests/Plugins/HranaHttpClientCancellationTests.swift b/TableProTests/Plugins/HranaHttpClientCancellationTests.swift new file mode 100644 index 0000000000..79e47233bb --- /dev/null +++ b/TableProTests/Plugins/HranaHttpClientCancellationTests.swift @@ -0,0 +1,133 @@ +import Foundation +import TableProPluginKit +import Testing + +/// Plays a Hrana pipeline endpoint: a statement whose SQL is `SELECT hang` never answers until +/// cancelled, any other statement answers at once, so a test can hold several requests in flight +/// and watch what a cancel does to each. +private final class HranaStubProtocol: URLProtocol, @unchecked Sendable { + private static let lock = NSLock() + nonisolated(unsafe) private static var lastRequest: URLRequest? + + static var recorded: URLRequest? { + lock.withLock { lastRequest } + } + + override class func canInit(with request: URLRequest) -> Bool { true } + override class func canonicalRequest(for request: URLRequest) -> URLRequest { request } + + override func startLoading() { + Self.lock.withLock { Self.lastRequest = request } + guard !Self.bodyText(of: request).contains("SELECT hang"), + let url = request.url, + let response = HTTPURLResponse(url: url, statusCode: 200, httpVersion: nil, headerFields: nil) + else { return } + let row = #"[{"type":"integer","value":"1"}]"# + let result = #"{"cols":[{"name":"n","decltype":"INTEGER"}],"rows":[\#(row)],"affected_row_count":0}"# + let body = #"{"results":[{"type":"ok","response":{"type":"execute","result":\#(result)}}]}"# + client?.urlProtocol(self, didReceive: response, cacheStoragePolicy: .notAllowed) + client?.urlProtocol(self, didLoad: Data(body.utf8)) + client?.urlProtocolDidFinishLoading(self) + } + + override func stopLoading() {} + + private static func bodyText(of request: URLRequest) -> String { + guard let stream = request.httpBodyStream else { + return String(bytes: request.httpBody ?? Data(), encoding: .utf8) ?? "" + } + stream.open() + defer { stream.close() } + var data = Data() + var buffer = [UInt8](repeating: 0, count: 1_024) + while stream.hasBytesAvailable { + let read = stream.read(&buffer, maxLength: buffer.count) + guard read > 0 else { break } + data.append(buffer, count: read) + } + return String(bytes: data, encoding: .utf8) ?? "" + } +} + +@Suite("libSQL Hrana HTTP client cancellation", .serialized) +struct HranaHttpClientCancellationTests { + private func connectedClient() throws -> HranaHttpClient { + let url = try #require(URL(string: "https://db.turso.test")) + let client = HranaHttpClient(baseUrl: url, authToken: "token") + let configuration = URLSessionConfiguration.ephemeral + configuration.protocolClasses = [HranaStubProtocol.self] + client.createSession(configuration: configuration) + return client + } + + private func waitUntil(_ condition: () -> Bool) async { + for _ in 0 ..< 300 where !condition() { + try? await Task.sleep(for: .milliseconds(10)) + } + } + + @Test("A statement carries the token and the query timeout and returns its rows") + func roundTrip() async throws { + let client = try connectedClient() + client.setQueryTimeout(300) + + let result = try await client.execute(sql: "SELECT 1") + + #expect(result.cols.map(\.name) == ["n"]) + #expect(result.rows.first?.first?.stringValue == "1") + let request = try #require(HranaStubProtocol.recorded) + #expect(request.url?.path == "/v2/pipeline") + #expect(request.value(forHTTPHeaderField: "Authorization") == "Bearer token") + #expect(request.timeoutInterval == HttpQueryTimeout(serverTimeoutSeconds: 300).requestTimeoutInterval) + #expect(client.inFlightCount == 0) + } + + @Test("Cancelling everything stops every request in flight, not just the latest") + func cancelAllStopsEveryRequest() async throws { + let client = try connectedClient() + let query = Task { try await client.execute(sql: "SELECT hang") } + let sidebarRead = Task { try await client.execute(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 2 } + + client.cancelAll() + + await #expect(throws: CancellationError.self) { try await query.value } + await #expect(throws: CancellationError.self) { try await sidebarRead.value } + #expect(client.inFlightCount == 0) + } + + @Test("A request that finishes leaves the others cancellable") + func finishedRequestKeepsOthersTracked() async throws { + let client = try connectedClient() + let query = Task { try await client.execute(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 1 } + + _ = try await client.execute(sql: "SELECT 1") + client.cancelAll() + + await #expect(throws: CancellationError.self) { try await query.value } + } + + @Test("Cancelling the awaiting task cancels its request") + func taskCancellation() async throws { + let client = try connectedClient() + let query = Task { try await client.execute(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 1 } + + query.cancel() + + await #expect(throws: CancellationError.self) { try await query.value } + #expect(client.inFlightCount == 0) + } + + @Test("Disconnecting cancels every request in flight") + func invalidateSessionCancels() async throws { + let client = try connectedClient() + let query = Task { try await client.execute(sql: "SELECT hang") } + await waitUntil { client.inFlightCount == 1 } + + client.invalidateSession() + + await #expect(throws: CancellationError.self) { try await query.value } + } +} diff --git a/project.yml b/project.yml index 5f0887862e..37156df5e3 100644 --- a/project.yml +++ b/project.yml @@ -494,7 +494,9 @@ targets: - Plugins/SQLiteDriverPlugin/SQLiteCreateTableDDL.swift - Plugins/SQLiteDriverPlugin/SQLiteDefaultValue.swift - Plugins/LibSQLDriverPlugin/LibSQLDefaultValue.swift + - Plugins/LibSQLDriverPlugin/HranaHttpClient.swift - Plugins/CloudflareD1DriverPlugin/CloudflareD1DefaultValue.swift + - Plugins/CloudflareD1DriverPlugin/D1HttpClient.swift - Plugins/MSSQLDriverPlugin/MSSQLCheckConstraintDefinition.swift - Plugins/PostgreSQLDriverPlugin/ColumnQueryShape.swift - Plugins/PostgreSQLDriverPlugin/LibPQByteaDecoder.swift