Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

### Fixed

- Connection screen stuck on Connecting for good after a cancelled connect on iPhone and iPad.
- Edited connection host, port or credentials ignored until relaunch on iPhone and iPad.
- SSH tunnel handshake with no timeout on iPhone and iPad, against a server that accepts TCP and then stalls.
- Export and Transfer To preselecting a same-named table from another schema, or nothing at all.
- Delete queuing a table drop from the menu bar with none of the confirmation the sidebar asks for.
- Truncate Table offered from the menu bar for a view, which the server then refuses.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ public final class ConnectionManager: @unchecked Sendable {
private var sessions: [UUID: ConnectionSession] = [:]
private var teardowns: [UUID: Task<Void, Never>] = [:]
private var blockingTeardowns: Set<UUID> = []
private var attemptGenerations: [UUID: Int] = [:]

public init(
driverFactory: DriverFactory,
Expand All @@ -22,20 +23,25 @@ public final class ConnectionManager: @unchecked Sendable {
}

public func connect(_ connection: DatabaseConnection) async throws -> ConnectionSession {
await disconnect(connection.id)
let generation = beginAttempt(for: connection.id)
await awaitTeardown(of: connection.id)
guard isCurrentAttempt(generation, for: connection.id) else { throw CancellationError() }
let password = try secureStore.retrieve(forKey: Self.passwordKey(for: connection.id))

var effectiveHost = connection.host
var effectivePort = connection.port
var tunnelId: UUID?
if connection.sshEnabled, let ssh = connection.sshConfiguration {
guard let provider = sshProvider else {
throw ConnectionError.sshNotSupported
}
let tunnel = try await provider.createTunnel(
config: ssh,
connectionId: connection.id,
remoteHost: connection.host,
remotePort: connection.port
)
tunnelId = tunnel.id
effectiveHost = tunnel.localHost
effectivePort = tunnel.localPort
}
Expand All @@ -54,16 +60,45 @@ public final class ConnectionManager: @unchecked Sendable {
activeDatabase: connection.database,
status: .connected
)
storeSession(session, for: connection.id)
guard adoptSession(session, for: connection.id, generation: generation) else {
try? await driver.disconnect()
if let tunnelId, let provider = sshProvider {
try? await provider.closeTunnel(id: tunnelId)
}
throw CancellationError()
}
return session
} catch {
if connection.sshEnabled, let provider = sshProvider {
try? await provider.closeTunnel(for: connection.id)
if let tunnelId, let provider = sshProvider {
try? await provider.closeTunnel(id: tunnelId)
}
throw error
}
}

/// `Task.cancel()` cannot retire an attempt: the drivers block in C calls that never observe it.
public func invalidateAttempt(for connectionId: UUID) {
lock.lock()
defer { lock.unlock() }
attemptGenerations[connectionId, default: 0] += 1
}

private func beginAttempt(for connectionId: UUID) -> Int {
lock.lock()
defer { lock.unlock() }
let next = (attemptGenerations[connectionId] ?? 0) + 1
attemptGenerations[connectionId] = next
return next
}

private func adoptSession(_ session: ConnectionSession, for id: UUID, generation: Int) -> Bool {
lock.lock()
defer { lock.unlock() }
guard attemptGenerations[id] == generation else { return false }
sessions[id] = session
return true
}

public func storePassword(_ password: String, for connectionId: UUID) throws {
try secureStore.store(password, forKey: Self.passwordKey(for: connectionId))
}
Expand All @@ -77,11 +112,22 @@ public final class ConnectionManager: @unchecked Sendable {
}

public func disconnect(_ connectionId: UUID) async {
invalidateAttempt(for: connectionId)
await awaitTeardown(of: connectionId)
}

private func awaitTeardown(of connectionId: UUID) async {
guard let teardown = claimTeardown(for: connectionId) else { return }
await teardown.value
finishTeardown(teardown, for: connectionId)
}

private func isCurrentAttempt(_ generation: Int, for connectionId: UUID) -> Bool {
lock.lock()
defer { lock.unlock() }
return attemptGenerations[connectionId] == generation
}

public var hasSuspensionBlockingResources: Bool {
!suspensionBlockingIds().isEmpty
}
Expand Down Expand Up @@ -152,10 +198,4 @@ public final class ConnectionManager: @unchecked Sendable {
defer { lock.unlock() }
return sessions[connectionId]
}

private func storeSession(_ session: ConnectionSession, for id: UUID) {
lock.lock()
sessions[id] = session
lock.unlock()
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,18 +4,23 @@ import TableProModels
public protocol SSHProvider: Sendable {
func createTunnel(
config: SSHConfiguration,
connectionId: UUID,
remoteHost: String,
remotePort: Int
) async throws -> SSHTunnel

func closeTunnel(for connectionId: UUID) async throws

func closeTunnel(id: UUID) async throws
}

public struct SSHTunnel: Sendable {
public let id: UUID
public let localHost: String
public let localPort: Int

public init(localHost: String, localPort: Int) {
public init(id: UUID = UUID(), localHost: String, localPort: Int) {
self.id = id
self.localHost = localHost
self.localPort = localPort
}
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
import Foundation
@testable import TableProDatabase
@testable import TableProModels
import Testing

// MARK: - Mock Types

final class MockDatabaseDriver: DatabaseDriver, @unchecked Sendable {
var isConnected = false
var shouldFailConnect = false
var holdsSuspensionBlockingResource = false
var onConnect: (@Sendable () -> Void)?
var beforeConnect: (@Sendable () async -> Void)?
var beforeDisconnect: (@Sendable () async -> Void)?
private(set) var disconnectCount = 0

func connect() async throws {
await beforeConnect?()
if shouldFailConnect { throw NSError(domain: "test", code: 1) }
onConnect?()
isConnected = true
}

func disconnect() async throws {
await beforeDisconnect?()
isConnected = false
disconnectCount += 1
}
func ping() async throws -> Bool { isConnected }

func execute(query: String) async throws -> QueryResult {
QueryResult(columns: [], rows: [], rowsAffected: 0, executionTime: 0, isTruncated: false, statusMessage: nil)
}

func cancelCurrentQuery() async throws {}
func fetchTables(schema: String?) async throws -> [TableInfo] { [] }
func fetchColumns(table: String, schema: String?) async throws -> [ColumnInfo] { [] }
func fetchIndexes(table: String, schema: String?) async throws -> [IndexInfo] { [] }
func fetchForeignKeys(table: String, schema: String?) async throws -> [ForeignKeyInfo] { [] }
func fetchDatabases() async throws -> [String] { [] }
func switchDatabase(to name: String) async throws {}
var supportsSchemas: Bool { false }
func switchSchema(to name: String) async throws {}
func fetchSchemas() async throws -> [String] { [] }
var currentSchema: String? { nil }
var supportsTransactions: Bool { false }
func beginTransaction() async throws {}
func commitTransaction() async throws {}
func rollbackTransaction() async throws {}
var serverVersion: String? { nil }
}

final class MockDriverFactory: DriverFactory, @unchecked Sendable {
var drivers: [String: any DatabaseDriver] = [:]

func createDriver(for connection: DatabaseConnection, password: String?) throws -> any DatabaseDriver {
guard let driver = drivers[connection.type.rawValue] else {
throw ConnectionError.driverNotFound(connection.type.rawValue)
}
return driver
}

func supportedTypes() -> [DatabaseType] { [] }
}

final class MockSecureStore: SecureStore, Sendable {
private let passwords: [String: String]

init(passwords: [String: String] = [:]) {
self.passwords = passwords
}

func store(_ value: String, forKey key: String) throws {}

func retrieve(forKey key: String) throws -> String? {
passwords[key]
}

func delete(forKey key: String) throws {}
}

// MARK: - Synchronisation Helpers

actor Gate {
private var isOpen = false
private var hasEntered = false
private var blocked: [CheckedContinuation<Void, Never>] = []
private var observers: [CheckedContinuation<Void, Never>] = []

func enter() async {
hasEntered = true
for observer in observers { observer.resume() }
observers.removeAll()
guard !isOpen else { return }
await withCheckedContinuation { blocked.append($0) }
}

func open() {
isOpen = true
for waiter in blocked { waiter.resume() }
blocked.removeAll()
}

func waitUntilEntered() async {
guard !hasEntered else { return }
await withCheckedContinuation { observers.append($0) }
}
}

actor Barrier {
private static let pollInterval: UInt64 = 20_000_000
private static let maxPolls = 250

private let expected: Int
private var arrived = 0
private(set) var overlapped = 0

init(expected: Int) {
self.expected = expected
}

func arriveAndWait() async {
arrived += 1
var polls = 0
while arrived < expected, polls < Self.maxPolls {
try? await Task.sleep(nanoseconds: Self.pollInterval)
polls += 1
}
guard arrived >= expected else { return }
overlapped += 1
}
}

// MARK: - Mock SSH Provider

final class MockSSHProvider: SSHProvider, @unchecked Sendable {
var closedTunnels: Set<UUID> = []
var closedTunnelIds: Set<UUID> = []
var openedTunnelIds: [UUID] = []
var tunnelledConnectionIds: [UUID] = []

func createTunnel(
config: SSHConfiguration,
connectionId: UUID,
remoteHost: String,
remotePort: Int
) async throws -> SSHTunnel {
tunnelledConnectionIds.append(connectionId)
let id = UUID()
openedTunnelIds.append(id)
return SSHTunnel(id: id, localHost: "127.0.0.1", localPort: 33_306)
}

func closeTunnel(for connectionId: UUID) async throws {
closedTunnels.insert(connectionId)
}

func closeTunnel(id: UUID) async throws {
closedTunnelIds.insert(id)
}
}
Loading
Loading