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
48 changes: 37 additions & 11 deletions Sources/Arrow/ArrowReader.swift
Original file line number Diff line number Diff line change
Expand Up @@ -274,19 +274,27 @@ public class ArrowReader { // swiftlint:disable:this type_body_length
}

offset += Int(MemoryLayout<UInt32>.size)
streamData = input[offset...]
streamData = input[(input.startIndex + offset)...]
var dataBuffer = ByteBuffer(
data: streamData,
allowReadingUnalignedBuffers: useUnalignedBuffers
)
let message: org_apache_arrow_flatbuf_Message = getRoot(byteBuffer: &dataBuffer)
switch message.headerType {
case .recordbatch:
let rbMessage = message.header(type: org_apache_arrow_flatbuf_RecordBatch.self)!
guard let rbMessage = message.header(type: org_apache_arrow_flatbuf_RecordBatch.self) else {
return .failure(.invalid("RecordBatch header not found"))
}
guard let schemaMsg = schemaMessage else {
return .failure(.invalid("Schema must be defined before RecordBatch"))
}
guard let schema = result.schema else {
return .failure(.invalid("Schema not loaded"))
}
let recordBatchResult = loadRecordBatch(
rbMessage,
schema: schemaMessage!,
arrowSchema: result.schema!,
schema: schemaMsg,
arrowSchema: schema,
data: input,
messageEndOffset: (Int64(offset) + Int64(length)))
switch recordBatchResult {
Expand Down Expand Up @@ -335,7 +343,10 @@ public class ArrowReader { // swiftlint:disable:this type_body_length
data: footerData,
allowReadingUnalignedBuffers: useUnalignedBuffers)
let footer: org_apache_arrow_flatbuf_Footer = getRoot(byteBuffer: &footerBuffer)
let schemaResult = loadSchema(footer.schema!)
guard let footerSchema = footer.schema else {
return .failure(.invalid("Footer schema not found"))
}
let schemaResult = loadSchema(footerSchema)
switch schemaResult {
case .success(let schema):
result.schema = schema
Expand Down Expand Up @@ -368,11 +379,16 @@ public class ArrowReader { // swiftlint:disable:this type_body_length
let message: org_apache_arrow_flatbuf_Message = getRoot(byteBuffer: &mbb)
switch message.headerType {
case .recordbatch:
let rbMessage = message.header(type: org_apache_arrow_flatbuf_RecordBatch.self)!
guard let rbMessage = message.header(type: org_apache_arrow_flatbuf_RecordBatch.self) else {
return .failure(.invalid("RecordBatch header not found"))
}
guard let schema = result.schema else {
return .failure(.invalid("Schema not loaded"))
}
let recordBatchResult = loadRecordBatch(
rbMessage,
schema: footer.schema!,
arrowSchema: result.schema!,
schema: footerSchema,
arrowSchema: schema,
data: fileData,
messageEndOffset: messageEndOffset)
switch recordBatchResult {
Expand Down Expand Up @@ -421,7 +437,9 @@ public class ArrowReader { // swiftlint:disable:this type_body_length
let message: org_apache_arrow_flatbuf_Message = getRoot(byteBuffer: &mbb)
switch message.headerType {
case .schema:
let sMessage = message.header(type: org_apache_arrow_flatbuf_Schema.self)!
guard let sMessage = message.header(type: org_apache_arrow_flatbuf_Schema.self) else {
return .failure(.invalid("Schema header not found"))
}
switch loadSchema(sMessage) {
case .success(let schema):
result.schema = schema
Expand All @@ -431,9 +449,17 @@ public class ArrowReader { // swiftlint:disable:this type_body_length
return .failure(error)
}
case .recordbatch:
let rbMessage = message.header(type: org_apache_arrow_flatbuf_RecordBatch.self)!
guard let rbMessage = message.header(type: org_apache_arrow_flatbuf_RecordBatch.self) else {
return .failure(.invalid("RecordBatch header not found"))
}
guard let messageSchema = result.messageSchema else {
return .failure(.invalid("Schema must be defined before RecordBatch"))
}
guard let schema = result.schema else {
return .failure(.invalid("Schema not loaded"))
}
let recordBatchResult = loadRecordBatch(
rbMessage, schema: result.messageSchema!, arrowSchema: result.schema!,
rbMessage, schema: messageSchema, arrowSchema: schema,
data: dataBody, messageEndOffset: 0)
switch recordBatchResult {
case .success(let recordBatch):
Expand Down
45 changes: 45 additions & 0 deletions Tests/ArrowTests/IPCTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -262,6 +262,50 @@ final class IPCStreamReaderTests: XCTestCase {
throw error
}
}

func testReadStreamingRecordBatchBeforeSchema() throws {
// Build a minimal streaming message: a RecordBatch header with no
// preceding Schema message. This should fail gracefully instead
// of crashing on a force unwrap.
let schema = makeSchema()
let recordBatch = try makeRecordBatch()
let arrowWriter = ArrowWriter()
let writerInfo = ArrowWriter.Info(.recordbatch, schema: schema, batches: [recordBatch])

switch arrowWriter.writeStreaming(writerInfo) {
case .success(let writeData):
// Mirror the parsing logic in ArrowReader.readStreaming to advance
// past exactly one message (the Schema message), leaving the
// RecordBatch message intact and correctly positioned for
// readStreaming to parse on its own.
var offset = 0
var length = getUInt32(writeData, offset: offset)
if length == CONTINUATIONMARKER {
offset += Int(MemoryLayout<UInt32>.size)
length = getUInt32(writeData, offset: offset)
}
offset += Int(MemoryLayout<UInt32>.size)

var dataBuffer = ByteBuffer(
data: writeData[offset...],
allowReadingUnalignedBuffers: false)
let message: org_apache_arrow_flatbuf_Message = getRoot(byteBuffer: &dataBuffer)
XCTAssertEqual(message.headerType, .schema)

offset += Int(message.bodyLength + Int64(length))
let truncatedData = writeData[offset...]

let arrowReader = ArrowReader()
switch arrowReader.readStreaming(truncatedData) {
case .success:
XCTFail("Expected failure when RecordBatch precedes Schema")
case .failure:
break // Correct: should fail gracefully, not crash
}
case .failure(let error):
throw error
}
}
}

final class IPCFileReaderTests: XCTestCase { // swiftlint:disable:this type_body_length
Expand Down Expand Up @@ -671,5 +715,6 @@ final class IPCFileReaderTests: XCTestCase { // swiftlint:disable:this type_body
throw error
}
}

}
// swiftlint:disable:this file_length