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
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,8 @@ import com.typesafe.config.Config
final class BigQueryStorageSettings private (
val host: String,
val port: Int,
val rootCa: Option[String] = None) {
val rootCa: Option[String] = None,
val arrowAllocatorBytes: Long = BigQueryStorageSettings.DefaultArrowAllocatorBytes) {

/**
* Endpoint hostname where the gRPC connection is made.
Expand All @@ -38,19 +39,34 @@ final class BigQueryStorageSettings private (
def withRootCa(rootCa: String): BigQueryStorageSettings =
copy(rootCa = Some(rootCa))

private def copy(host: String = host, port: Int = port, rootCa: Option[String] = rootCa) =
new BigQueryStorageSettings(host, port, rootCa)
/**
* Maximum bytes the Arrow root allocator may reserve per batch.
* The allocator is created per batch and closed after reading, so this bounds
* native memory for a single batch rather than the lifetime of the stream.
*/
def withArrowAllocatorBytes(bytes: Long): BigQueryStorageSettings =
copy(arrowAllocatorBytes = bytes)

private def copy(
host: String = host,
port: Int = port,
rootCa: Option[String] = rootCa,
arrowAllocatorBytes: Long = arrowAllocatorBytes) =
new BigQueryStorageSettings(host, port, rootCa, arrowAllocatorBytes)

override def toString: String =
"BigQueryStorageSettings(" +
s"host=$host, " +
s"port=$port, " +
s"rootCa=$rootCa" +
s"rootCa=$rootCa, " +
s"arrowAllocatorBytes=$arrowAllocatorBytes" +
")"
}

object BigQueryStorageSettings {

val DefaultArrowAllocatorBytes: Long = 512L * 1024 * 1024 // 512 MB

/**
* Create settings for unsecure (no tls), unauthenticated (no root ca)
* and unauthorized (no call credentials) endpoint.
Expand All @@ -73,7 +89,12 @@ object BigQueryStorageSettings {
case _ => bigQueryConfig
}

Seq(setRootCa).foldLeft(bigQueryConfig) {
val setAllocatorBytes = (bigQueryConfig: BigQueryStorageSettings) =>
if (config.hasPath("arrowAllocatorBytes"))
bigQueryConfig.withArrowAllocatorBytes(config.getBytes("arrowAllocatorBytes"))
else bigQueryConfig

Seq(setRootCa, setAllocatorBytes).foldLeft(bigQueryConfig) {
case (c, f) => f(c)
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,19 +32,28 @@ import scala.jdk.CollectionConverters._

object ArrowSource {

def readRecordsMerged(client: BigQueryReadClient, readSession: ReadSession): Source[List[BigQueryRecord], NotUsed] =
def readRecordsMerged(client: BigQueryReadClient,
readSession: ReadSession,
allocatorBytes: Long): Source[List[BigQueryRecord], NotUsed] =
readMerged(client, readSession)
.map(a => new SimpleRowReader(readSession.schema.arrowSchema.get).read(a))
.map { a =>
val reader = new SimpleRowReader(readSession.schema.arrowSchema.get, allocatorBytes)
try reader.read(a)
finally reader.close()
}

def readMerged(client: BigQueryReadClient, session: ReadSession): Source[ArrowRecordBatch, NotUsed] =
read(client, session)
.reduce((a, b) => a.merge(b))
read(client, session).reduce((a, b) => a.merge(b))

def readRecords(client: BigQueryReadClient, session: ReadSession): Seq[Source[BigQueryRecord, NotUsed]] =
def readRecords(client: BigQueryReadClient, session: ReadSession,
allocatorBytes: Long): Seq[Source[BigQueryRecord, NotUsed]] =
read(client, session)
.map { a =>
a.map(new SimpleRowReader(session.schema.arrowSchema.get).read(_))
.mapConcat(c => c)
a.map { batch =>
val reader = new SimpleRowReader(session.schema.arrowSchema.get, allocatorBytes)
try reader.read(batch)
finally reader.close()
}.mapConcat(c => c)
}

def read(client: BigQueryReadClient, session: ReadSession): Seq[Source[ArrowRecordBatch, NotUsed]] =
Expand All @@ -56,9 +65,9 @@ object ArrowSource {

}

final class SimpleRowReader(val schema: ArrowSchema) extends AutoCloseable {
final class SimpleRowReader(val schema: ArrowSchema, allocatorBytes: Long) extends AutoCloseable {

val allocator = new RootAllocator(Long.MaxValue)
val allocator = new RootAllocator(allocatorBytes)

val sd = MessageSerializer.deserializeSchema(
new ReadChannel(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ object BigQueryArrowStorage {
tableId,
readOptions,
maxNumStreams,
(_, client, session) => ArrowSource.readRecordsMerged(client, session))
(_, client, session, allocatorBytes) => ArrowSource.readRecordsMerged(client, session, allocatorBytes))
.flatMapConcat(a => a)

def readRecords(projectId: String,
Expand All @@ -55,7 +55,7 @@ object BigQueryArrowStorage {
tableId,
readOptions,
maxNumStreams,
(_, client, session) => ArrowSource.readRecords(client, session))
(_, client, session, allocatorBytes) => ArrowSource.readRecords(client, session, allocatorBytes))

def readMerged(projectId: String,
datasetId: String,
Expand All @@ -67,7 +67,7 @@ object BigQueryArrowStorage {
tableId,
readOptions,
maxNumStreams,
(schema, client, session) => (schema, ArrowSource.readMerged(client, session)))
(schema, client, session, _) => (schema, ArrowSource.readMerged(client, session)))

def read(projectId: String,
datasetId: String,
Expand All @@ -79,20 +79,22 @@ object BigQueryArrowStorage {
tableId,
readOptions,
maxNumStreams,
(schema, client, session) => (schema, ArrowSource.read(client, session)))
(schema, client, session, _) => (schema, ArrowSource.read(client, session)))

private def readAndMapTo[T](projectId: String,
datasetId: String,
tableId: String,
readOptions: Option[TableReadOptions],
maxNumStreams: Int,
fx: (ArrowSchema, BigQueryReadClient, ReadSession) => T): Source[T, Future[NotUsed]] =
fx: (ArrowSchema, BigQueryReadClient, ReadSession, Long) => T): Source[T, Future[NotUsed]] =
Source.fromMaterializer { (mat, attr) =>
val client = reader(mat.system, attr).client
val rdr = reader(mat.system, attr)
val client = rdr.client
val allocatorBytes = rdr.settings.arrowAllocatorBytes
readSession(client, projectId, datasetId, tableId, DataFormat.ARROW, readOptions, maxNumStreams)
.map { session =>
session.schema match {
case ReadSession.Schema.ArrowSchema(schema) => fx(schema, client, session)
case ReadSession.Schema.ArrowSchema(schema) => fx(schema, client, session, allocatorBytes)
case other => throw new IllegalArgumentException(s"Only Arrow format is supported, received: $other")
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,8 @@ import com.google.cloud.bigquery.storage.v1.storage.BigQueryReadClient
/**
* Holds the gRPC scala reader client instance.
*/
final class GrpcBigQueryStorageReader private (settings: BigQueryStorageSettings, sys: ClassicActorSystemProvider) {
final class GrpcBigQueryStorageReader private[scaladsl] (val settings: BigQueryStorageSettings,
sys: ClassicActorSystemProvider) {

@ApiMayChange
final val client = BigQueryReadClient(PekkoGrpcSettings.fromBigQuerySettings(settings)(sys))(sys)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,8 @@ class BigQueryArrowStorageSpec

"BigQueryArrowStorage.readArrow" should {

val reader = new SimpleRowReader(ArrowSchema(serializedSchema = GCPSerializedArrowSchema))
val reader = new SimpleRowReader(ArrowSchema(serializedSchema = GCPSerializedArrowSchema),
BigQueryStorageSettings.DefaultArrowAllocatorBytes)
val expectedRecords = reader.read(ArrowRecordBatch(GCPSerializedArrowTenRecordBatch, 10))

"stream the results for a query in records merged" in {
Expand Down Expand Up @@ -83,7 +84,7 @@ class BigQueryArrowStorageSpec
.futureValue
.head

val rowReader = new SimpleRowReader(schema)
val rowReader = new SimpleRowReader(schema, BigQueryStorageSettings.DefaultArrowAllocatorBytes)
val records = rowReader.read(recordBatch)

records shouldBe expectedRecords
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,8 @@ class BigQueryStorageSpec
}

"stream the results for a query using arrow deserializer" in {
val reader = new SimpleRowReader(ArrowSchema(serializedSchema = GCPSerializedArrowSchema))
val reader = new SimpleRowReader(ArrowSchema(serializedSchema = GCPSerializedArrowSchema),
BigQueryStorageSettings.DefaultArrowAllocatorBytes)
val expectedRecords = reader.read(ArrowRecordBatch(GCPSerializedArrowTenRecordBatch, 10))

implicit val um: ArrowByteStringDecoder = new ArrowByteStringDecoder(ArrowSchema(GCPSerializedArrowSchema))
Expand Down