Skip to content

Commit 7021deb

Browse files
committed
Add custom SQL function support
Register custom scalar functions with a SQLite store and invoke them from predicates and sort descriptors, plus a raw-SQL escape hatch for app-managed virtual tables (e.g. R*Tree). - SQLiteDatabase/SQLiteViewContext.register(function:) bridges a DatabaseFunction to Connection.createFunction, converting between AttributeValue and SQLite Binding at the boundary. - SQLiteDatabase.execute(_:_:) runs raw SQL so apps can create and maintain their own virtual tables; the library never creates or syncs them itself. - Predicate.swift compiles .function expressions to SQL function calls in WHERE clauses; Database.swift emits them in ORDER BY.
1 parent 46ebd93 commit 7021deb

5 files changed

Lines changed: 313 additions & 6 deletions

File tree

‎Sources/CoreModelSQLite/AttributeValue.swift‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,26 @@ internal extension AttributeValue {
4343
}
4444
}
4545

46+
/// Decode from a SQLite binding value with no declared attribute type (e.g. a
47+
/// raw argument passed into a custom SQL function), inferring the value's shape
48+
/// from the binding's runtime type.
49+
init(binding: Binding?) {
50+
switch binding {
51+
case .none:
52+
self = .null
53+
case let value as Int64:
54+
self = .int64(value)
55+
case let value as Double:
56+
self = .double(value)
57+
case let value as String:
58+
self = .string(value)
59+
case let value as Blob:
60+
self = .data(Data(value.bytes))
61+
default:
62+
self = .null
63+
}
64+
}
65+
4666
/// Decode from a SQLite binding value, interpreting it according to the declared attribute type.
4767
init(binding: Binding?, type: AttributeType) throws {
4868
guard let binding else {

‎Sources/CoreModelSQLite/Database.swift‎

Lines changed: 46 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -111,6 +111,37 @@ extension SQLiteDatabase: ModelStorage {
111111
try connection.delete(entity, for: ids, model: model)
112112
invalidateCache(for: [entity])
113113
}
114+
115+
public func register(function: DatabaseFunction) async throws {
116+
connection.register(function: function)
117+
}
118+
}
119+
120+
public extension SQLiteDatabase {
121+
122+
/// Execute raw SQL against the underlying connection — e.g. to create and
123+
/// maintain an R*Tree or other virtual table. CoreModelSQLite does not create,
124+
/// sync, or otherwise know about any virtual table itself; that is entirely the
125+
/// caller's responsibility.
126+
func execute(_ sql: String, _ bindings: [Binding?] = []) throws {
127+
try connection.run(sql, bindings)
128+
}
129+
}
130+
131+
internal extension SQLite.Connection {
132+
133+
/// Registers a `DatabaseFunction` with this connection, bridging SQLite's untyped
134+
/// `Binding` values to/from `AttributeValue` at the boundary.
135+
func register(function: DatabaseFunction) {
136+
createFunction(
137+
function.name,
138+
argumentCount: function.argumentCount.map { UInt($0) },
139+
deterministic: function.deterministic
140+
) { arguments in
141+
let values: [AttributeValue?] = arguments.map { AttributeValue(binding: $0) }
142+
return function.evaluate(values)?.binding
143+
}
144+
}
114145
}
115146

116147
internal extension SQLiteDatabase {
@@ -352,13 +383,23 @@ internal extension FetchRequest {
352383
bindings += fragment.bindings
353384
}
354385
if sortDescriptors.isEmpty == false {
355-
let terms = try sortDescriptors.map { sort -> String in
356-
guard entity.hasColumn(for: sort.property) else {
357-
throw SQLiteDatabaseError.invalidProperty(sort.property, entity.id)
386+
let placeholderPredicate = FetchRequest.Predicate.value(true)
387+
let fragments = try sortDescriptors.map { sort -> SQLFragment in
388+
switch sort.term {
389+
case let .property(property):
390+
guard entity.hasColumn(for: property) else {
391+
throw SQLiteDatabaseError.invalidProperty(property, entity.id)
392+
}
393+
let sql = property.rawValue.quotedIdentifier + (sort.ascending ? " ASC" : " DESC")
394+
return SQLFragment(sql: sql, bindings: [])
395+
case let .function(function):
396+
let functionFragment = try function.sqlFragment(for: entity, predicate: placeholderPredicate)
397+
let sql = functionFragment.sql + (sort.ascending ? " ASC" : " DESC")
398+
return SQLFragment(sql: sql, bindings: functionFragment.bindings)
358399
}
359-
return sort.property.rawValue.quotedIdentifier + (sort.ascending ? " ASC" : " DESC")
360400
}
361-
sql += " ORDER BY " + terms.joined(separator: ", ")
401+
sql += " ORDER BY " + fragments.map(\.sql).joined(separator: ", ")
402+
bindings += fragments.flatMap(\.bindings)
362403
} else {
363404
// match CoreData's default behavior of sorting by object ID when no sort descriptors are provided
364405
sql += " ORDER BY \(SQLiteDatabase.primaryKeyColumn.quotedIdentifier) ASC"

‎Sources/CoreModelSQLite/Predicate.swift‎

Lines changed: 65 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -74,6 +74,35 @@ internal extension FetchRequest.Predicate.Comparison {
7474
predicate: FetchRequest.Predicate
7575
) throws -> SQLFragment {
7676

77+
// `function(...) <operator> constant` comparisons compile to a SQL function call.
78+
if case let .function(function) = left {
79+
guard modifier == nil else {
80+
throw SQLiteDatabaseError.invalidPredicate(predicate)
81+
}
82+
let functionFragment = try function.sqlFragment(for: entity, predicate: predicate)
83+
switch type {
84+
case .lessThan, .lessThanOrEqualTo, .greaterThan, .greaterThanOrEqualTo:
85+
let value = try right.constantBinding(predicate: predicate)
86+
return SQLFragment(
87+
sql: "\(functionFragment.sql) \(type.rawValue) ?",
88+
bindings: functionFragment.bindings + [value]
89+
)
90+
case .equalTo, .notEqualTo:
91+
let value = try right.constantBinding(predicate: predicate)
92+
let sqlOperator = (type == .equalTo) ? "=" : "<>"
93+
guard let value else {
94+
let nullOperator = (type == .equalTo) ? "IS NULL" : "IS NOT NULL"
95+
return SQLFragment(sql: "\(functionFragment.sql) \(nullOperator)", bindings: functionFragment.bindings)
96+
}
97+
return SQLFragment(
98+
sql: "\(functionFragment.sql) \(sqlOperator) ?",
99+
bindings: functionFragment.bindings + [value]
100+
)
101+
default:
102+
throw SQLiteDatabaseError.invalidPredicate(predicate)
103+
}
104+
}
105+
77106
// Only `keyPath <operator> constant` comparisons map directly to columns.
78107
guard case let .keyPath(keyPath) = left else {
79108
throw SQLiteDatabaseError.invalidPredicate(predicate)
@@ -207,8 +236,43 @@ private extension SQLFragment {
207236
}
208237
}
209238

239+
internal extension FetchRequest.Predicate.FunctionExpression {
240+
241+
/// Translate a function call expression into a SQL function-call fragment,
242+
/// e.g. `myFunction("lat", "lon", ?, ?)`.
243+
func sqlFragment(
244+
for entity: EntityDescription,
245+
predicate: FetchRequest.Predicate
246+
) throws -> SQLFragment {
247+
let argumentFragments = try arguments.map {
248+
try $0.argumentSQLFragment(for: entity, predicate: predicate)
249+
}
250+
return SQLFragment(
251+
sql: "\(name)(" + argumentFragments.map(\.sql).joined(separator: ", ") + ")",
252+
bindings: argumentFragments.flatMap(\.bindings)
253+
)
254+
}
255+
}
256+
210257
private extension FetchRequest.Predicate.Expression {
211258

259+
/// The expression as a SQL fragment suitable for use as a function argument
260+
/// (a column reference, a constant placeholder, or a nested function call).
261+
func argumentSQLFragment(
262+
for entity: EntityDescription,
263+
predicate: FetchRequest.Predicate
264+
) throws -> SQLFragment {
265+
switch self {
266+
case let .keyPath(keyPath):
267+
let column = try entity.validateColumn(PropertyKey(rawValue: keyPath.rawValue), predicate: predicate)
268+
return SQLFragment(sql: column, bindings: [])
269+
case let .function(function):
270+
return try function.sqlFragment(for: entity, predicate: predicate)
271+
case .attribute, .relationship:
272+
return SQLFragment(sql: "?", bindings: [try constantBinding(predicate: predicate)])
273+
}
274+
}
275+
212276
/// The expression as a single constant binding.
213277
func constantBinding(predicate: FetchRequest.Predicate) throws -> Binding? {
214278
switch self {
@@ -223,7 +287,7 @@ private extension FetchRequest.Predicate.Expression {
223287
case .toMany:
224288
throw SQLiteDatabaseError.invalidPredicate(predicate)
225289
}
226-
case .keyPath:
290+
case .keyPath, .function:
227291
throw SQLiteDatabaseError.invalidPredicate(predicate)
228292
}
229293
}

‎Sources/CoreModelSQLite/ViewContext.swift‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,4 +54,13 @@ public final class SQLiteViewContext: ViewContext {
5454
public func count(_ fetchRequest: FetchRequest) throws -> UInt {
5555
try connection.count(fetchRequest, model: model)
5656
}
57+
58+
/// Registers a custom function on this context's read-only connection.
59+
///
60+
/// Function registration is per-connection: a function registered with a
61+
/// paired ``SQLiteDatabase`` must also be registered here to be usable from
62+
/// queries run through this view context.
63+
public func register(function: DatabaseFunction) throws {
64+
connection.register(function: function)
65+
}
5766
}
Lines changed: 173 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,173 @@
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

Comments
 (0)