Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
- Literal backticks in cloudflared, cloud-sql-proxy, SSH config, remote command and dump tool install messages.
- Tunnel command preview showing port 0 or the wrong host when Port is blank or the connection uses a host list.
- SSH tab host-list warning naming replica set failover for Redis and Kafka, and implying Sentinel works through a tunnel.
- MongoDB exports dropping fields first seen after the 200th document, and exporting a field null in the first 200 as text.
- iPhone and iPad reading a Safe Mode level they do not recognize from iCloud as Off.

### Security
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -86,4 +86,25 @@ struct SyncRecordMapperTests {
let decoded = try #require(SyncRecordMapper.toConnection(record))
#expect(decoded.safeModeLevel == .confirmWrites)
}

@Test("An unrecognized wire value keeps the legacy read-only restriction")
func unknownWireValuePreservesReadOnly() throws {
let record = makeRawRecord(safeModeLevelRaw: "someFutureLevel", isReadOnly: true)
let decoded = try #require(SyncRecordMapper.toConnection(record))
#expect(decoded.safeModeLevel == .readOnly)
}

@Test("A rename preserves an unrecognized wire value and requires confirmation")
func renamePreservesUnknownWireValue() throws {
let record = makeRawRecord(safeModeLevelRaw: "someFutureLevel")
var connection = try #require(SyncRecordMapper.toConnection(record))
connection.name = "Renamed"

SyncRecordMapper.updateRecord(record, with: connection)

#expect(record["safeModeLevel"] as? String == "someFutureLevel")
let decoded = try #require(SyncRecordMapper.toConnection(record))
#expect(decoded.name == "Renamed")
#expect(decoded.safeModeLevel == .confirmWrites)
}
}
2 changes: 1 addition & 1 deletion Plugins/MongoDBDriverPlugin/BsonDocumentFlattener.swift
Original file line number Diff line number Diff line change
Expand Up @@ -461,7 +461,7 @@ struct BsonDocumentFlattener {
return counts.max(by: { $0.value < $1.value })?.key ?? .string
}

private static func valueKind(for value: Any, representation: MongoDBUuidRepresentation) -> BsonValueKind {
static func valueKind(for value: Any, representation: MongoDBUuidRepresentation) -> BsonValueKind {
if value is NSNull { return .null }

switch value {
Expand Down
86 changes: 61 additions & 25 deletions Plugins/MongoDBDriverPlugin/MongoDBConnection+SyncHelpers.swift
Original file line number Diff line number Diff line change
Expand Up @@ -153,7 +153,8 @@ extension MongoDBConnection {
}

func aggregateSync(
client: OpaquePointer, database: String, collection: String, pipeline: String
client: OpaquePointer, database: String, collection: String, pipeline: String,
optionsJson: String? = nil
) throws -> (docs: [[String: Any]], isTruncated: Bool) {
try checkCancelled()

Expand All @@ -167,7 +168,9 @@ extension MongoDBConnection {

let timeoutMS = queryTimeoutMS
var optsBson: OpaquePointer?
if timeoutMS > 0 {
if let optionsJson {
optsBson = jsonToBson(optionsJson)
} else if timeoutMS > 0 {
optsBson = jsonToBson("{\"maxTimeMS\": \(timeoutMS)}")
}
defer { if let opts = optsBson { bson_destroy(opts) } }
Expand Down Expand Up @@ -415,12 +418,14 @@ extension MongoDBConnection {
func iterateCursorStreaming(
cursor: OpaquePointer,
continuation: AsyncThrowingStream<PluginStreamElement, Error>.Continuation,
streamState: MongoStreamState
streamState: MongoStreamState,
census: () throws -> MongoFieldCensus?
) {
var docPtr: OpaquePointer?
var sample: [[String: Any]] = []
var projection: MongoStreamProjection?
var emitted = 0
var unannounced = Set<String>()

while mongoc_cursor_next(cursor, &docPtr) {
if Task.isCancelled {
Expand All @@ -432,24 +437,34 @@ extension MongoDBConnection {
guard let doc = docPtr else { continue }
let dict = bsonToDict(doc)

if let projection {
continuation.yield(.rows([projection.row(for: dict, convert: streamCellValue)]))
emitted += 1
if emitted >= PluginRowLimits.emergencyMax {
logger.warning("Streamed result truncated at \(PluginRowLimits.emergencyMax) documents")
break
if projection == nil, sample.count >= MongoStreamProjection.sampleSize {
do {
projection = openStream(sample: sample, census: try census(), continuation: continuation)
} catch {
cleanup(streamState)
continuation.finish(throwing: error)
return
}
emitted += sample.count
sample = []
}

guard let projection else {
sample.append(dict)
continue
}

sample.append(dict)
if sample.count >= MongoStreamProjection.sampleSize {
projection = openStream(sample: sample, continuation: continuation)
emitted += sample.count
sample = []
continuation.yield(.rows([projection.row(for: dict, convert: streamCellValue)]))
unannounced.formUnion(projection.unannouncedFields(in: dict))
emitted += 1
if emitted >= PluginRowLimits.emergencyMax {
logger.warning("Streamed result truncated at \(PluginRowLimits.emergencyMax) documents")
break
}
}

logUnannouncedFields(unannounced)

var error = bson_error_t()
if mongoc_cursor_error(cursor, &error) {
cleanup(streamState)
Expand All @@ -458,27 +473,48 @@ extension MongoDBConnection {
}

if projection == nil {
_ = openStream(sample: sample, continuation: continuation)
_ = openStream(sample: sample, census: nil, continuation: continuation)
}

cleanup(streamState)
continuation.finish()
}

func fieldCensus(
_ request: MongoFieldCensus.Request?,
client: OpaquePointer,
database: String,
collection: String
) throws -> MongoFieldCensus? {
guard let request else { return nil }
do {
let groups = try aggregateSync(
client: client, database: database, collection: collection,
pipeline: request.pipeline, optionsJson: request.optionsJson
)
return MongoFieldCensus(groups: groups.docs)
} catch let error as MongoDBError {
logger.warning(
"Field census failed, the stream header holds the sampled fields only: \(error.message)"
)
return nil
}
}

private func logUnannouncedFields(_ fields: Set<String>) {
guard !fields.isEmpty else { return }
let names = fields.sorted().joined(separator: ", ")
logger.warning(
"Stream header left out \(fields.count) fields first seen after the sample: \(names, privacy: .private)"
)
}

private func openStream(
sample: [[String: Any]],
census: MongoFieldCensus?,
continuation: AsyncThrowingStream<PluginStreamElement, Error>.Continuation
) -> MongoStreamProjection {
let columns = BsonDocumentFlattener.unionColumns(from: sample)
let kinds = BsonDocumentFlattener.columnKinds(
for: columns, documents: sample, representation: uuidRepresentation
)
let columnTypeNames = kinds.map {
BsonDocumentFlattener.typeName(for: $0, representation: uuidRepresentation)
}
let projection = MongoStreamProjection(
columns: columns, columnTypeNames: columnTypeNames, kinds: kinds
)
let projection = MongoStreamProjection(sample: sample, census: census, representation: uuidRepresentation)

continuation.yield(.header(projection.header))

Expand Down
18 changes: 14 additions & 4 deletions Plugins/MongoDBDriverPlugin/MongoDBConnection.swift
Original file line number Diff line number Diff line change
Expand Up @@ -795,7 +795,8 @@ final class MongoDBConnection: @unchecked Sendable {
database: String,
collection: String,
filter: String,
optionsJson: String
optionsJson: String,
census: MongoFieldCensus.Request?
) -> AsyncThrowingStream<PluginStreamElement, Error> {
#if canImport(CLibMongoc)
let queue = self.queue
Expand Down Expand Up @@ -844,7 +845,11 @@ final class MongoDBConnection: @unchecked Sendable {
streamState.collection = col
streamState.lock.unlock()

iterateCursorStreaming(cursor: cursor, continuation: continuation, streamState: streamState)
iterateCursorStreaming(
cursor: cursor, continuation: continuation, streamState: streamState
) {
try self.fieldCensus(census, client: client, database: database, collection: collection)
}
} catch {
continuation.finish(throwing: error)
}
Expand All @@ -859,7 +864,8 @@ final class MongoDBConnection: @unchecked Sendable {
database: String,
collection: String,
pipeline: String,
optionsJson: String? = nil
optionsJson: String?,
census: MongoFieldCensus.Request?
) -> AsyncThrowingStream<PluginStreamElement, Error> {
#if canImport(CLibMongoc)
let queue = self.queue
Expand Down Expand Up @@ -919,7 +925,11 @@ final class MongoDBConnection: @unchecked Sendable {
streamState.collection = col
streamState.lock.unlock()

iterateCursorStreaming(cursor: cursor, continuation: continuation, streamState: streamState)
iterateCursorStreaming(
cursor: cursor, continuation: continuation, streamState: streamState
) {
try self.fieldCensus(census, client: client, database: database, collection: collection)
}
} catch {
continuation.finish(throwing: error)
}
Expand Down
9 changes: 7 additions & 2 deletions Plugins/MongoDBDriverPlugin/MongoDBPluginDriver.swift
Original file line number Diff line number Diff line change
Expand Up @@ -835,18 +835,23 @@ final class MongoDBPluginDriver: PluginDatabaseDriver, @unchecked Sendable {
switch try await runtime.exportPlan(for: trimmed, database: db) {
case .cursor(let plan, let databaseSwitch, let writes):
if let databaseSwitch { self.currentDb = databaseSwitch }
let census = MongoFieldCensus.request(
for: plan, limit: PluginRowLimits.emergencyMax, timeoutMS: timeout
)
let inner = plan.isFind
? conn.streamFind(
database: plan.database, collection: plan.collection,
filter: plan.filter,
optionsJson: plan.options.findOptionsJson(
limit: PluginRowLimits.emergencyMax, timeoutMS: timeout
)
),
census: census
)
: conn.streamAggregate(
database: plan.database, collection: plan.collection,
pipeline: plan.pipeline,
optionsJson: plan.options.aggregateOptionsJson(timeoutMS: timeout)
optionsJson: plan.options.aggregateOptionsJson(timeoutMS: timeout),
census: census
)
do {
for try await element in inner {
Expand Down
101 changes: 101 additions & 0 deletions Plugins/MongoDBDriverPlugin/MongoFieldCensus.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
import Foundation
import TableProPluginKit

struct MongoFieldCensus {
struct Tally {
let field: String
let documents: Int64
let example: Any
}

struct Request: Sendable {
let pipeline: String
let optionsJson: String?
}

private static let tallyStages = [
"{\"$project\": {\"_id\": 0, \"pair\": {\"$objectToArray\": \"$$ROOT\"}}}",
"{\"$unwind\": \"$pair\"}",
"{\"$group\": {\"_id\": {\"$arrayToObject\": [[[\"$pair.k\", {\"$type\": \"$pair.v\"}]]]}, "
+ "\"name\": {\"$first\": \"$pair.k\"}, "
+ "\"documents\": {\"$sum\": 1}, \"example\": {\"$first\": \"$pair.v\"}}}"
]

private static let writingStages: Set<String> = ["$out", "$merge"]

let tallies: [Tally]

init(tallies: [Tally]) {
self.tallies = tallies
}

init(groups: [[String: Any]]) {
self.init(tallies: groups.compactMap(Self.tally(from:)))
}

var fields: [String] {
Set(tallies.map(\.field)).sorted()
}

func kinds(representation: MongoDBUuidRepresentation) -> [String: BsonValueKind] {
var votes: [String: [BsonValueKind: Int64]] = [:]
for tally in tallies where !(tally.example is NSNull) {
let kind = BsonDocumentFlattener.valueKind(for: tally.example, representation: representation)
votes[tally.field, default: [:]][kind, default: 0] += tally.documents
}
return votes.compactMapValues { fieldVotes in
fieldVotes.max { $0.value < $1.value }?.key
}
}

static func request(for plan: MongoScriptCursorPlan, limit ceiling: Int, timeoutMS: Int32) -> Request? {
let sourceStages = plan.isFind
? findStages(of: plan, limit: ceiling)
: aggregateStages(of: plan.pipeline, limit: ceiling)
guard let sourceStages else { return nil }
return Request(
pipeline: "[\((sourceStages + tallyStages).joined(separator: ", "))]",
optionsJson: plan.options.aggregateOptionsJson(timeoutMS: timeoutMS)
)
}

private static func findStages(of plan: MongoScriptCursorPlan, limit ceiling: Int) -> [String] {
var stages = ["{\"$match\": \(plan.filter)}"]
let isPaged = plan.options.skip != nil || plan.options.limit != nil
if isPaged, let sort = document(plan.sort) { stages.append("{\"$sort\": \(sort)}") }
if plan.skip > 0 { stages.append("{\"$skip\": \(plan.skip)}") }
stages += limitStage(plan.options.effectiveLimit(ceiling: ceiling))
if let projection = document(plan.projection) { stages.append("{\"$project\": \(projection)}") }
return stages
}

private static func aggregateStages(of pipeline: String, limit ceiling: Int) -> [String]? {
guard let data = pipeline.data(using: .utf8),
let stages = (try? JSONSerialization.jsonObject(with: data)) as? [[String: Any]],
!stages.contains(where: { !writingStages.isDisjoint(with: $0.keys) }) else {
return nil
}
let trimmed = pipeline.trimmingCharacters(in: .whitespacesAndNewlines)
let inner = String(trimmed.dropFirst().dropLast()).trimmingCharacters(in: .whitespacesAndNewlines)
return (inner.isEmpty ? [] : [inner]) + limitStage(ceiling)
}

private static func limitStage(_ limit: Int) -> [String] {
limit > 0 ? ["{\"$limit\": \(limit)}"] : []
}

private static func document(_ text: String?) -> String? {
guard let text else { return nil }
let compact = text.filter { !$0.isWhitespace }
guard !compact.isEmpty, compact != "null", compact != "{}" else { return nil }
return text
}

private static func tally(from group: [String: Any]) -> Tally? {
guard let field = group["name"] as? String,
let documents = MongoScriptJson.numeric(group["documents"]) else {
return nil
}
return Tally(field: field, documents: documents, example: group["example"] ?? NSNull())
}
}
Loading
Loading