From d4c618ced8fe7b6699939169a1025d439d15191e Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Tue, 24 Mar 2026 01:51:16 -0700 Subject: [PATCH 1/2] Overhaul fetch APIs --- README.md | 41 +- Sources/iMessage/Database.swift | 502 ++++++++++++++++++------ Sources/iMessage/FetchRequest.swift | 129 ++++++ Tests/iMessageTests/DatabaseTests.swift | 107 +++++ 4 files changed, 651 insertions(+), 128 deletions(-) create mode 100644 Sources/iMessage/FetchRequest.swift diff --git a/README.md b/README.md index 7d12d68..873aa69 100644 --- a/README.md +++ b/README.md @@ -26,7 +26,7 @@ Add Madrid as a dependency to your `Package.swift`: ```swift dependencies: [ - .package(url: "https://github.com/mattt/Madrid.git", from: "0.2.0") + .package(url: "https://github.com/mattt/Madrid.git", from: "0.4.0") ] ``` @@ -54,25 +54,34 @@ import iMessage // Create a database (uses `~/Library/Messages/chat.db` by default) let db = try iMessage.Database() -// Fetch recent messages -let recentMessages = try db.fetchMessages(limit: 10) - -// Fetch messages from select individuals in time range -let pastWeek = Date.now.addingTimeInterval(-7*24*60*60)..( + predicate: .and([ + .participantHandles([ + "johnny.appleseed@mac.com", + "+18002752273", + ]), + .dateRange(pastWeek), + ]), + sortDescriptors: [.date(.descending)], + limit: 10 +) + +// Execute request +for message in try db.fetch(request) { print("From: \(message.sender)") - print("Content: \(message.content)") - print("Sent at: \(message.timestamp)") + print("Content: \(message.text)") + print("Sent at: \(message.date)") } ``` +> [!TIP] +> Legacy convenience APIs +> (`fetchMessages(for:with:in:limit:)`, `fetchChats(with:in:limit:)`) +> are still available, +> but are deprecated in favor of `fetch(_:)`. + ### Decoding TypedStream Data ```swift diff --git a/Sources/iMessage/Database.swift b/Sources/iMessage/Database.swift index 9d6fdeb..5aea472 100644 --- a/Sources/iMessage/Database.swift +++ b/Sources/iMessage/Database.swift @@ -58,6 +58,11 @@ public final class Database { case queryError(String) } + /// Backward-compatible alias for a message fetch request. + public typealias MessageFetchRequest = FetchRequest + /// Backward-compatible alias for a chat fetch request. + public typealias ChatFetchRequest = FetchRequest + private init( _ filename: String, flags: Flags = .default @@ -132,60 +137,24 @@ public final class Database { return results } - /// Fetches chats, optionally filtered by participants and date range. + /// Fetches chats using a predicate-style request. /// - /// - Parameters: - /// - participantHandles: An optional set of handles that must be present in the chat. - /// - dateRange: An optional date range used to filter chat activity. - /// - limit: The maximum number of chats to return. - /// - Returns: Chats ordered by most recent message date, descending. + /// - Parameter request: The typed chat fetch request. + /// - Returns: Chats matching the predicate and sort descriptors. /// - Throws: ``Error/queryError(_:)`` when SQL preparation or execution fails. - public func fetchChats( - with participantHandles: Set? = nil, - in dateRange: Range? = nil, - limit: Int = 100 - ) throws -> [Chat] { - try withTransaction { - var conditions: [String] = [] - var parameters: [any Bindable] = [] - - // Add date range if specified - if let dateRange = dateRange { - if let upperBound = dateRange.upperBound.nanosecondsSinceReferenceDate { - conditions.append("m.date < ?") - parameters.append(Int64(upperBound)) - } - if let lowerBound = dateRange.lowerBound.nanosecondsSinceReferenceDate { - conditions.append("m.date >= ?") - parameters.append(Int64(lowerBound)) - } - } + public func fetch(_ request: ChatFetchRequest) throws -> [Chat] { + try validatePagination(limit: request.limit, offset: request.offset) - // Add participants filter if specified - if let handles = participantHandles, !handles.isEmpty { - conditions.append( - """ - c.ROWID IN ( - SELECT chat_id - FROM chat_handle_join chj - JOIN handle h ON chj.handle_id = h.ROWID - WHERE h.id IN (\(String(repeating: "?,", count: handles.count).dropLast())) - GROUP BY chat_id - HAVING COUNT(DISTINCT handle_id) = ? - ) - """ - ) + return try withTransaction { + let compiledPredicate = try compileChatPredicate(request.predicate) + let orderByClause = chatOrderByClause(request.sortDescriptors) - // Add each participant as a value - handles.forEach { handle in - parameters.append(handle.rawValue) - } - // Add the count of participants - parameters.append(Int32(handles.count)) - } + var parameters = compiledPredicate.parameters + parameters.append(try bindableInt32(request.limit, name: "limit")) + parameters.append(try bindableInt32(request.offset, name: "offset")) let query = """ - SELECT + SELECT c.guid, c.display_name, c.service_name, @@ -193,25 +162,24 @@ public final class Database { FROM chat c LEFT JOIN chat_message_join cmj ON c.ROWID = cmj.chat_id LEFT JOIN message m ON cmj.message_id = m.ROWID - \(conditions.isEmpty ? "" : "WHERE \(conditions.joined(separator: " AND "))") + \(compiledPredicate.whereClause.map { "WHERE \($0)" } ?? "") GROUP BY c.ROWID - ORDER BY last_message_date DESC + ORDER BY \(orderByClause) LIMIT ? + OFFSET ? """ - parameters.append(Int32(limit)) - return try execute(query, parameters: parameters) { statement in - // Safely handle potentially null columns guard let guidText = sqlite3_column_text(statement, 0) else { return nil } let chatId = Chat.ID(rawValue: String(cString: guidText)) let displayName = sqlite3_column_text(statement, 1).map { String(cString: $0) } - let lastMessageDate = Date( - nanosecondsSinceReferenceDate: sqlite3_column_int64(statement, 3) - ) + let rawLastMessageDate = sqlite3_column_int64(statement, 3) + let lastMessageDate = + sqlite3_column_type(statement, 3) == SQLITE_NULL + ? nil + : Date(nanosecondsSinceReferenceDate: rawLastMessageDate) - // Fetch participants for this chat let participants = try fetchParticipants(for: chatId) return Chat( @@ -224,64 +192,63 @@ public final class Database { } } - /// Fetches messages, optionally filtered by chat, participants, and date range. + @available(*, deprecated, message: "Use fetch(_:) with a ChatFetchRequest.") + public func fetchChats(_ request: ChatFetchRequest) throws -> [Chat] { + try fetch(request) + } + + /// Fetches chats, optionally filtered by participants and date range. /// /// - Parameters: - /// - chatId: An optional chat identifier to scope the query. - /// - participantHandles: An optional set of sender handles to include. - /// - dateRange: An optional date range used to filter message dates. - /// - limit: The maximum number of messages to return. - /// - Returns: Messages ordered by message date, descending. + /// - participantHandles: An optional set of handles that must be present in the chat. + /// - dateRange: An optional date range used to filter chat activity. + /// - limit: The maximum number of chats to return. + /// - Returns: Chats ordered by most recent message date, descending. /// - Throws: ``Error/queryError(_:)`` when SQL preparation or execution fails. - public func fetchMessages( - for chatId: Chat.ID? = nil, + @available( + *, + deprecated, + message: "Use fetch(_:) with a ChatFetchRequest predicate instead." + ) + public func fetchChats( with participantHandles: Set? = nil, in dateRange: Range? = nil, limit: Int = 100 - ) throws -> [Message] { - try withTransaction { - var conditions: [String] = [] - var parameters: [any Bindable] = [] - - // Add chat filter if specified - if let chatId = chatId { - conditions.append("c.guid = ?") - parameters.append(chatId.rawValue) - } + ) throws -> [Chat] { + var predicates: [ChatPredicate] = [] + if let participantHandles = participantHandles, !participantHandles.isEmpty { + predicates.append(.participantHandles(participantHandles, match: .all)) + } + if let dateRange = dateRange { + predicates.append(.dateRange(dateRange)) + } - // Add participants filter if specified - if let handles = participantHandles, !handles.isEmpty { - conditions.append( - """ - m.ROWID IN ( - SELECT m.ROWID - FROM message m - JOIN handle h ON m.handle_id = h.ROWID - WHERE h.id IN (\(String(repeating: "?,", count: handles.count).dropLast())) - ) - """ - ) + return try fetch( + ChatFetchRequest( + predicate: .and(predicates), + limit: limit + ) + ) + } - // Add each participant as a value - handles.forEach { handle in - parameters.append(handle.rawValue) - } - } + /// Fetches messages using a predicate-style request. + /// + /// - Parameter request: The typed message fetch request. + /// - Returns: Messages matching the predicate and sort descriptors. + /// - Throws: ``Error/queryError(_:)`` when SQL preparation or execution fails. + public func fetch(_ request: MessageFetchRequest) throws -> [Message] { + try validatePagination(limit: request.limit, offset: request.offset) - // Add date range if specified - if let dateRange = dateRange { - if let upperBound = dateRange.upperBound.nanosecondsSinceReferenceDate { - conditions.append("m.date < ?") - parameters.append(upperBound) - } - if let lowerBound = dateRange.lowerBound.nanosecondsSinceReferenceDate { - conditions.append("m.date >= ?") - parameters.append(lowerBound) - } - } + return try withTransaction { + let compiledPredicate = try compileMessagePredicate(request.predicate) + let orderByClause = messageOrderByClause(request.sortDescriptors) + + var parameters = compiledPredicate.parameters + parameters.append(try bindableInt32(request.limit, name: "limit")) + parameters.append(try bindableInt32(request.offset, name: "offset")) let query = """ - SELECT + SELECT m.guid, m.text, HEX(m.attributedBody), @@ -291,26 +258,23 @@ public final class Database { m.service, m.date_read FROM message m - \(chatId != nil ? "JOIN chat_message_join cmj ON m.ROWID = cmj.message_id" : "") - \(chatId != nil ? "JOIN chat c ON cmj.chat_id = c.ROWID" : "") + \(compiledPredicate.requiresChatJoin ? "JOIN chat_message_join cmj ON m.ROWID = cmj.message_id" : "") + \(compiledPredicate.requiresChatJoin ? "JOIN chat c ON cmj.chat_id = c.ROWID" : "") LEFT JOIN handle h ON m.handle_id = h.ROWID - \(conditions.isEmpty ? "" : "WHERE \(conditions.joined(separator: " AND "))") - ORDER BY m.date DESC + \(compiledPredicate.whereClause.map { "WHERE \($0)" } ?? "") + ORDER BY \(orderByClause) LIMIT ? + OFFSET ? """ - parameters.append(Int32(limit)) - return try execute(query, parameters: parameters) { statement in let messageID: Message.ID if let messageIdText = sqlite3_column_text(statement, 0) { messageID = Message.ID(rawValue: String(cString: messageIdText)) } else { messageID = "N/A" - // FIXME } - // Handle text let text: String if let rawText = sqlite3_column_text(statement, 1) { text = String(cString: rawText) @@ -352,6 +316,311 @@ public final class Database { } } + @available(*, deprecated, message: "Use fetch(_:) with a MessageFetchRequest.") + public func fetchMessages(_ request: MessageFetchRequest) throws -> [Message] { + try fetch(request) + } + + /// Fetches messages, optionally filtered by chat, participants, and date range. + /// + /// - Parameters: + /// - chatId: An optional chat identifier to scope the query. + /// - participantHandles: An optional set of sender handles to include. + /// - dateRange: An optional date range used to filter message dates. + /// - limit: The maximum number of messages to return. + /// - Returns: Messages ordered by message date, descending. + /// - Throws: ``Error/queryError(_:)`` when SQL preparation or execution fails. + @available( + *, + deprecated, + message: "Use fetch(_:) with a MessageFetchRequest predicate instead." + ) + public func fetchMessages( + for chatId: Chat.ID? = nil, + with participantHandles: Set? = nil, + in dateRange: Range? = nil, + limit: Int = 100 + ) throws -> [Message] { + var predicates: [MessagePredicate] = [] + if let chatId = chatId { + predicates.append(.chatID(chatId)) + } + if let participantHandles = participantHandles, !participantHandles.isEmpty { + predicates.append(.participantHandles(participantHandles)) + } + if let dateRange = dateRange { + predicates.append(.dateRange(dateRange)) + } + + return try fetch( + MessageFetchRequest( + predicate: .and(predicates), + limit: limit + ) + ) + } + + private struct CompiledPredicate { + let whereClause: String? + let parameters: [any Bindable] + let requiresChatJoin: Bool + } + + private func compileMessagePredicate( + _ predicate: MessagePredicate + ) throws -> CompiledPredicate { + switch predicate { + case .all: + return CompiledPredicate(whereClause: nil, parameters: [], requiresChatJoin: false) + case .none: + return CompiledPredicate(whereClause: "1 = 0", parameters: [], requiresChatJoin: false) + case .chatID(let chatID): + return CompiledPredicate( + whereClause: "c.guid = ?", + parameters: [chatID.rawValue], + requiresChatJoin: true + ) + case .participantHandles(let handles): + if handles.isEmpty { + return CompiledPredicate( + whereClause: "1 = 0", + parameters: [], + requiresChatJoin: false + ) + } + let handleValues = orderedHandleValues(handles) + let placeholders = placeholders(handles.count) + let condition = """ + m.ROWID IN ( + SELECT m2.ROWID + FROM message m2 + JOIN handle h ON m2.handle_id = h.ROWID + WHERE h.id IN (\(placeholders)) + ) + """ + return CompiledPredicate( + whereClause: condition, + parameters: toBindableStrings(handleValues), + requiresChatJoin: false + ) + case .dateRange(let dateRange): + let upperBound = try requireNanoseconds(dateRange.upperBound, label: "upperBound") + let lowerBound = try requireNanoseconds(dateRange.lowerBound, label: "lowerBound") + return CompiledPredicate( + whereClause: "m.date < ? AND m.date >= ?", + parameters: [upperBound, lowerBound], + requiresChatJoin: false + ) + case .and(let predicates): + if predicates.isEmpty { + return CompiledPredicate(whereClause: nil, parameters: [], requiresChatJoin: false) + } + return try combineMessagePredicates(predicates, joiner: "AND") + case .or(let predicates): + if predicates.isEmpty { + return CompiledPredicate(whereClause: "1 = 0", parameters: [], requiresChatJoin: false) + } + return try combineMessagePredicates(predicates, joiner: "OR") + case .not(let predicate): + let compiled = try compileMessagePredicate(predicate) + let whereClause = compiled.whereClause ?? "1 = 1" + return CompiledPredicate( + whereClause: "NOT (\(whereClause))", + parameters: compiled.parameters, + requiresChatJoin: compiled.requiresChatJoin + ) + } + } + + private func combineMessagePredicates( + _ predicates: [MessagePredicate], + joiner: String + ) throws -> CompiledPredicate { + var whereParts: [String] = [] + var parameters: [any Bindable] = [] + var requiresChatJoin = false + + for predicate in predicates { + let compiled = try compileMessagePredicate(predicate) + if let whereClause = compiled.whereClause { + whereParts.append("(\(whereClause))") + parameters.append(contentsOf: compiled.parameters) + } + requiresChatJoin = requiresChatJoin || compiled.requiresChatJoin + } + + let joinedWhereClause = whereParts.isEmpty ? nil : whereParts.joined(separator: " \(joiner) ") + return CompiledPredicate( + whereClause: joinedWhereClause, + parameters: parameters, + requiresChatJoin: requiresChatJoin + ) + } + + private func compileChatPredicate( + _ predicate: ChatPredicate + ) throws -> CompiledPredicate { + switch predicate { + case .all: + return CompiledPredicate(whereClause: nil, parameters: [], requiresChatJoin: false) + case .none: + return CompiledPredicate(whereClause: "1 = 0", parameters: [], requiresChatJoin: false) + case .participantHandles(let handles, let match): + if handles.isEmpty { + let whereClause = match == .all ? nil : "1 = 0" + return CompiledPredicate( + whereClause: whereClause, + parameters: [], + requiresChatJoin: false + ) + } + + let handleValues = orderedHandleValues(handles) + let placeholders = placeholders(handles.count) + switch match { + case .any: + let condition = """ + c.ROWID IN ( + SELECT chat_id + FROM chat_handle_join chj + JOIN handle h ON chj.handle_id = h.ROWID + WHERE h.id IN (\(placeholders)) + ) + """ + return CompiledPredicate( + whereClause: condition, + parameters: toBindableStrings(handleValues), + requiresChatJoin: false + ) + case .all: + let condition = """ + c.ROWID IN ( + SELECT chat_id + FROM chat_handle_join chj + JOIN handle h ON chj.handle_id = h.ROWID + WHERE h.id IN (\(placeholders)) + GROUP BY chat_id + HAVING COUNT(DISTINCT handle_id) = ? + ) + """ + var parameters = toBindableStrings(handleValues) + parameters.append(try bindableInt32(handles.count, name: "participant count")) + return CompiledPredicate( + whereClause: condition, + parameters: parameters, + requiresChatJoin: false + ) + } + case .dateRange(let dateRange): + let upperBound = try requireNanoseconds(dateRange.upperBound, label: "upperBound") + let lowerBound = try requireNanoseconds(dateRange.lowerBound, label: "lowerBound") + return CompiledPredicate( + whereClause: "m.date < ? AND m.date >= ?", + parameters: [upperBound, lowerBound], + requiresChatJoin: false + ) + case .and(let predicates): + if predicates.isEmpty { + return CompiledPredicate(whereClause: nil, parameters: [], requiresChatJoin: false) + } + return try combineChatPredicates(predicates, joiner: "AND") + case .or(let predicates): + if predicates.isEmpty { + return CompiledPredicate(whereClause: "1 = 0", parameters: [], requiresChatJoin: false) + } + return try combineChatPredicates(predicates, joiner: "OR") + case .not(let predicate): + let compiled = try compileChatPredicate(predicate) + let whereClause = compiled.whereClause ?? "1 = 1" + return CompiledPredicate( + whereClause: "NOT (\(whereClause))", + parameters: compiled.parameters, + requiresChatJoin: false + ) + } + } + + private func combineChatPredicates( + _ predicates: [ChatPredicate], + joiner: String + ) throws -> CompiledPredicate { + var whereParts: [String] = [] + var parameters: [any Bindable] = [] + + for predicate in predicates { + let compiled = try compileChatPredicate(predicate) + if let whereClause = compiled.whereClause { + whereParts.append("(\(whereClause))") + parameters.append(contentsOf: compiled.parameters) + } + } + + return CompiledPredicate( + whereClause: whereParts.isEmpty ? nil : whereParts.joined(separator: " \(joiner) "), + parameters: parameters, + requiresChatJoin: false + ) + } + + private func messageOrderByClause(_ descriptors: [MessageSortDescriptor]) -> String { + let descriptors = descriptors.isEmpty ? [.date(.descending), .id(.descending)] : descriptors + return descriptors.map { descriptor in + switch descriptor { + case .date(let order): + return "m.date \(order.sqlKeyword)" + case .id(let order): + return "m.guid \(order.sqlKeyword)" + } + }.joined(separator: ", ") + } + + private func chatOrderByClause(_ descriptors: [ChatSortDescriptor]) -> String { + let descriptors = descriptors.isEmpty ? [.lastMessageDate(.descending), .id(.ascending)] : descriptors + return descriptors.map { descriptor in + switch descriptor { + case .lastMessageDate(let order): + return "last_message_date \(order.sqlKeyword)" + case .id(let order): + return "c.guid \(order.sqlKeyword)" + } + }.joined(separator: ", ") + } + + private func placeholders(_ count: Int) -> String { + Array(repeating: "?", count: count).joined(separator: ",") + } + + private func orderedHandleValues(_ handles: Set) -> [String] { + handles.map(\.rawValue).sorted() + } + + private func toBindableStrings(_ values: [String]) -> [any Bindable] { + values.map { $0 as any Bindable } + } + + private func requireNanoseconds(_ date: Date, label: String) throws -> Int64 { + guard let value = date.nanosecondsSinceReferenceDate else { + throw Error.queryError("Could not represent \(label) as Int64 nanoseconds.") + } + return value + } + + private func validatePagination(limit: Int, offset: Int) throws { + guard limit >= 0 else { + throw Error.queryError("limit must be >= 0") + } + guard offset >= 0 else { + throw Error.queryError("offset must be >= 0") + } + } + + private func bindableInt32(_ value: Int, name: String) throws -> Int32 { + guard value <= Int(Int32.max), value >= Int(Int32.min) else { + throw Error.queryError("\(name) is out of Int32 range.") + } + return Int32(value) + } + /// Fetches participants for a chat. /// /// - Parameters: @@ -461,6 +730,15 @@ public final class Database { // MARK: - +private extension SortOrder { + var sqlKeyword: String { + switch self { + case .ascending: "ASC" + case .descending: "DESC" + } + } +} + private protocol Bindable { func bind(to statement: OpaquePointer, at index: Int32) } diff --git a/Sources/iMessage/FetchRequest.swift b/Sources/iMessage/FetchRequest.swift new file mode 100644 index 0000000..c87289d --- /dev/null +++ b/Sources/iMessage/FetchRequest.swift @@ -0,0 +1,129 @@ +import Foundation + +/// Sort direction for fetch requests. +public enum SortOrder: String, Sendable, Hashable, CaseIterable { + /// Sort values from smallest to largest. + case ascending + /// Sort values from largest to smallest. + case descending +} + +/// Participant matching mode for participant-based predicates. +public enum ParticipantMatch: String, Sendable, Hashable, CaseIterable { + /// Match when any provided participant is present. + case any + /// Match only when all provided participants are present. + case all +} + +/// Predicate tree used to filter messages. +public indirect enum MessagePredicate: Sendable, Hashable { + /// Match all messages. + case all + /// Match no messages. + case none + /// Match messages that belong to the specified chat. + case chatID(Chat.ID) + /// Match messages sent by any of the provided handles. + case participantHandles(Set) + /// Match messages in the half-open date range. + case dateRange(Range) + /// Match messages that satisfy every nested predicate. + case and([MessagePredicate]) + /// Match messages that satisfy at least one nested predicate. + case or([MessagePredicate]) + /// Match messages that do not satisfy the nested predicate. + case not(MessagePredicate) +} + +/// Typed message sort descriptor. +public enum MessageSortDescriptor: Sendable, Hashable { + /// Sort by message date. + case date(SortOrder) + /// Sort by stable message identifier. + case id(SortOrder) +} + +/// Predicate tree used to filter chats. +public indirect enum ChatPredicate: Sendable, Hashable { + /// Match all chats. + case all + /// Match no chats. + case none + /// Match chats by participant handles using the selected mode. + case participantHandles(Set, match: ParticipantMatch) + /// Match chats that contain message activity in the half-open date range. + case dateRange(Range) + /// Match chats that satisfy every nested predicate. + case and([ChatPredicate]) + /// Match chats that satisfy at least one nested predicate. + case or([ChatPredicate]) + /// Match chats that do not satisfy the nested predicate. + case not(ChatPredicate) +} + +/// Typed chat sort descriptor. +public enum ChatSortDescriptor: Sendable, Hashable { + /// Sort by each chat's latest message date. + case lastMessageDate(SortOrder) + /// Sort by stable chat identifier. + case id(SortOrder) +} + +/// Describes result-specific behavior for ``FetchRequest``. +public protocol FetchRequestResult { + /// Predicate type accepted for this result type. + associatedtype Predicate: Sendable + /// Sort descriptor type accepted for this result type. + associatedtype SortDescriptor: Sendable + + /// Default predicate when callers do not provide one. + static var defaultFetchPredicate: Predicate { get } + /// Default sort descriptors when callers do not provide any. + static var defaultFetchSortDescriptors: [SortDescriptor] { get } +} + +/// A generic typed fetch request for queryable result models. +public struct FetchRequest: Sendable { + /// Predicate used to filter rows. + public var predicate: Result.Predicate + /// Sort descriptors applied in order. + public var sortDescriptors: [Result.SortDescriptor] + /// Maximum number of rows to return. + public var limit: Int + /// Number of rows to skip before returning. + public var offset: Int + + /// Creates a fetch request. + /// + /// - Parameters: + /// - predicate: The filter predicate to apply. + /// - sortDescriptors: The sort descriptors applied in order. + /// - limit: The maximum number of rows to return. + /// - offset: The number of rows to skip before returning. + public init( + predicate: Result.Predicate = Result.defaultFetchPredicate, + sortDescriptors: [Result.SortDescriptor] = Result.defaultFetchSortDescriptors, + limit: Int = 100, + offset: Int = 0 + ) { + self.predicate = predicate + self.sortDescriptors = sortDescriptors + self.limit = limit + self.offset = offset + } +} + +extension Message: FetchRequestResult { + public static var defaultFetchPredicate: MessagePredicate { .all } + public static var defaultFetchSortDescriptors: [MessageSortDescriptor] { + [.date(.descending), .id(.descending)] + } +} + +extension Chat: FetchRequestResult { + public static var defaultFetchPredicate: ChatPredicate { .all } + public static var defaultFetchSortDescriptors: [ChatSortDescriptor] { + [.lastMessageDate(.descending), .id(.ascending)] + } +} diff --git a/Tests/iMessageTests/DatabaseTests.swift b/Tests/iMessageTests/DatabaseTests.swift index 148d9a9..3cf263c 100644 --- a/Tests/iMessageTests/DatabaseTests.swift +++ b/Tests/iMessageTests/DatabaseTests.swift @@ -128,4 +128,111 @@ struct DatabaseTests { ) #expect(!rangeMessages.isEmpty) } + + @Test + func testMessageFetchRequestSortAndPagination() async throws { + let request = Database.MessageFetchRequest( + predicate: .all, + sortDescriptors: [ + .date(.ascending), + .id(.ascending), + ], + limit: 2, + offset: 1 + ) + + let messages = try db.fetchMessages(request) + #expect(messages.count == 2) + #expect(messages[0].id.rawValue == "msg-guid-5") + #expect(messages[1].id.rawValue == "msg-guid-1") + } + + @Test + func testMessagePredicateComposition() async throws { + let request = Database.MessageFetchRequest( + predicate: .or([ + .chatID("chat-guid-2"), + .participantHandles(["person@example.com"]), + ]), + limit: 10 + ) + + let messages = try db.fetchMessages(request) + let messageIDs = Set(messages.map(\.id.rawValue)) + #expect(messageIDs == ["msg-guid-3", "msg-guid-4", "msg-guid-5"]) + } + + @Test + func testChatParticipantMatchModes() async throws { + let anyMatchRequest = Database.ChatFetchRequest( + predicate: .participantHandles(["+1234567890", "third@example.com"], match: .any), + limit: 10 + ) + let anyMatch = try db.fetchChats(anyMatchRequest) + #expect(Set(anyMatch.map(\.id.rawValue)) == ["chat-guid-1", "chat-guid-2"]) + + let allMatchRequest = Database.ChatFetchRequest( + predicate: .participantHandles(["+1234567890", "third@example.com"], match: .all), + limit: 10 + ) + let allMatch = try db.fetchChats(allMatchRequest) + #expect(allMatch.count == 1) + #expect(allMatch[0].id.rawValue == "chat-guid-2") + } + + @Test + func testPredicateEmptyCompoundSemantics() async throws { + let allMessages = try db.fetchMessages( + Database.MessageFetchRequest(predicate: .and([]), limit: 10) + ) + #expect(allMessages.count == 5) + + let noMessages = try db.fetchMessages( + Database.MessageFetchRequest(predicate: .or([]), limit: 10) + ) + #expect(noMessages.isEmpty) + } + + @Test + func testLegacyWrappersMatchTypedRequests() async throws { + let legacyMessages = try db.fetchMessages( + for: "chat-guid-1", + with: ["+1234567890", "person@example.com"], + limit: 10 + ) + let requestMessages = try db.fetchMessages( + Database.MessageFetchRequest( + predicate: .and([ + .chatID("chat-guid-1"), + .participantHandles(["+1234567890", "person@example.com"]), + ]), + limit: 10 + ) + ) + #expect(legacyMessages.map(\.id) == requestMessages.map(\.id)) + + let legacyChats = try db.fetchChats( + with: ["+1234567890", "person@example.com"], + limit: 10 + ) + let requestChats = try db.fetchChats( + Database.ChatFetchRequest( + predicate: .participantHandles(["+1234567890", "person@example.com"], match: .all), + limit: 10 + ) + ) + #expect(legacyChats.map(\.id) == requestChats.map(\.id)) + } + + @Test + func testInvalidPaginationThrows() async throws { + do { + _ = try db.fetchMessages( + Database.MessageFetchRequest(limit: -1) + ) + Issue.record("Expected negative limit to throw") + } catch Database.Error.queryError { + // Expected. + } + } } From 40bc90f0d354cd5c76a0d158f9ab744d2241c268 Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Tue, 24 Mar 2026 02:09:49 -0700 Subject: [PATCH 2/2] Incorporate feedback from review --- Sources/iMessage/Database.swift | 8 +++ Tests/iMessageTests/DatabaseTests.swift | 96 ++++++++++++++++++++++--- 2 files changed, 95 insertions(+), 9 deletions(-) diff --git a/Sources/iMessage/Database.swift b/Sources/iMessage/Database.swift index 5aea472..44175df 100644 --- a/Sources/iMessage/Database.swift +++ b/Sources/iMessage/Database.swift @@ -436,6 +436,7 @@ public final class Database { _ predicates: [MessagePredicate], joiner: String ) throws -> CompiledPredicate { + let isOR = joiner == "OR" var whereParts: [String] = [] var parameters: [any Bindable] = [] var requiresChatJoin = false @@ -445,6 +446,9 @@ public final class Database { if let whereClause = compiled.whereClause { whereParts.append("(\(whereClause))") parameters.append(contentsOf: compiled.parameters) + } else if isOR { + // OR with a match-all branch is itself match-all. + return CompiledPredicate(whereClause: nil, parameters: [], requiresChatJoin: false) } requiresChatJoin = requiresChatJoin || compiled.requiresChatJoin } @@ -544,6 +548,7 @@ public final class Database { _ predicates: [ChatPredicate], joiner: String ) throws -> CompiledPredicate { + let isOR = joiner == "OR" var whereParts: [String] = [] var parameters: [any Bindable] = [] @@ -552,6 +557,9 @@ public final class Database { if let whereClause = compiled.whereClause { whereParts.append("(\(whereClause))") parameters.append(contentsOf: compiled.parameters) + } else if isOR { + // OR with a match-all branch is itself match-all. + return CompiledPredicate(whereClause: nil, parameters: [], requiresChatJoin: false) } } diff --git a/Tests/iMessageTests/DatabaseTests.swift b/Tests/iMessageTests/DatabaseTests.swift index 3cf263c..b6c6e61 100644 --- a/Tests/iMessageTests/DatabaseTests.swift +++ b/Tests/iMessageTests/DatabaseTests.swift @@ -141,7 +141,7 @@ struct DatabaseTests { offset: 1 ) - let messages = try db.fetchMessages(request) + let messages = try db.fetch(request) #expect(messages.count == 2) #expect(messages[0].id.rawValue == "msg-guid-5") #expect(messages[1].id.rawValue == "msg-guid-1") @@ -157,7 +157,7 @@ struct DatabaseTests { limit: 10 ) - let messages = try db.fetchMessages(request) + let messages = try db.fetch(request) let messageIDs = Set(messages.map(\.id.rawValue)) #expect(messageIDs == ["msg-guid-3", "msg-guid-4", "msg-guid-5"]) } @@ -168,29 +168,107 @@ struct DatabaseTests { predicate: .participantHandles(["+1234567890", "third@example.com"], match: .any), limit: 10 ) - let anyMatch = try db.fetchChats(anyMatchRequest) + let anyMatch = try db.fetch(anyMatchRequest) #expect(Set(anyMatch.map(\.id.rawValue)) == ["chat-guid-1", "chat-guid-2"]) let allMatchRequest = Database.ChatFetchRequest( predicate: .participantHandles(["+1234567890", "third@example.com"], match: .all), limit: 10 ) - let allMatch = try db.fetchChats(allMatchRequest) + let allMatch = try db.fetch(allMatchRequest) #expect(allMatch.count == 1) #expect(allMatch[0].id.rawValue == "chat-guid-2") } @Test func testPredicateEmptyCompoundSemantics() async throws { - let allMessages = try db.fetchMessages( + let allMessages = try db.fetch( Database.MessageFetchRequest(predicate: .and([]), limit: 10) ) #expect(allMessages.count == 5) - let noMessages = try db.fetchMessages( + let noMessages = try db.fetch( Database.MessageFetchRequest(predicate: .or([]), limit: 10) ) #expect(noMessages.isEmpty) + + let orWithAllMessages = try db.fetch( + Database.MessageFetchRequest( + predicate: .or([ + .all, + .chatID("chat-guid-1"), + ]), + limit: 10 + ) + ) + #expect(orWithAllMessages.count == 5) + + let orWithAllChats = try db.fetch( + Database.ChatFetchRequest( + predicate: .or([ + .all, + .participantHandles(["nonexistent@example.com"], match: .all), + ]), + limit: 10 + ) + ) + #expect(orWithAllChats.count == 2) + + let nestedOrWithAllMessages = try db.fetch( + Database.MessageFetchRequest( + predicate: .or([ + .and([]), + .chatID("chat-guid-1"), + ]), + limit: 10 + ) + ) + #expect(nestedOrWithAllMessages.count == 5) + + let nestedOrWithAllChats = try db.fetch( + Database.ChatFetchRequest( + predicate: .or([ + .and([.all]), + .none, + ]), + limit: 10 + ) + ) + #expect(nestedOrWithAllChats.count == 2) + } + + @Test + func testOrWithMatchAllDoesNotDuplicateMessagesFromChatJoinRows() async throws { + // Create a duplicate chat join row for one message. + try db.execute( + """ + INSERT INTO chat_message_join (chat_id, message_id) + VALUES (2, 1); + """ + ) + + let allMessages = try db.fetch( + Database.MessageFetchRequest( + predicate: .all, + sortDescriptors: [.id(.ascending)], + limit: 20 + ) + ) + #expect(allMessages.count == 5) + + let orWithAll = try db.fetch( + Database.MessageFetchRequest( + predicate: .or([ + .all, + .chatID("chat-guid-1"), + ]), + sortDescriptors: [.id(.ascending)], + limit: 20 + ) + ) + + #expect(orWithAll.count == 5) + #expect(Set(orWithAll.map(\.id)) == Set(allMessages.map(\.id))) } @Test @@ -200,7 +278,7 @@ struct DatabaseTests { with: ["+1234567890", "person@example.com"], limit: 10 ) - let requestMessages = try db.fetchMessages( + let requestMessages = try db.fetch( Database.MessageFetchRequest( predicate: .and([ .chatID("chat-guid-1"), @@ -215,7 +293,7 @@ struct DatabaseTests { with: ["+1234567890", "person@example.com"], limit: 10 ) - let requestChats = try db.fetchChats( + let requestChats = try db.fetch( Database.ChatFetchRequest( predicate: .participantHandles(["+1234567890", "person@example.com"], match: .all), limit: 10 @@ -227,7 +305,7 @@ struct DatabaseTests { @Test func testInvalidPaginationThrows() async throws { do { - _ = try db.fetchMessages( + _ = try db.fetch( Database.MessageFetchRequest(limit: -1) ) Issue.record("Expected negative limit to throw")