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
5 changes: 5 additions & 0 deletions flight/flight-core/pom.xml
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,11 @@ under the License.
<artifactId>grpc-services</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.apache.arrow</groupId>
<artifactId>arrow-compression</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.grpc</groupId>
<artifactId>grpc-inprocess</artifactId>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,14 +17,37 @@
package org.apache.arrow.flight;

import io.grpc.stub.AbstractStub;
import java.util.Arrays;
import java.util.concurrent.TimeUnit;
import java.util.stream.Collectors;
import org.apache.arrow.vector.compression.CompressionCodec;
import org.apache.arrow.vector.compression.CompressionUtil;

/** Common call options. */
public class CallOptions {
public static CallOption timeout(long duration, TimeUnit unit) {
return new Timeout(duration, unit);
}

/**
* Advertise support for IPC body compression codecs in preference order.
*
* <p>Servers that do not support negotiation ignore this option and return uncompressed IPC.
*/
public static CallOption acceptIpcCompression(CompressionUtil.CodecType... codecs) {
if (codecs.length == 0) {
throw new IllegalArgumentException("At least one IPC compression codec is required");
}
for (CompressionUtil.CodecType codec : codecs) {
CompressionCodec.Factory.INSTANCE.createCodec(codec);
}
final FlightCallHeaders headers = new FlightCallHeaders();
headers.insert(
FlightConstants.IPC_ACCEPT_COMPRESSION_HEADER,
Arrays.stream(codecs).map(IpcCompression::codecName).collect(Collectors.joining(",")));
return new HeaderCallOption(headers);
}

static <T extends AbstractStub<T>> T wrapStub(T stub, CallOption[] options) {
for (CallOption option : options) {
if (option instanceof GrpcCallOption) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
import org.apache.arrow.vector.FieldVector;
import org.apache.arrow.vector.VectorSchemaRoot;
import org.apache.arrow.vector.VectorUnloader;
import org.apache.arrow.vector.compression.CompressionCodec;
import org.apache.arrow.vector.dictionary.Dictionary;
import org.apache.arrow.vector.dictionary.DictionaryProvider;
import org.apache.arrow.vector.ipc.message.ArrowDictionaryBatch;
Expand Down Expand Up @@ -57,6 +58,18 @@ static Schema generateSchemaMessages(
final IpcOption option,
final Consumer<ArrowMessage> messageCallback)
throws Exception {
return generateSchemaMessages(
originalSchema, descriptor, provider, option, null, messageCallback);
}

static Schema generateSchemaMessages(
final Schema originalSchema,
final FlightDescriptor descriptor,
final DictionaryProvider provider,
final IpcOption option,
final CompressionCodec compressionCodec,
final Consumer<ArrowMessage> messageCallback)
throws Exception {
final Set<Long> dictionaryIds = new HashSet<>();
final Schema schema = generateSchema(originalSchema, provider, dictionaryIds);
MetadataV4UnionChecker.checkForUnion(schema.getFields().iterator(), option.metadataVersion);
Expand All @@ -77,7 +90,9 @@ static Schema generateSchemaMessages(
Collections.singletonList(vector.getField()),
Collections.singletonList(vector),
count);
final VectorUnloader unloader = new VectorUnloader(dictRoot);
final VectorUnloader unloader =
new VectorUnloader(
dictRoot, /* includeNullCount */ true, compressionCodec, /* alignBuffers */ true);
try (final ArrowDictionaryBatch dictionaryBatch =
new ArrowDictionaryBatch(id, unloader.getRecordBatch());
final ArrowMessage message = new ArrowMessage(dictionaryBatch, option)) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,13 +26,15 @@
import io.grpc.protobuf.ProtoUtils;
import io.grpc.stub.ServerCalls;
import io.grpc.stub.StreamObserver;
import java.util.List;
import java.util.Set;
import java.util.concurrent.ExecutorService;
import org.apache.arrow.flight.auth.ServerAuthHandler;
import org.apache.arrow.flight.impl.Flight;
import org.apache.arrow.flight.impl.Flight.PutResult;
import org.apache.arrow.flight.impl.FlightServiceGrpc;
import org.apache.arrow.memory.BufferAllocator;
import org.apache.arrow.vector.compression.CompressionUtil;

/** Extends the basic flight service to override some methods for more efficient implementations. */
class FlightBindingService implements BindableService {
Expand All @@ -53,8 +55,18 @@ public FlightBindingService(
FlightProducer producer,
ServerAuthHandler authHandler,
ExecutorService executor) {
this(allocator, producer, authHandler, executor, java.util.Collections.emptyList());
}

public FlightBindingService(
BufferAllocator allocator,
FlightProducer producer,
ServerAuthHandler authHandler,
ExecutorService executor,
List<CompressionUtil.CodecType> ipcCompressionCodecs) {
this.allocator = allocator;
this.delegate = new FlightService(allocator, producer, authHandler, executor);
this.delegate =
new FlightService(allocator, producer, authHandler, executor, ipcCompressionCodecs);
}

public static MethodDescriptor<Flight.Ticket, ArrowMessage> getDoGetDescriptor(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,9 @@ public interface FlightConstants {

String SERVICE = "arrow.flight.protocol.FlightService";

/** Request header listing supported IPC body compression codecs in preference order. */
String IPC_ACCEPT_COMPRESSION_HEADER = "arrow-ipc-accept-compression";

FlightServerMiddleware.Key<ServerHeaderMiddleware> HEADER_KEY =
FlightServerMiddleware.Key.of("org.apache.arrow.flight.ServerHeaderMiddleware");

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@
import java.net.URI;
import java.net.URISyntaxException;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
Expand All @@ -55,6 +57,8 @@
import org.apache.arrow.memory.BufferAllocator;
import org.apache.arrow.util.Preconditions;
import org.apache.arrow.util.VisibleForTesting;
import org.apache.arrow.vector.compression.CompressionCodec;
import org.apache.arrow.vector.compression.CompressionUtil;

/**
* Generic server of flight data that is customized via construction with delegate classes for the
Expand Down Expand Up @@ -197,6 +201,7 @@ public static final class Builder {
private final List<KeyFactory<?>> interceptors;
// Keep track of inserted interceptors
private final Set<String> interceptorKeys;
private List<CompressionUtil.CodecType> ipcCompressionCodecs = Collections.emptyList();

Builder() {
builderOptions = new HashMap<>();
Expand Down Expand Up @@ -321,7 +326,7 @@ public FlightServer build() {
}

final FlightBindingService flightService =
new FlightBindingService(allocator, producer, authHandler, exec);
new FlightBindingService(allocator, producer, authHandler, exec, ipcCompressionCodecs);
builder
.executor(exec)
.maxInboundMessageSize(maxInboundMessageSize)
Expand Down Expand Up @@ -392,6 +397,22 @@ public Builder backpressureThreshold(int backpressureThreshold) {
return this;
}

/**
* Enable negotiated IPC body compression for server response streams.
*
* <p>The client preference order wins. Clients that do not advertise support continue to
* receive uncompressed IPC.
*/
public Builder ipcCompression(CompressionUtil.CodecType... codecs) {
Preconditions.checkArgument(codecs.length > 0, "At least one codec is required");
for (CompressionUtil.CodecType codec : codecs) {
IpcCompression.codecName(codec);
CompressionCodec.Factory.INSTANCE.createCodec(codec);
}
ipcCompressionCodecs = Collections.unmodifiableList(Arrays.asList(codecs.clone()));
return this;
}

/**
* A small utility function to ensure that InputStream attributes. are closed if they are not
* null
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
import io.grpc.stub.ServerCallStreamObserver;
import io.grpc.stub.StreamObserver;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Future;
Expand All @@ -38,6 +39,8 @@
import org.apache.arrow.flight.impl.FlightServiceGrpc.FlightServiceImplBase;
import org.apache.arrow.memory.BufferAllocator;
import org.apache.arrow.util.AutoCloseables;
import org.apache.arrow.vector.compression.CompressionCodec;
import org.apache.arrow.vector.compression.CompressionUtil;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

Expand All @@ -51,16 +54,27 @@ class FlightService extends FlightServiceImplBase {
private final FlightProducer producer;
private final ServerAuthHandler authHandler;
private final ExecutorService executors;
private final List<CompressionUtil.CodecType> ipcCompressionCodecs;

FlightService(
BufferAllocator allocator,
FlightProducer producer,
ServerAuthHandler authHandler,
ExecutorService executors) {
this(allocator, producer, authHandler, executors, Collections.emptyList());
}

FlightService(
BufferAllocator allocator,
FlightProducer producer,
ServerAuthHandler authHandler,
ExecutorService executors,
List<CompressionUtil.CodecType> ipcCompressionCodecs) {
this.allocator = allocator;
this.producer = producer;
this.authHandler = authHandler;
this.executors = new ContextPropagatingExecutorService(executors);
this.ipcCompressionCodecs = ipcCompressionCodecs;
}

private CallContext makeContext(ServerCallStreamObserver<?> responseObserver) {
Expand Down Expand Up @@ -107,10 +121,14 @@ public void doGetCustom(
final ServerCallStreamObserver<ArrowMessage> responseObserver =
(ServerCallStreamObserver<ArrowMessage>) responseObserverSimple;

final CallContext context = makeContext(responseObserver);
final GetListener listener =
new GetListener(responseObserver, this::handleExceptionWithMiddleware);
new GetListener(
responseObserver,
this::handleExceptionWithMiddleware,
IpcCompression.negotiate(context, ipcCompressionCodecs));
try {
producer.getStream(makeContext(responseObserver), new Ticket(ticket), listener);
producer.getStream(context, new Ticket(ticket), listener);
} catch (Exception ex) {
listener.error(ex);
}
Expand Down Expand Up @@ -155,7 +173,14 @@ private static class GetListener extends OutboundStreamListenerImpl

public GetListener(
ServerCallStreamObserver<ArrowMessage> responseObserver, Consumer<Throwable> errorHandler) {
super(null, responseObserver);
this(responseObserver, errorHandler, null);
}

public GetListener(
ServerCallStreamObserver<ArrowMessage> responseObserver,
Consumer<Throwable> errorHandler,
CompressionCodec compressionCodec) {
super(null, responseObserver, compressionCodec);
this.errorHandler = errorHandler;
this.completed = false;
this.serverCallResponseObserver = responseObserver;
Expand Down Expand Up @@ -327,7 +352,14 @@ private static class ExchangeListener extends GetListener {

public ExchangeListener(
ServerCallStreamObserver<ArrowMessage> responseObserver, Consumer<Throwable> errorHandler) {
super(responseObserver, errorHandler);
this(responseObserver, errorHandler, null);
}

public ExchangeListener(
ServerCallStreamObserver<ArrowMessage> responseObserver,
Consumer<Throwable> errorHandler,
CompressionCodec compressionCodec) {
super(responseObserver, errorHandler, compressionCodec);
this.resource = null;
super.setOnCancelHandler(
() -> {
Expand Down Expand Up @@ -387,8 +419,12 @@ public StreamObserver<ArrowMessage> doExchangeCustom(
StreamObserver<ArrowMessage> responseObserverSimple) {
final ServerCallStreamObserver<ArrowMessage> responseObserver =
(ServerCallStreamObserver<ArrowMessage>) responseObserverSimple;
final CallContext context = makeContext(responseObserver);
final ExchangeListener listener =
new ExchangeListener(responseObserver, this::handleExceptionWithMiddleware);
new ExchangeListener(
responseObserver,
this::handleExceptionWithMiddleware,
IpcCompression.negotiate(context, ipcCompressionCodecs));
final FlightStream fs =
new FlightStream(
allocator,
Expand All @@ -405,7 +441,7 @@ public StreamObserver<ArrowMessage> doExchangeCustom(
executors.submit(
() -> {
try {
producer.doExchange(makeContext(responseObserver), fs, listener);
producer.doExchange(context, fs, listener);
} catch (Exception ex) {
listener.error(ex);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
import org.apache.arrow.vector.FieldVector;
import org.apache.arrow.vector.VectorLoader;
import org.apache.arrow.vector.VectorSchemaRoot;
import org.apache.arrow.vector.compression.CompressionUtil;
import org.apache.arrow.vector.dictionary.Dictionary;
import org.apache.arrow.vector.dictionary.DictionaryProvider;
import org.apache.arrow.vector.ipc.message.ArrowDictionaryBatch;
Expand Down Expand Up @@ -86,6 +87,7 @@ public void close() throws Exception {}
private volatile Throwable ex;
private volatile ArrowBuf applicationMetadata = null;
@VisibleForTesting volatile MetadataVersion metadataVersion = null;
@VisibleForTesting volatile CompressionUtil.CodecType compressionType = null;

/**
* Constructs a new instance.
Expand Down Expand Up @@ -272,6 +274,9 @@ public boolean next() {
// Ensure we have the root
root.get().clear();
try (ArrowRecordBatch arb = msg.asRecordBatch()) {
compressionType =
CompressionUtil.CodecType.fromCompressionType(
arb.getBodyCompression().getCodec());
loader.load(arb);
}
updateMetadata(msg);
Expand Down
Loading
Loading