|
| 1 | +import Foundation |
| 2 | +import Testing |
| 3 | +import CoreModel |
| 4 | +import SQLite |
| 5 | +@testable import CoreModelSQLite |
| 6 | + |
| 7 | +/// A Haversine distance function, in meters, written directly in the test — CoreModelSQLite |
| 8 | +/// itself has no notion of "distance" or geo data; this exercises the generic |
| 9 | +/// `DatabaseFunction`/`.function` expression mechanism using a realistic example. |
| 10 | +private func haversineDistance(_ lat1: Double, _ lon1: Double, _ lat2: Double, _ lon2: Double) -> Double { |
| 11 | + let earthRadius = 6_371_000.0 |
| 12 | + let dLat = (lat2 - lat1) * .pi / 180 |
| 13 | + let dLon = (lon2 - lon1) * .pi / 180 |
| 14 | + let a = sin(dLat / 2) * sin(dLat / 2) |
| 15 | + + cos(lat1 * .pi / 180) * cos(lat2 * .pi / 180) * sin(dLon / 2) * sin(dLon / 2) |
| 16 | + let c = 2 * atan2(a.squareRoot(), (1 - a).squareRoot()) |
| 17 | + return earthRadius * c |
| 18 | +} |
| 19 | + |
| 20 | +private let distanceFunction = DatabaseFunction(name: "distance", argumentCount: 4) { arguments in |
| 21 | + guard case let .double(lat1) = arguments[0], |
| 22 | + case let .double(lon1) = arguments[1], |
| 23 | + case let .double(lat2) = arguments[2], |
| 24 | + case let .double(lon2) = arguments[3] |
| 25 | + else { return nil } |
| 26 | + return .double(haversineDistance(lat1, lon1, lat2, lon2)) |
| 27 | +} |
| 28 | + |
| 29 | +private func makeGeoDatabase() async throws -> SQLiteDatabase { |
| 30 | + let model = Model(entities: [ |
| 31 | + EntityDescription( |
| 32 | + id: "Site", |
| 33 | + attributes: [ |
| 34 | + .init(id: "name", type: .string), |
| 35 | + .init(id: "latitude", type: .double), |
| 36 | + .init(id: "longitude", type: .double) |
| 37 | + ], |
| 38 | + relationships: [] |
| 39 | + ) |
| 40 | + ]) |
| 41 | + let database = try SQLiteDatabase(path: temporaryDatabasePath(named: "GeoTests"), model: model) |
| 42 | + try await database.register(function: distanceFunction) |
| 43 | + return database |
| 44 | +} |
| 45 | + |
| 46 | +/// Known coordinates, distances computed against Raleigh, NC (35.7796, -78.6382), the |
| 47 | +/// reference point every test below filters/sorts by. |
| 48 | +private let referenceLatitude = 35.7796 |
| 49 | +private let referenceLongitude = -78.6382 |
| 50 | + |
| 51 | +private let sites: [(id: ObjectID, name: String, latitude: Double, longitude: Double)] = [ |
| 52 | + ("raleigh", "Raleigh, NC", 35.7796, -78.6382), // ~0 m |
| 53 | + ("durham", "Durham, NC", 35.9940, -78.8986), // ~30 km |
| 54 | + ("charlotte", "Charlotte, NC", 35.2271, -80.8431), // ~210 km |
| 55 | + ("atlanta", "Atlanta, GA", 33.7490, -84.3880), // ~530 km |
| 56 | + ("nyc", "New York, NY", 40.7128, -74.0060) // ~660 km |
| 57 | +] |
| 58 | + |
| 59 | +private func insertSites(_ database: SQLiteDatabase) async throws { |
| 60 | + for site in sites { |
| 61 | + try await database.insert(ModelData( |
| 62 | + entity: "Site", |
| 63 | + id: site.id, |
| 64 | + attributes: [ |
| 65 | + "name": .string(site.name), |
| 66 | + "latitude": .double(site.latitude), |
| 67 | + "longitude": .double(site.longitude) |
| 68 | + ] |
| 69 | + )) |
| 70 | + } |
| 71 | +} |
| 72 | + |
| 73 | +private func distanceExpression() -> FetchRequest.Predicate.Expression { |
| 74 | + .function(.init(name: "distance", arguments: [ |
| 75 | + .keyPath("latitude"), |
| 76 | + .keyPath("longitude"), |
| 77 | + .attribute(.double(referenceLatitude)), |
| 78 | + .attribute(.double(referenceLongitude)) |
| 79 | + ])) |
| 80 | +} |
| 81 | + |
| 82 | +/// Independently computed oracle, not exercising any library code, to check results against. |
| 83 | +private func expectedIDs(within radiusMeters: Double) -> Set<ObjectID> { |
| 84 | + Set(sites.filter { haversineDistance($0.latitude, $0.longitude, referenceLatitude, referenceLongitude) <= radiusMeters }.map(\.id)) |
| 85 | +} |
| 86 | + |
| 87 | +@Test func functionPredicateFiltersByRadius() async throws { |
| 88 | + let database = try await makeGeoDatabase() |
| 89 | + try await insertSites(database) |
| 90 | + |
| 91 | + let radius = 250_000.0 |
| 92 | + let request = FetchRequest( |
| 93 | + entity: "Site", |
| 94 | + predicate: .comparison(.init(left: distanceExpression(), right: .attribute(.double(radius)), type: .lessThanOrEqualTo)) |
| 95 | + ) |
| 96 | + let ids = Set(try await database.fetchID(request)) |
| 97 | + #expect(ids == expectedIDs(within: radius)) |
| 98 | + #expect(ids == ["raleigh", "durham", "charlotte"]) |
| 99 | +} |
| 100 | + |
| 101 | +@Test func functionSortOrdersByDistance() async throws { |
| 102 | + let database = try await makeGeoDatabase() |
| 103 | + try await insertSites(database) |
| 104 | + |
| 105 | + let request = FetchRequest( |
| 106 | + entity: "Site", |
| 107 | + sortDescriptors: [.init(term: .function(.init(name: "distance", arguments: [ |
| 108 | + .keyPath("latitude"), .keyPath("longitude"), |
| 109 | + .attribute(.double(referenceLatitude)), .attribute(.double(referenceLongitude)) |
| 110 | + ])), ascending: true)] |
| 111 | + ) |
| 112 | + let ids = try await database.fetchID(request) |
| 113 | + #expect(ids == ["raleigh", "durham", "charlotte", "atlanta", "nyc"]) |
| 114 | +} |
| 115 | + |
| 116 | +@Test func functionFilterAndSortWithLimit() async throws { |
| 117 | + let database = try await makeGeoDatabase() |
| 118 | + try await insertSites(database) |
| 119 | + |
| 120 | + let radius = 700_000.0 |
| 121 | + let request = FetchRequest( |
| 122 | + entity: "Site", |
| 123 | + sortDescriptors: [.init(term: .function(.init(name: "distance", arguments: [ |
| 124 | + .keyPath("latitude"), .keyPath("longitude"), |
| 125 | + .attribute(.double(referenceLatitude)), .attribute(.double(referenceLongitude)) |
| 126 | + ])), ascending: true)], |
| 127 | + predicate: .comparison(.init(left: distanceExpression(), right: .attribute(.double(radius)), type: .lessThanOrEqualTo)), |
| 128 | + fetchLimit: 2 |
| 129 | + ) |
| 130 | + let ids = try await database.fetchID(request) |
| 131 | + #expect(ids == ["raleigh", "durham"]) |
| 132 | +} |
| 133 | + |
| 134 | +@MainActor |
| 135 | +@Test func functionRegisteredOnViewContext() async throws { |
| 136 | + let path = temporaryDatabasePath(named: "GeoViewContextTests") |
| 137 | + let model = Model(entities: [ |
| 138 | + EntityDescription( |
| 139 | + id: "Site", |
| 140 | + attributes: [ |
| 141 | + .init(id: "name", type: .string), |
| 142 | + .init(id: "latitude", type: .double), |
| 143 | + .init(id: "longitude", type: .double) |
| 144 | + ], |
| 145 | + relationships: [] |
| 146 | + ) |
| 147 | + ]) |
| 148 | + let database = try SQLiteDatabase(path: path, model: model) |
| 149 | + try await database.register(function: distanceFunction) |
| 150 | + try await insertSites(database) |
| 151 | + |
| 152 | + let viewContext = try SQLiteViewContext(.uri(path), model: model) |
| 153 | + try viewContext.register(function: distanceFunction) |
| 154 | + |
| 155 | + let radius = 250_000.0 |
| 156 | + let request = FetchRequest( |
| 157 | + entity: "Site", |
| 158 | + predicate: .comparison(.init(left: distanceExpression(), right: .attribute(.double(radius)), type: .lessThanOrEqualTo)) |
| 159 | + ) |
| 160 | + let ids = Set(try viewContext.fetchID(request)) |
| 161 | + #expect(ids == expectedIDs(within: radius)) |
| 162 | +} |
| 163 | + |
| 164 | +@Test func unrelatedEntitiesUnaffected() async throws { |
| 165 | + // regression: entities/queries with no `.function` usage behave exactly as before |
| 166 | + let database = try makeDatabase() |
| 167 | + let people = (0..<3).map { index in |
| 168 | + ModelData(entity: "Person", id: ObjectID(rawValue: "person\(index)"), attributes: ["name": .string("Person \(index)")]) |
| 169 | + } |
| 170 | + try await database.insert(people) |
| 171 | + let count = try await database.count(FetchRequest(entity: "Person")) |
| 172 | + #expect(count == 3) |
| 173 | +} |
0 commit comments