Skip to content

Commit 85cf682

Browse files
authored
Merge pull request #7 from PureSwift/feature/function
Add custom SQL function support
2 parents 46ebd93 + 403a05f commit 85cf682

6 files changed

Lines changed: 474 additions & 7 deletions

File tree

‎Package.swift‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,7 @@ let package = Package(
3434
dependencies: [
3535
.package(
3636
url: "https://github.com/PureSwift/CoreModel",
37-
from: "2.7.2"
37+
from: "2.8.0"
3838
),
3939
sqliteDependency
4040
],

‎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: 53 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -111,6 +111,44 @@ extension SQLiteDatabase: ModelStorage {
111111
try connection.delete(entity, for: ids, model: model)
112112
invalidateCache(for: [entity])
113113
}
114+
115+
/// Registers a custom scalar function so it can be invoked from a predicate or sort
116+
/// descriptor via ``FetchRequest/Predicate/Expression/function(_:)``.
117+
///
118+
/// - Important: Only supported on Apple platforms. The underlying SQLite.swift
119+
/// `createFunction` registers the callback through `@convention(block)` +
120+
/// `unsafeBitCast`, which is unreliable on non-Apple platforms and can corrupt
121+
/// SQLite's heap (upstream https://github.com/stephencelis/SQLite.swift/issues/1071).
122+
public func register(function: DatabaseFunction) async throws {
123+
connection.register(function: function)
124+
}
125+
}
126+
127+
public extension SQLiteDatabase {
128+
129+
/// Execute raw SQL against the underlying connection — e.g. to create and
130+
/// maintain an R*Tree or other virtual table. CoreModelSQLite does not create,
131+
/// sync, or otherwise know about any virtual table itself; that is entirely the
132+
/// caller's responsibility.
133+
func execute(_ sql: String, _ bindings: [Binding?] = []) throws {
134+
try connection.run(sql, bindings)
135+
}
136+
}
137+
138+
internal extension SQLite.Connection {
139+
140+
/// Registers a `DatabaseFunction` with this connection, bridging SQLite's untyped
141+
/// `Binding` values to/from `AttributeValue` at the boundary.
142+
func register(function: DatabaseFunction) {
143+
createFunction(
144+
function.name,
145+
argumentCount: function.argumentCount.map { UInt($0) },
146+
deterministic: function.deterministic
147+
) { arguments in
148+
let values: [AttributeValue?] = arguments.map { AttributeValue(binding: $0) }
149+
return function.evaluate(values)?.binding
150+
}
151+
}
114152
}
115153

116154
internal extension SQLiteDatabase {
@@ -352,13 +390,23 @@ internal extension FetchRequest {
352390
bindings += fragment.bindings
353391
}
354392
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)
393+
let placeholderPredicate = FetchRequest.Predicate.value(true)
394+
let fragments = try sortDescriptors.map { sort -> SQLFragment in
395+
switch sort.term {
396+
case let .property(property):
397+
guard entity.hasColumn(for: property) else {
398+
throw SQLiteDatabaseError.invalidProperty(property, entity.id)
399+
}
400+
let sql = property.rawValue.quotedIdentifier + (sort.ascending ? " ASC" : " DESC")
401+
return SQLFragment(sql: sql, bindings: [])
402+
case let .function(function):
403+
let functionFragment = try function.sqlFragment(for: entity, predicate: placeholderPredicate)
404+
let sql = functionFragment.sql + (sort.ascending ? " ASC" : " DESC")
405+
return SQLFragment(sql: sql, bindings: functionFragment.bindings)
358406
}
359-
return sort.property.rawValue.quotedIdentifier + (sort.ascending ? " ASC" : " DESC")
360407
}
361-
sql += " ORDER BY " + terms.joined(separator: ", ")
408+
sql += " ORDER BY " + fragments.map(\.sql).joined(separator: ", ")
409+
bindings += fragments.flatMap(\.bindings)
362410
} else {
363411
// match CoreData's default behavior of sorting by object ID when no sort descriptors are provided
364412
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
}

0 commit comments

Comments
 (0)