From 98c5c489f37f0e54dfc86303f3feead44b5f67cc Mon Sep 17 00:00:00 2001 From: Stephen Celis Date: Tue, 18 Aug 2026 14:04:04 -0700 Subject: [PATCH] Use cached statements GRDB has a `cachedStatement` API that we can leverage to improve performance when the underlying SQL matches. --- Sources/SQLiteData/FetchAll+Sections.swift | 12 +-- Sources/SQLiteData/FetchAll.swift | 15 ++-- Sources/SQLiteData/FetchOne.swift | 39 +++++----- .../SQLiteData/Internal/StatementKey.swift | 7 +- .../StructuredQueries+GRDB/QueryCursor.swift | 59 ++++++++++---- .../Statement+GRDB.swift | 29 +++---- .../Statement+Sections.swift | 4 +- .../SQLiteDataTests/StatementCacheTests.swift | 76 +++++++++++++++++++ 8 files changed, 178 insertions(+), 63 deletions(-) create mode 100644 Tests/SQLiteDataTests/StatementCacheTests.swift diff --git a/Sources/SQLiteData/FetchAll+Sections.swift b/Sources/SQLiteData/FetchAll+Sections.swift index f7287284..e3e7adc7 100644 --- a/Sources/SQLiteData/FetchAll+Sections.swift +++ b/Sources/SQLiteData/FetchAll+Sections.swift @@ -1556,7 +1556,7 @@ struct FetchAllSectionedStatementValueRequest< Key: QueryRepresentable >: FetchKeyRequest where Value.QueryOutput: Sendable, Key.QueryOutput: Hashable & Sendable { - let statement: SQLQueryExpression<(Value, Key)> + let prepared: PreparedQuery init( statement: Select<(), Value, ()>, @@ -1564,7 +1564,7 @@ where Value.QueryOutput: Sendable, Key.QueryOutput: Hashable & Sendable { ) where Value: StructuredQueriesCore.Table { let prefix: Select<(Value, Key), Value, ()> = sectionedColumns(of: Value.self, sectionBy) let sectioned: Select<(Value, Key), Value, ()> = prefix + statement - self.statement = SQLQueryExpression(sectioned) + self.prepared = PreparedQuery(sectioned.query) } init( @@ -1575,21 +1575,21 @@ where Value.QueryOutput: Sendable, Key.QueryOutput: Hashable & Sendable { sectionedOrder(of: From.self, sectionBy) + statement let sectioned: Select<(Value, Key), From, (repeat each J)> = ordered + sectionedColumn(of: From.self, sectionBy) - self.statement = SQLQueryExpression(sectioned) + self.prepared = PreparedQuery(sectioned.query) } func fetch(_ db: Database) throws -> ResultsSectionCollection { try ResultsSectionCollection( - cursor: QuerySectionedCursor(db: db, query: statement.queryFragment) + cursor: QuerySectionedCursor(db: db, prepared: prepared, cached: true) ) } static func == (lhs: Self, rhs: Self) -> Bool { - lhs.statement.query == rhs.statement.query + lhs.prepared == rhs.prepared } func hash(into hasher: inout Hasher) { - hasher.combine(statement.query) + hasher.combine(prepared) } } diff --git a/Sources/SQLiteData/FetchAll.swift b/Sources/SQLiteData/FetchAll.swift index 99c42b05..839e8063 100644 --- a/Sources/SQLiteData/FetchAll.swift +++ b/Sources/SQLiteData/FetchAll.swift @@ -597,12 +597,15 @@ extension FetchAll: Equatable where Element: Equatable { } #endif -struct FetchAllStatementValueRequest: StatementKeyRequest { - let statement: SQLQueryExpression - init(statement: some StructuredQueriesCore.Statement) { - self.statement = SQLQueryExpression(statement) +struct FetchAllStatementValueRequest: StatementKeyRequest { + let prepared: PreparedQuery + init(statement: some StructuredQueriesCore.Statement) { + self.prepared = PreparedQuery(statement.query) } - func fetch(_ db: Database) throws -> [Value.QueryOutput] { - try statement.fetchAll(db) + func fetch(_ db: Database) throws -> [QueryValue.QueryOutput] { + let cursor = try QueryValueCursor(db: db, prepared: prepared, cached: true) + var output: [QueryValue.QueryOutput] = [] + try cursor.forEach { output.append($0) } + return output } } diff --git a/Sources/SQLiteData/FetchOne.swift b/Sources/SQLiteData/FetchOne.swift index 2944682e..3fa698c6 100644 --- a/Sources/SQLiteData/FetchOne.swift +++ b/Sources/SQLiteData/FetchOne.swift @@ -1495,38 +1495,39 @@ extension FetchOne: Equatable where Value: Equatable { } #endif -private struct FetchOneStatementValueRequest: StatementKeyRequest { - let statement: SQLQueryExpression - init(statement: some StructuredQueriesCore.Statement) { - self.statement = SQLQueryExpression(statement) +private struct FetchOneStatementValueRequest: StatementKeyRequest { + let prepared: PreparedQuery + init(statement: some StructuredQueriesCore.Statement) { + self.prepared = PreparedQuery(statement.query) } - func fetch(_ db: Database) throws -> Value.QueryOutput { - guard let result = try statement.fetchOne(db) + func fetch(_ db: Database) throws -> QueryValue.QueryOutput { + guard + let result = try QueryValueCursor(db: db, prepared: prepared, cached: true).next() else { throw NotFound() } return result } } -private struct FetchOneStatementOptionalValueRequest: +private struct FetchOneStatementOptionalValueRequest: StatementKeyRequest { - let statement: SQLQueryExpression - init(statement: some StructuredQueriesCore.Statement) { - self.statement = SQLQueryExpression(statement) + let prepared: PreparedQuery + init(statement: some StructuredQueriesCore.Statement) { + self.prepared = PreparedQuery(statement.query) } - func fetch(_ db: Database) throws -> Value.QueryOutput? { - try statement.fetchOne(db) + func fetch(_ db: Database) throws -> QueryValue.QueryOutput? { + try QueryValueCursor(db: db, prepared: prepared, cached: true).next() } } private struct FetchOneStatementOptionalProtocolRequest< - Value: QueryRepresentable & StructuredQueriesCore._OptionalProtocol ->: StatementKeyRequest where Value.QueryOutput: StructuredQueriesCore._OptionalProtocol { - let statement: SQLQueryExpression - init(statement: some StructuredQueriesCore.Statement) { - self.statement = SQLQueryExpression(statement) + QueryValue: QueryRepresentable & StructuredQueriesCore._OptionalProtocol +>: StatementKeyRequest where QueryValue.QueryOutput: StructuredQueriesCore._OptionalProtocol { + let prepared: PreparedQuery + init(statement: some StructuredQueriesCore.Statement) { + self.prepared = PreparedQuery(statement.query) } - func fetch(_ db: Database) throws -> Value.QueryOutput { - try statement.fetchOne(db) ?? ._none + func fetch(_ db: Database) throws -> QueryValue.QueryOutput { + try QueryValueCursor(db: db, prepared: prepared, cached: true).next() ?? ._none } } diff --git a/Sources/SQLiteData/Internal/StatementKey.swift b/Sources/SQLiteData/Internal/StatementKey.swift index b9172ac3..f2de074f 100644 --- a/Sources/SQLiteData/Internal/StatementKey.swift +++ b/Sources/SQLiteData/Internal/StatementKey.swift @@ -2,15 +2,16 @@ import StructuredQueriesCore protocol StatementKeyRequest: FetchKeyRequest { associatedtype QueryValue - var statement: SQLQueryExpression { get } + var prepared: PreparedQuery { get } } extension StatementKeyRequest { static func == (lhs: Self, rhs: Self) -> Bool { - lhs.statement.query == rhs.statement.query + lhs.prepared.sql == rhs.prepared.sql && lhs.prepared.bindings == rhs.prepared.bindings } func hash(into hasher: inout Hasher) { - hasher.combine(statement.query) + hasher.combine(prepared.sql) + hasher.combine(prepared.bindings) } } diff --git a/Sources/SQLiteData/StructuredQueries+GRDB/QueryCursor.swift b/Sources/SQLiteData/StructuredQueries+GRDB/QueryCursor.swift index c6ef826d..12096d63 100644 --- a/Sources/SQLiteData/StructuredQueries+GRDB/QueryCursor.swift +++ b/Sources/SQLiteData/StructuredQueries+GRDB/QueryCursor.swift @@ -14,12 +14,18 @@ public class QueryCursor: DatabaseCursor { var decoder: SQLiteQueryDecoder @usableFromInline - init(db: Database, query: QueryFragment) throws { - (_statement, decoder) = try db.prepare(query: query) + init(db: Database, prepared: PreparedQuery, cached: Bool) throws { + (_statement, decoder) = try db.prepare(prepared, cached: cached) + } + + @usableFromInline + convenience init(db: Database, query: QueryFragment, cached: Bool) throws { + try self.init(db: db, prepared: PreparedQuery(query), cached: cached) } deinit { sqlite3_reset(_statement.sqliteStatement) + sqlite3_clear_bindings(_statement.sqliteStatement) } public func _element(sqliteStatement _: SQLiteStatement) throws -> Element { @@ -84,8 +90,8 @@ final class QueryValueCursor: QueryCursor { // NB: Required to workaround a "Legacy previews execution" bug // https://github.com/pointfreeco/sqlite-data/pull/60 @usableFromInline - override init(db: Database, query: QueryFragment) throws { - try super.init(db: db, query: query) + override init(db: Database, prepared: PreparedQuery, cached: Bool) throws { + try super.init(db: db, prepared: prepared, cached: cached) } @inlinable @@ -175,15 +181,38 @@ final class QueryVoidCursor: QueryCursor { } } -extension Database { - @inlinable - func prepare(query: QueryFragment) throws -> (GRDB.Statement, SQLiteQueryDecoder) { +@usableFromInline +struct PreparedQuery: Hashable, Sendable { + @usableFromInline + let sql: String + + @usableFromInline + let bindings: [QueryBinding] + + @usableFromInline + init(_ query: QueryFragment) { var (sql, bindings) = query.prepare { _ in "?" } if sql.isEmpty { sql = "SELECT 1 WHERE 0 -- Empty query generated by StructuredQueries" } - let statement = try makeStatement(sql: sql) - for (index, binding) in zip(Int32(1)..., bindings) { + self.sql = sql + self.bindings = bindings + } +} + +extension Database { + @usableFromInline + func prepare( + _ prepared: PreparedQuery, cached: Bool + ) throws -> (GRDB.Statement, SQLiteQueryDecoder) { + let statement: GRDB.Statement + if cached { + statement = try cachedStatement(sql: prepared.sql) + sqlite3_reset(statement.sqliteStatement) + } else { + statement = try makeStatement(sql: prepared.sql) + } + for (index, binding) in zip(Int32(1)..., prepared.bindings) { try binding.bind(to: statement.sqliteStatement, at: index) } return ( diff --git a/Sources/SQLiteData/StructuredQueries+GRDB/Statement+GRDB.swift b/Sources/SQLiteData/StructuredQueries+GRDB/Statement+GRDB.swift index 862a41e3..b9cf623d 100644 --- a/Sources/SQLiteData/StructuredQueries+GRDB/Statement+GRDB.swift +++ b/Sources/SQLiteData/StructuredQueries+GRDB/Statement+GRDB.swift @@ -18,7 +18,7 @@ extension StructuredQueriesCore.Statement { /// - Parameter db: A database connection. @inlinable public func execute(_ db: Database) throws where QueryValue == () { - try QueryVoidCursor(db: db, query: query).next() + try QueryVoidCursor(db: db, query: query, cached: true).next() } /// Returns an array of all values fetched from the database. @@ -41,7 +41,7 @@ extension StructuredQueriesCore.Statement { @inlinable public func fetchAll(_ db: Database) throws -> [QueryValue.QueryOutput] where QueryValue: QueryRepresentable { - let cursor = try QueryValueCursor(db: db, query: query) + let cursor = try QueryValueCursor(db: db, query: query, cached: true) var output: [QueryValue.QueryOutput] = [] try cursor.forEach { output.append($0) } return output @@ -69,7 +69,7 @@ extension StructuredQueriesCore.Statement { @inlinable public func fetchOne(_ db: Database) throws -> QueryValue.QueryOutput? where QueryValue: QueryRepresentable { - try fetchCursor(db).next() + try QueryValueCursor(db: db, query: query, cached: true).next() } /// Returns a cursor to all values fetched from the database. @@ -92,7 +92,7 @@ extension StructuredQueriesCore.Statement { @inlinable public func fetchCursor(_ db: Database) throws -> QueryCursor where QueryValue: QueryRepresentable { - try QueryValueCursor(db: db, query: query) + try QueryValueCursor(db: db, query: query, cached: false) } } @@ -108,7 +108,7 @@ extension StructuredQueriesCore.Statement { _ db: Database ) throws -> [(repeat (each Value).QueryOutput)] where QueryValue == (repeat each Value) { - let cursor = try fetchCursor(db) + let cursor = try QueryPackCursor(db: db, query: query, cached: true) return try Array(cursor) } @@ -122,7 +122,7 @@ extension StructuredQueriesCore.Statement { _ db: Database ) throws -> (repeat (each Value).QueryOutput)? where QueryValue == (repeat each Value) { - let cursor = try fetchCursor(db) + let cursor = try QueryPackCursor(db: db, query: query, cached: true) return try cursor.next() } @@ -136,7 +136,7 @@ extension StructuredQueriesCore.Statement { _ db: Database ) throws -> QueryCursor<(repeat (each Value).QueryOutput)> where QueryValue == (repeat each Value) { - try QueryPackCursor(db: db, query: query) + try QueryPackCursor(db: db, query: query, cached: false) } } @@ -160,7 +160,7 @@ extension SelectStatement where QueryValue == (), Joins == () { @_documentation(visibility: private) @inlinable public func fetchAll(_ db: Database) throws -> [From.QueryOutput] { - let cursor = try QueryValueCursor(db: db, query: query) + let cursor = try QueryValueCursor(db: db, query: query, cached: true) var output: [From.QueryOutput] = [] try cursor.forEach { output.append($0) } return output @@ -173,7 +173,7 @@ extension SelectStatement where QueryValue == (), Joins == () { @_documentation(visibility: private) @inlinable public func fetchOne(_ db: Database) throws -> From.QueryOutput? { - try asSelect().limit(1).fetchCursor(db).next() + try QueryValueCursor(db: db, query: asSelect().limit(1).query, cached: true).next() } /// Returns a cursor to all values fetched from the database. @@ -183,7 +183,7 @@ extension SelectStatement where QueryValue == (), Joins == () { @_documentation(visibility: private) @inlinable public func fetchCursor(_ db: Database) throws -> QueryCursor { - try QueryValueCursor(db: db, query: query) + try QueryValueCursor(db: db, query: query, cached: false) } } @@ -218,7 +218,7 @@ extension SelectStatement where QueryValue == () { _ db: Database ) throws -> [(From.QueryOutput, repeat (each J).QueryOutput)] where Joins == (repeat each J) { - try Array(fetchCursor(db)) + try Array(QueryPackCursor(db: db, query: query, cached: true)) } /// Returns a single value fetched from the database. @@ -231,7 +231,10 @@ extension SelectStatement where QueryValue == () { _ db: Database ) throws -> (From.QueryOutput, repeat (each J).QueryOutput)? where Joins == (repeat each J) { - try asSelect().limit(1).fetchCursor(db).next() + try QueryPackCursor( + db: db, query: asSelect().limit(1).query, cached: true + ) + .next() } /// Returns a cursor to all values fetched from the database. @@ -244,6 +247,6 @@ extension SelectStatement where QueryValue == () { _ db: Database ) throws -> QueryCursor<(From.QueryOutput, repeat (each J).QueryOutput)> where Joins == (repeat each J) { - try QueryPackCursor(db: db, query: query) + try QueryPackCursor(db: db, query: query, cached: false) } } diff --git a/Sources/SQLiteData/StructuredQueries+GRDB/Statement+Sections.swift b/Sources/SQLiteData/StructuredQueries+GRDB/Statement+Sections.swift index 1d70822c..73a53176 100644 --- a/Sources/SQLiteData/StructuredQueries+GRDB/Statement+Sections.swift +++ b/Sources/SQLiteData/StructuredQueries+GRDB/Statement+Sections.swift @@ -141,5 +141,7 @@ private func sectionedResults ResultsSectionCollection where Key.QueryOutput: Hashable { - try ResultsSectionCollection(cursor: QuerySectionedCursor(db: db, query: query)) + try ResultsSectionCollection( + cursor: QuerySectionedCursor(db: db, query: query, cached: true) + ) } diff --git a/Tests/SQLiteDataTests/StatementCacheTests.swift b/Tests/SQLiteDataTests/StatementCacheTests.swift new file mode 100644 index 00000000..2fa14f6d --- /dev/null +++ b/Tests/SQLiteDataTests/StatementCacheTests.swift @@ -0,0 +1,76 @@ +import DependenciesTestSupport +import Foundation +import SQLiteData +import Testing + +@Suite(.dependency(\.defaultDatabase, try .database())) +struct StatementCacheTests { + @Dependency(\.defaultDatabase) var database + + @Test func `rebinding across cached statement reuse`() throws { + try database.read { db in + for id in 1...10 { + let record = try Record.where { $0.id.eq(id) }.fetchOne(db) + #expect(record?.id == id) + #expect(record?.value == "value \(id)") + } + } + } + + @Test func `escaping cursor does not share cached statement`() throws { + try database.read { db in + let query = Record.where { $0.id.lte(3) } + let cursor = try query.fetchCursor(db) + #expect(try cursor.next()?.id == 1) + #expect(try query.fetchAll(db).map(\.id) == [1, 2, 3]) + #expect(try cursor.next()?.id == 2) + #expect(try cursor.next()?.id == 3) + #expect(try cursor.next() == nil) + } + } + + @Test func `cached execute reuse`() throws { + try database.write { db in + for id in 11...20 { + try Record.insert { Record(id: id, value: "value \(id)") }.execute(db) + } + #expect(try Record.all.fetchCount(db) == 20) + } + } + + @Test func `schema change between cached uses`() throws { + try database.write { db in + #expect(try Record.all.fetchCount(db) == 10) + try #sql(#"ALTER TABLE "records" ADD COLUMN "extra" TEXT"#).execute(db) + #expect(try Record.all.fetchCount(db) == 10) + #expect(try Record.where { $0.id.eq(5) }.fetchOne(db)?.value == "value 5") + } + } +} + +@Table +private struct Record: Equatable { + let id: Int + var value: String +} + +extension DatabaseWriter where Self == DatabaseQueue { + fileprivate static func database() throws -> DatabaseQueue { + let database = try DatabaseQueue() + try database.write { db in + try #sql( + """ + CREATE TABLE "records" ( + "id" INTEGER PRIMARY KEY, + "value" TEXT NOT NULL + ) + """ + ) + .execute(db) + for id in 1...10 { + try Record.insert { Record(id: id, value: "value \(id)") }.execute(db) + } + } + return database + } +}