GH-750 Add support for pluggable protobufs

This initial support adds plugin extension to support CloudEvent proto as well as the example
Additional plugins could be provided in the same ay as CloudEvent plugin extension

Resolves #750
This commit is contained in:
Oleg Zhurakousky
2021-10-11 14:03:24 +02:00
parent 346ff53539
commit 7fc755e157
41 changed files with 2659 additions and 312 deletions

View File

@@ -17,7 +17,7 @@
<properties>
<grpc.version>1.16.1</grpc.version>
<checkstyle.skip>true</checkstyle.skip>
<disable.checks>true</disable.checks>
</properties>
<dependencies>
@@ -65,12 +65,6 @@
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-checkstyle-plugin</artifactId>
<!-- <executions> -->
<!-- <execution> -->
<!-- <id>checkstyle-validation</id> -->
<!-- <phase>none</phase> -->
<!-- </execution> -->
<!-- </executions> -->
</plugin>
<plugin>
<groupId>org.xolstice.maven.plugins</groupId>

View File

@@ -0,0 +1,61 @@
/*
* Copyright 2021-2021 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.cloud.function.grpc;
import org.springframework.messaging.Message;
import com.google.protobuf.GeneratedMessageV3;
/**
*
* @author Oleg Zhurakousky
*
* @param <T> instance of {@link GeneratedMessageV3}
*/
public abstract class AbstractGrpcMessageConverter<T extends GeneratedMessageV3> implements GrpcMessageConverter<T> {
@Override
public Message<byte[]> toSpringMessage(T grpcMessage) {
if (this.supports(grpcMessage)) {
return this.doToSpringMessage(grpcMessage);
}
return null;
}
@Override
public T fromSpringMessage(Message<byte[]> springMessage, Class<T> grpcClass) {
if (this.supports(grpcClass)) {
return this.doFromSpringMessage(springMessage);
}
return null;
}
protected abstract Message<byte[]> doToSpringMessage(T grpcMessage);
protected abstract T doFromSpringMessage(Message<byte[]> springMessage);
protected boolean supports(T grpcMessage) {
// String fieldName = grpcMessage.getAllFields().keySet().iterator().next().getFullName();
// fieldName = fieldName.substring(0, fieldName.lastIndexOf("."));
// System.out.println(grpcMessage.getClass().getName());
// return fieldName.contains(grpcMessage.getClass().getSimpleName());
return this.supports(grpcMessage.getClass());
}
protected abstract boolean supports(Class<? extends GeneratedMessageV3> grpcClass);
}

View File

@@ -16,13 +16,19 @@
package org.springframework.cloud.function.grpc;
import java.util.List;
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.cloud.function.context.FunctionCatalog;
import org.springframework.cloud.function.context.FunctionProperties;
import org.springframework.cloud.function.grpc.MessagingServiceGrpc.MessagingServiceImplBase;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.util.Assert;
import com.google.protobuf.GeneratedMessageV3;
import io.grpc.BindableService;
/**
*
@@ -35,13 +41,25 @@ import org.springframework.context.annotation.Configuration;
class GrpcAutoConfiguration {
@Bean
public GrpcServer grpcServer(FunctionGrpcProperties grpcProperties, MessagingServiceImplBase grpcMessagingService) {
return new GrpcServer(grpcProperties, grpcMessagingService);
public GrpcServer grpcServer(FunctionGrpcProperties grpcProperties, BindableService[] grpcMessagingServices) {
Assert.notEmpty(grpcMessagingServices, "'grpcMessagingServices' must not be null or empty");
return new GrpcServer(grpcProperties, grpcMessagingServices);
}
@Bean
public GrpcServerMessageHandler grpcMessageService(FunctionProperties funcProperties, FunctionCatalog functionCatalog) {
return new GrpcServerMessageHandler(funcProperties, functionCatalog);
public BindableService grpcSpringMessageHandler(MessageHandlingHelper helper) {
return new GrpcServerMessageHandler(helper);
}
@Bean
public MessageHandlingHelper grpcMessageHandlingHelper(List<GrpcMessageConverter<?>> grpcConverters,
FunctionProperties funcProperties, FunctionCatalog functionCatalog) {
return new MessageHandlingHelper(grpcConverters, functionCatalog, funcProperties);
}
@Bean
public GrpcSpringMessageConverter grpcSpringMessageConverter() {
return new GrpcSpringMessageConverter();
}
}

View File

@@ -0,0 +1,34 @@
/*
* Copyright 2021-2021 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.cloud.function.grpc;
import com.google.protobuf.GeneratedMessageV3;
import org.springframework.messaging.Message;
/**
*
* @author Oleg Zhurakousky
*
* @param <T> instance of {@link GeneratedMessageV3}
*/
public interface GrpcMessageConverter<T extends GeneratedMessageV3> {
Message<byte[]> toSpringMessage(T grpcMessage);
T fromSpringMessage(Message<byte[]> springMessage, Class<T> grpcClass);
}

View File

@@ -19,38 +19,46 @@ package org.springframework.cloud.function.grpc;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import io.grpc.BindableService;
import io.grpc.Server;
import io.grpc.ServerBuilder;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.cloud.function.grpc.MessagingServiceGrpc.MessagingServiceImplBase;
import org.springframework.context.SmartLifecycle;
/**
*
* @author Oleg Zhurakousky
*
*/
class GrpcServer implements SmartLifecycle {
private Log logger = LogFactory.getLog(GrpcServer.class);
private final FunctionGrpcProperties grpcProperties;
private final MessagingServiceImplBase grpcMessageService;
private final BindableService[] grpcMessageServices;
private final ExecutorService executor = Executors.newSingleThreadExecutor();
private Server server;
GrpcServer(FunctionGrpcProperties grpcProperties, MessagingServiceImplBase grpcMessageService) {
GrpcServer(FunctionGrpcProperties grpcProperties, BindableService[] grpcMessageServices) {
this.grpcProperties = grpcProperties;
this.grpcMessageService = grpcMessageService;
this.grpcMessageServices = grpcMessageServices;
}
@Override
public void start() {
this.executor.execute(() -> {
try {
this.server = ServerBuilder.forPort(this.grpcProperties.getPort())
.addService(this.grpcMessageService)
.build();
ServerBuilder<?> serverBuilder = ServerBuilder.forPort(this.grpcProperties.getPort());
for (int i = 0; i < this.grpcMessageServices.length; i++) {
BindableService bindableService = this.grpcMessageServices[i];
serverBuilder.addService(bindableService);
}
this.server = serverBuilder.build();
logger.info("Starting gRPC server");
this.server.start();

View File

@@ -32,14 +32,13 @@
package org.springframework.cloud.function.grpc;
import java.util.Map;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.LinkedBlockingQueue;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicBoolean;
//
import io.grpc.Status;
import io.grpc.stub.ServerCallStreamObserver;
import io.grpc.stub.StreamObserver;
@@ -49,7 +48,7 @@ import org.reactivestreams.Publisher;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Sinks;
import reactor.core.publisher.Sinks.Many;
//
import org.springframework.cloud.function.context.FunctionCatalog;
import org.springframework.cloud.function.context.FunctionProperties;
import org.springframework.cloud.function.context.catalog.SimpleFunctionRegistry.FunctionInvocationWrapper;
@@ -59,291 +58,53 @@ import org.springframework.messaging.Message;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import com.google.protobuf.GeneratedMessageV3;
//
//import com.google.protobuf.GeneratedMessage;
/**
*
* @author Oleg Zhurakousky
* @since 3.2
*
*/
class GrpcServerMessageHandler extends MessagingServiceImplBase implements SmartLifecycle {
@SuppressWarnings("rawtypes")
class GrpcServerMessageHandler extends MessagingServiceImplBase {
private Log logger = LogFactory.getLog(GrpcServerMessageHandler.class);
private final ExecutorService executor;
private final FunctionProperties funcProperties;
private final FunctionCatalog functionCatalog;
private final MessageHandlingHelper helper;
private boolean running;
GrpcServerMessageHandler(FunctionProperties funcProperties, FunctionCatalog functionCatalog) {
this.functionCatalog = functionCatalog;
this.funcProperties = funcProperties;
this.executor = Executors.newCachedThreadPool();
GrpcServerMessageHandler(MessageHandlingHelper<GeneratedMessageV3> helper) {
this.helper = helper;
}
@Override
@SuppressWarnings("unchecked")
public void requestReply(GrpcMessage request, StreamObserver<GrpcMessage> responseObserver) {
Message<byte[]> message = GrpcUtils.fromGrpcMessage(request);
FunctionInvocationWrapper function = this.resolveFunction(message.getHeaders());
public void requestReply(GrpcSpringMessage request, StreamObserver<GrpcSpringMessage> responseObserver) {
this.helper.requestReply(request, responseObserver);
}
Message<byte[]> replyMessage = (Message<byte[]>) function.apply(message);
@Override
@SuppressWarnings("unchecked")
public void serverStream(GrpcSpringMessage request, StreamObserver<GrpcSpringMessage> responseObserver) {
this.helper.serverStream(request, responseObserver);
}
GrpcMessage reply = GrpcUtils.toGrpcMessage(replyMessage);
responseObserver.onNext(reply);
responseObserver.onCompleted();
@Override
@SuppressWarnings("unchecked")
public StreamObserver<GrpcSpringMessage> clientStream(StreamObserver<GrpcSpringMessage> responseObserver) {
return this.helper.clientStream(responseObserver, GrpcSpringMessage.class);
}
@SuppressWarnings("unchecked")
@Override
public void serverStream(GrpcMessage request, StreamObserver<GrpcMessage> responseObserver) {
Message<byte[]> message = GrpcUtils.fromGrpcMessage(request);
FunctionInvocationWrapper function = this.resolveFunction(message.getHeaders());
Publisher<Message<byte[]>> replyStream = (Publisher<Message<byte[]>>) function.apply(message);
Flux.from(replyStream).doOnNext(replyMessage -> {
responseObserver.onNext(GrpcUtils.toGrpcMessage(replyMessage));
})
.doOnComplete(() -> responseObserver.onCompleted())
.subscribe();
}
@SuppressWarnings("unchecked")
@Override
public StreamObserver<GrpcMessage> clientStream(StreamObserver<GrpcMessage> responseObserver) {
ServerCallStreamObserver<GrpcMessage> serverCallStreamObserver = (ServerCallStreamObserver<GrpcMessage>) responseObserver;
serverCallStreamObserver.disableAutoInboundFlowControl();
FunctionInvocationWrapper function = this.resolveFunction(null);
AtomicBoolean wasReady = new AtomicBoolean(false);
serverCallStreamObserver.setOnReadyHandler(() -> {
if (serverCallStreamObserver.isReady() && !wasReady.get()) {
wasReady.set(true);
logger.info("gRPC Server receiving stream is ready.");
serverCallStreamObserver.request(1);
}
});
if (!function.isInputTypePublisher()) {
throw new UnsupportedOperationException("The client streaming is "
+ "not supported for functions that accept non-Publisher: "
+ function);
}
else if (function.isOutputTypePublisher()) {
throw new UnsupportedOperationException("The client streaming is "
+ "not supported for functions that return Publisher: "
+ function);
}
else {
Many<Message<byte[]>> inputStream = Sinks.many().unicast().onBackpressureBuffer();
Flux<Message<byte[]>> inputStreamFlux = inputStream.asFlux();
LinkedBlockingQueue<Message<byte[]>> resultRef = new LinkedBlockingQueue<>(1);
this.executor.execute(() -> {
Message<byte[]> replyMessage = (Message<byte[]>) function.apply(inputStreamFlux);
if (logger.isDebugEnabled()) {
logger.debug("Function invocation reply: " + replyMessage);
}
resultRef.offer(replyMessage);
});
return new StreamObserver<GrpcMessage>() {
@Override
public void onNext(GrpcMessage inputMessage) {
if (logger.isDebugEnabled()) {
logger.debug("gRPC Server receiving: " + inputMessage);
}
inputStream.tryEmitNext(GrpcUtils.fromGrpcMessage(inputMessage));
serverCallStreamObserver.request(1);
}
@Override
public void onError(Throwable t) {
t.printStackTrace();
responseObserver.onCompleted();
}
@Override
public void onCompleted() {
logger.info("gRPC Server has finished receiving data.");
inputStream.tryEmitComplete();
try {
responseObserver.onNext(GrpcUtils.toGrpcMessage(resultRef.poll(Integer.MAX_VALUE, TimeUnit.MILLISECONDS)));
}
catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
responseObserver.onCompleted();
}
};
}
}
@Override
public StreamObserver<GrpcMessage> biStream(StreamObserver<GrpcMessage> responseObserver) {
ServerCallStreamObserver<GrpcMessage> serverCallStreamObserver = (ServerCallStreamObserver<GrpcMessage>) responseObserver;
serverCallStreamObserver.disableAutoInboundFlowControl();
FunctionInvocationWrapper function = this.resolveFunction(null);
AtomicBoolean wasReady = new AtomicBoolean(false);
serverCallStreamObserver.setOnReadyHandler(() -> {
if (serverCallStreamObserver.isReady() && !wasReady.get()) {
wasReady.set(true);
logger.info("gRPC Server receiving stream is ready.");
serverCallStreamObserver.request(1);
}
});
if (function.isInputTypePublisher()) {
if (function.isOutputTypePublisher()) {
return this.biStreamReactive(responseObserver, serverCallStreamObserver);
}
throw new UnsupportedOperationException("The bi-directional streaming is "
+ "not supported for functions that accept Publisher but return non-Publisher: "
+ function);
}
else {
if (!function.isOutputTypePublisher()) {
return this.biStreamImperative(responseObserver, serverCallStreamObserver, wasReady);
}
throw new UnsupportedOperationException("The bidirection streaming is "
+ "not supported for functions that accept non-Publisher but return Publisher: "
+ function);
}
}
private StreamObserver<GrpcMessage> biStreamImperative(StreamObserver<GrpcMessage> responseObserver,
ServerCallStreamObserver<GrpcMessage> serverCallStreamObserver, AtomicBoolean wasReady) {
return new StreamObserver<GrpcMessage>() {
@SuppressWarnings("unchecked")
@Override
public void onNext(GrpcMessage request) {
try {
Message<byte[]> message = GrpcUtils.fromGrpcMessage(request);
FunctionInvocationWrapper function = resolveFunction(message.getHeaders());
Message<byte[]> replyMessage = (Message<byte[]>) function.apply(message);
GrpcMessage reply = GrpcUtils.toGrpcMessage(replyMessage);
responseObserver.onNext(reply);
// Check the provided ServerCallStreamObserver to see if it is still
// ready to accept more messages.
if (serverCallStreamObserver.isReady()) {
serverCallStreamObserver.request(1);
}
else {
wasReady.set(false);
}
}
catch (Throwable throwable) {
throwable.printStackTrace();
responseObserver.onError(
Status.UNKNOWN.withDescription("Error handling request").withCause(throwable).asException());
}
}
@Override
public void onError(Throwable t) {
t.printStackTrace();
responseObserver.onCompleted();
}
@Override
public void onCompleted() {
logger.info("gRPC Server has finished receiving data.");
responseObserver.onCompleted();
}
};
}
@Override
public void start() {
this.running = true;
}
@Override
public void stop() {
this.executor.shutdown();
try {
Assert.isTrue(this.executor.awaitTermination(5000, TimeUnit.MILLISECONDS), "gRPC Server executor timed out while stopping, "
+ "since there are currently executing tasks");
}
catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
this.running = false;
}
@Override
public boolean isRunning() {
return this.running;
}
@SuppressWarnings("unchecked")
private StreamObserver<GrpcMessage> biStreamReactive(StreamObserver<GrpcMessage> responseObserver,
ServerCallStreamObserver<GrpcMessage> serverCallStreamObserver) {
Many<Message<byte[]>> inputStream = Sinks.many().unicast().onBackpressureBuffer();
Flux<Message<byte[]>> inputStreamFlux = inputStream.asFlux();
FunctionInvocationWrapper function = this.resolveFunction(null);
Publisher<Message<byte[]>> outputPublisher = (Publisher<Message<byte[]>>) function.apply(inputStreamFlux);
Flux.from(outputPublisher).subscribe(functionResult -> {
GrpcMessage outputMessage = GrpcUtils.toGrpcMessage(functionResult);
if (logger.isDebugEnabled()) {
logger.debug("gRPC Server replying: " + outputMessage);
}
responseObserver.onNext(outputMessage);
});
return new StreamObserver<GrpcMessage>() {
@Override
public void onNext(GrpcMessage inputMessage) {
if (logger.isDebugEnabled()) {
logger.debug("gRPC Server receiving: " + inputMessage);
}
inputStream.tryEmitNext(GrpcUtils.fromGrpcMessage(inputMessage));
serverCallStreamObserver.request(1);
}
@Override
public void onError(Throwable t) {
t.printStackTrace();
responseObserver.onCompleted();
}
@Override
public void onCompleted() {
logger.info("gRPC Server has finished receiving data.");
inputStream.tryEmitComplete();
responseObserver.onCompleted();
}
};
}
private FunctionInvocationWrapper resolveFunction(Map<String, Object> headers) {
String functionDefinition = funcProperties.getDefinition();
if (!CollectionUtils.isEmpty(headers) && headers.containsKey(FunctionProperties.FUNCTION_DEFINITION)) {
functionDefinition = (String) headers.get(FunctionProperties.FUNCTION_DEFINITION);
}
FunctionInvocationWrapper function = this.functionCatalog.lookup(functionDefinition, "application/json");
Assert.notNull(function, "Failed to lookup function " + funcProperties.getDefinition());
return function;
public StreamObserver<GrpcSpringMessage> biStream(StreamObserver<GrpcSpringMessage> responseObserver) {
return this.helper.biStream(responseObserver, GrpcSpringMessage.class);
}
}

View File

@@ -0,0 +1,59 @@
/*
* Copyright 2021-2021 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.cloud.function.grpc;
import java.util.HashMap;
import java.util.Map;
import org.springframework.messaging.Message;
import org.springframework.messaging.support.MessageBuilder;
import com.google.protobuf.ByteString;
import com.google.protobuf.GeneratedMessageV3;
/**
*
* @author Oleg Zhurakousky
*
*/
public class GrpcSpringMessageConverter extends AbstractGrpcMessageConverter<GrpcSpringMessage> {
@Override
protected Message<byte[]> doToSpringMessage(GrpcSpringMessage grpcMessage) {
return MessageBuilder.withPayload(grpcMessage.getPayload().toByteArray())
.copyHeaders(grpcMessage.getHeadersMap())
.build();
}
@Override
protected GrpcSpringMessage doFromSpringMessage(Message<byte[]> springMessage) {
Map<String, String> stringHeaders = new HashMap<>();
springMessage.getHeaders().forEach((k, v) -> {
stringHeaders.put(k, v.toString());
});
return GrpcSpringMessage.newBuilder()
.setPayload(ByteString.copyFrom(springMessage.getPayload()))
.putAllHeaders(stringHeaders)
.build();
}
@Override
protected boolean supports(Class<? extends GeneratedMessageV3> grpcClass) {
return grpcClass.isAssignableFrom(GrpcSpringMessage.class);
}
}

View File

@@ -16,6 +16,7 @@
package org.springframework.cloud.function.grpc;
import java.util.HashMap;
import java.util.Iterator;
import java.util.Map;
@@ -24,6 +25,7 @@ import java.util.concurrent.Executors;
import java.util.concurrent.LinkedBlockingQueue;
import java.util.concurrent.TimeUnit;
import com.google.protobuf.ByteString;
import io.grpc.ManagedChannel;
import io.grpc.ManagedChannelBuilder;
@@ -53,22 +55,22 @@ final class GrpcUtils {
}
public static GrpcMessage toGrpcMessage(byte[] payload, Map<String, String> headers) {
return GrpcMessage.newBuilder()
public static GrpcSpringMessage toGrpcSpringMessage(byte[] payload, Map<String, String> headers) {
return GrpcSpringMessage.newBuilder()
.setPayload(ByteString.copyFrom(payload))
.putAllHeaders(headers)
.build();
}
public static GrpcMessage toGrpcMessage(Message<byte[]> message) {
public static GrpcSpringMessage toGrpcSpringMessage(Message<byte[]> message) {
Map<String, String> stringHeaders = new HashMap<>();
message.getHeaders().forEach((k, v) -> {
stringHeaders.put(k, v.toString());
});
return toGrpcMessage(message.getPayload(), stringHeaders);
return toGrpcSpringMessage(message.getPayload(), stringHeaders);
}
public static Message<byte[]> fromGrpcMessage(GrpcMessage message) {
public static Message<byte[]> fromGrpcSpringMessage(GrpcSpringMessage message) {
return MessageBuilder.withPayload(message.getPayload().toByteArray())
.copyHeaders(message.getHeadersMap())
.build();
@@ -84,9 +86,9 @@ final class GrpcUtils {
MessagingServiceGrpc.MessagingServiceBlockingStub stub = MessagingServiceGrpc
.newBlockingStub(channel);
GrpcMessage response = stub.requestReply(toGrpcMessage(inputMessage));
GrpcSpringMessage response = stub.requestReply(toGrpcSpringMessage(inputMessage));
channel.shutdown();
return fromGrpcMessage(response);
return fromGrpcSpringMessage(response);
}
/**
@@ -121,7 +123,7 @@ final class GrpcUtils {
.newStub(channel);
Many<Message<byte[]>> sink = Sinks.many().unicast().onBackpressureBuffer();
ClientResponseObserver<GrpcMessage, GrpcMessage> clientResponseObserver = clientResponseObserver(inputStream, sink);
ClientResponseObserver<GrpcSpringMessage, GrpcSpringMessage> clientResponseObserver = clientResponseObserver(inputStream, sink);
stub.biStream(clientResponseObserver);
@@ -137,14 +139,14 @@ final class GrpcUtils {
MessagingServiceGrpc.MessagingServiceBlockingStub stub = MessagingServiceGrpc
.newBlockingStub(channel);
Iterator<GrpcMessage> serverStream = stub.serverStream(toGrpcMessage(inputMessage));
Iterator<GrpcSpringMessage> serverStream = stub.serverStream(toGrpcSpringMessage(inputMessage));
Many<Message<byte[]>> sink = Sinks.many().unicast().onBackpressureBuffer();
ExecutorService executor = Executors.newSingleThreadExecutor();
executor.execute(() -> {
while (serverStream.hasNext()) {
GrpcMessage grpcMessage = serverStream.next();
sink.tryEmitNext(GrpcUtils.fromGrpcMessage(grpcMessage));
GrpcSpringMessage grpcMessage = serverStream.next();
sink.tryEmitNext(GrpcUtils.fromGrpcSpringMessage(grpcMessage));
}
sink.tryEmitComplete();
});
@@ -182,13 +184,13 @@ final class GrpcUtils {
.usePlaintext().build();
LinkedBlockingQueue<Message<byte[]>> resultRef = new LinkedBlockingQueue<>(1);
StreamObserver<GrpcMessage> responseObserver = new StreamObserver<GrpcMessage>() {
StreamObserver<GrpcSpringMessage> responseObserver = new StreamObserver<GrpcSpringMessage>() {
@Override
public void onNext(GrpcMessage result) {
public void onNext(GrpcSpringMessage result) {
if (logger.isDebugEnabled()) {
logger.debug("Client received reply: " + result);
}
resultRef.offer(GrpcUtils.fromGrpcMessage(result));
resultRef.offer(GrpcUtils.fromGrpcSpringMessage(result));
}
@Override
@@ -204,14 +206,14 @@ final class GrpcUtils {
MessagingServiceGrpc.MessagingServiceStub asyncStub = MessagingServiceGrpc.newStub(channel);
StreamObserver<GrpcMessage> requestObserver = asyncStub.clientStream(responseObserver);
StreamObserver<GrpcSpringMessage> requestObserver = asyncStub.clientStream(responseObserver);
inputStream.doOnNext(message -> {
if (logger.isDebugEnabled()) {
logger.debug("Client sending: " + message);
}
try {
requestObserver.onNext(GrpcUtils.toGrpcMessage(message));
requestObserver.onNext(GrpcUtils.toGrpcSpringMessage(message));
}
catch (Exception e) {
requestObserver.onError(e);
@@ -229,13 +231,13 @@ final class GrpcUtils {
}
}
private static ClientResponseObserver<GrpcMessage, GrpcMessage> clientResponseObserver(Flux<Message<byte[]>> inputStream, Many<Message<byte[]>> sink) {
return new ClientResponseObserver<GrpcMessage, GrpcMessage>() {
private static ClientResponseObserver<GrpcSpringMessage, GrpcSpringMessage> clientResponseObserver(Flux<Message<byte[]>> inputStream, Many<Message<byte[]>> sink) {
return new ClientResponseObserver<GrpcSpringMessage, GrpcSpringMessage>() {
ClientCallStreamObserver<GrpcMessage> requestStreamObserver;
ClientCallStreamObserver<GrpcSpringMessage> requestStreamObserver;
@Override
public void beforeStart(ClientCallStreamObserver<GrpcMessage> requestStreamObserver) {
public void beforeStart(ClientCallStreamObserver<GrpcSpringMessage> requestStreamObserver) {
this.requestStreamObserver = requestStreamObserver;
requestStreamObserver.disableAutoInboundFlowControl();
@@ -247,7 +249,7 @@ final class GrpcUtils {
if (logger.isDebugEnabled()) {
logger.debug("Streaming message to function: " + request);
}
requestStreamObserver.onNext(GrpcUtils.toGrpcMessage(request));
requestStreamObserver.onNext(GrpcUtils.toGrpcSpringMessage(request));
})
.doOnComplete(() -> {
requestStreamObserver.onCompleted();
@@ -258,11 +260,11 @@ final class GrpcUtils {
}
@Override
public void onNext(GrpcMessage message) {
public void onNext(GrpcSpringMessage message) {
if (logger.isDebugEnabled()) {
logger.debug("Streaming message from function: " + message);
}
sink.tryEmitNext(fromGrpcMessage(message));
sink.tryEmitNext(fromGrpcSpringMessage(message));
requestStreamObserver.request(1);
}

View File

@@ -0,0 +1,352 @@
/*
* Copyright 2021-2021 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.cloud.function.grpc;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.LinkedBlockingQueue;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicBoolean;
import com.google.protobuf.GeneratedMessageV3;
import io.grpc.Status;
import io.grpc.stub.ServerCallStreamObserver;
import io.grpc.stub.StreamObserver;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Sinks;
import reactor.core.publisher.Sinks.Many;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.reactivestreams.Publisher;
import org.springframework.cloud.function.context.FunctionCatalog;
import org.springframework.cloud.function.context.FunctionProperties;
import org.springframework.cloud.function.context.catalog.SimpleFunctionRegistry.FunctionInvocationWrapper;
import org.springframework.context.SmartLifecycle;
import org.springframework.messaging.Message;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
/**
*
* @author Oleg Zhurakousky
*
*/
public class MessageHandlingHelper<T extends GeneratedMessageV3> implements SmartLifecycle {
private Log logger = LogFactory.getLog(MessageHandlingHelper.class);
private final List<GrpcMessageConverter<?>> grpcConverters;
private final FunctionProperties funcProperties;
private final FunctionCatalog functionCatalog;
private final ExecutorService executor;
private boolean running;
public MessageHandlingHelper(List<GrpcMessageConverter<?>> grpcConverters,
FunctionCatalog functionCatalog, FunctionProperties funcProperties) {
this.grpcConverters = grpcConverters;
this.funcProperties = funcProperties;
this.functionCatalog = functionCatalog;
this.executor = Executors.newCachedThreadPool();
}
@SuppressWarnings("unchecked")
public void requestReply(T request, StreamObserver<T> responseObserver) {
Message<byte[]> message = this.toSpringMessage(request);
FunctionInvocationWrapper function = this.resolveFunction(message.getHeaders());
Message<byte[]> replyMessage = (Message<byte[]>) function.apply(message);
GeneratedMessageV3 reply = this.toGrpcMessage(replyMessage, (Class<T>) request.getClass());
responseObserver.onNext((T) reply);
responseObserver.onCompleted();
}
@SuppressWarnings("unchecked")
public void serverStream(T request, StreamObserver<T> responseObserver) {
Message<byte[]> message = this.toSpringMessage(request);
FunctionInvocationWrapper function = this.resolveFunction(message.getHeaders());
Publisher<Message<byte[]>> replyStream = (Publisher<Message<byte[]>>) function.apply(message);
Flux.from(replyStream).doOnNext(replyMessage -> {
responseObserver.onNext(this.toGrpcMessage(replyMessage, (Class<T>) request.getClass()));
})
.doOnComplete(() -> responseObserver.onCompleted())
.subscribe();
}
@SuppressWarnings("unchecked")
public StreamObserver<T> clientStream(StreamObserver<T> responseObserver, Class<T> grpcMessageType) {
ServerCallStreamObserver<T> serverCallStreamObserver = (ServerCallStreamObserver<T>) responseObserver;
serverCallStreamObserver.disableAutoInboundFlowControl();
FunctionInvocationWrapper function = this.resolveFunction(null);
AtomicBoolean wasReady = new AtomicBoolean(false);
serverCallStreamObserver.setOnReadyHandler(() -> {
if (serverCallStreamObserver.isReady() && !wasReady.get()) {
wasReady.set(true);
logger.info("gRPC Server receiving stream is ready.");
serverCallStreamObserver.request(1);
}
});
if (!function.isInputTypePublisher()) {
throw new UnsupportedOperationException("The client streaming is "
+ "not supported for functions that accept non-Publisher: "
+ function);
}
else if (function.isOutputTypePublisher()) {
throw new UnsupportedOperationException("The client streaming is "
+ "not supported for functions that return Publisher: "
+ function);
}
else {
Many<Message<byte[]>> inputStream = Sinks.many().unicast().onBackpressureBuffer();
Flux<Message<byte[]>> inputStreamFlux = inputStream.asFlux();
LinkedBlockingQueue<Message<byte[]>> resultRef = new LinkedBlockingQueue<>(1);
this.executor.execute(() -> {
Message<byte[]> replyMessage = (Message<byte[]>) function.apply(inputStreamFlux);
if (logger.isDebugEnabled()) {
logger.debug("Function invocation reply: " + replyMessage);
}
resultRef.offer(replyMessage);
});
return new StreamObserver<T>() {
@Override
public void onNext(T inputMessage) {
if (logger.isDebugEnabled()) {
logger.debug("gRPC Server receiving: " + inputMessage);
}
inputStream.tryEmitNext(toSpringMessage(inputMessage));
serverCallStreamObserver.request(1);
}
@Override
public void onError(Throwable t) {
t.printStackTrace();
responseObserver.onCompleted();
}
@Override
public void onCompleted() {
logger.info("gRPC Server has finished receiving data.");
inputStream.tryEmitComplete();
try {
responseObserver.onNext(toGrpcMessage(resultRef.poll(Integer.MAX_VALUE, TimeUnit.MILLISECONDS), grpcMessageType));
}
catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
responseObserver.onCompleted();
}
};
}
}
public StreamObserver<T> biStream(StreamObserver<T> responseObserver, Class<T> grpcMessageType) {
ServerCallStreamObserver<T> serverCallStreamObserver = (ServerCallStreamObserver<T>) responseObserver;
serverCallStreamObserver.disableAutoInboundFlowControl();
FunctionInvocationWrapper function = this.resolveFunction(null);
AtomicBoolean wasReady = new AtomicBoolean(false);
serverCallStreamObserver.setOnReadyHandler(() -> {
if (serverCallStreamObserver.isReady() && !wasReady.get()) {
wasReady.set(true);
logger.info("gRPC Server receiving stream is ready.");
serverCallStreamObserver.request(1);
}
});
if (function.isInputTypePublisher()) {
if (function.isOutputTypePublisher()) {
return this.biStreamReactive(responseObserver, serverCallStreamObserver, grpcMessageType);
}
throw new UnsupportedOperationException("The bi-directional streaming is "
+ "not supported for functions that accept Publisher but return non-Publisher: "
+ function);
}
else {
if (!function.isOutputTypePublisher()) {
return this.biStreamImperative(responseObserver, serverCallStreamObserver, wasReady);
}
throw new UnsupportedOperationException("The bidirection streaming is "
+ "not supported for functions that accept non-Publisher but return Publisher: "
+ function);
}
}
@SuppressWarnings("unchecked")
private StreamObserver<T> biStreamReactive(StreamObserver<T> responseObserver,
ServerCallStreamObserver<T> serverCallStreamObserver, Class<T> grpcMessageType) {
Many<Message<byte[]>> inputStream = Sinks.many().unicast().onBackpressureBuffer();
Flux<Message<byte[]>> inputStreamFlux = inputStream.asFlux();
FunctionInvocationWrapper function = this.resolveFunction(null);
Publisher<Message<byte[]>> outputPublisher = (Publisher<Message<byte[]>>) function.apply(inputStreamFlux);
Flux.from(outputPublisher).subscribe(functionResult -> {
T outputMessage = toGrpcMessage(functionResult, grpcMessageType);
if (logger.isDebugEnabled()) {
logger.debug("gRPC Server replying: " + outputMessage);
}
responseObserver.onNext(outputMessage);
});
return new StreamObserver<T>() {
@Override
public void onNext(T inputMessage) {
if (logger.isDebugEnabled()) {
logger.debug("gRPC Server receiving: " + inputMessage);
}
//GRPC_MESSAGE_TYPE = (Class<T>) inputMessage.getClass();
inputStream.tryEmitNext(toSpringMessage(inputMessage));
serverCallStreamObserver.request(1);
}
@Override
public void onError(Throwable t) {
t.printStackTrace();
responseObserver.onCompleted();
}
@Override
public void onCompleted() {
logger.info("gRPC Server has finished receiving data.");
inputStream.tryEmitComplete();
responseObserver.onCompleted();
}
};
}
private StreamObserver<T> biStreamImperative(StreamObserver<T> responseObserver,
ServerCallStreamObserver<T> serverCallStreamObserver,
AtomicBoolean wasReady) {
return new StreamObserver<T>() {
@SuppressWarnings("unchecked")
@Override
public void onNext(T request) {
try {
Message<byte[]> message = toSpringMessage(request);
FunctionInvocationWrapper function = resolveFunction(
message.getHeaders());
Message<byte[]> replyMessage = (Message<byte[]>) function
.apply(message);
T reply = toGrpcMessage(replyMessage, (Class<T>) request.getClass());
responseObserver.onNext(reply);
// Check the provided ServerCallStreamObserver to see if it is still
// ready to accept more messages.
if (serverCallStreamObserver.isReady()) {
serverCallStreamObserver.request(1);
}
else {
wasReady.set(false);
}
}
catch (Throwable throwable) {
throwable.printStackTrace();
responseObserver.onError(
Status.UNKNOWN.withDescription("Error handling request")
.withCause(throwable).asException());
}
}
@Override
public void onError(Throwable t) {
t.printStackTrace();
responseObserver.onCompleted();
}
@Override
public void onCompleted() {
logger.info("gRPC Server has finished receiving data.");
responseObserver.onCompleted();
}
};
}
@SuppressWarnings({ "rawtypes", "unchecked" })
private T toGrpcMessage(Message<byte[]> request, Class<T> grpcClass) {
for (GrpcMessageConverter converter : this.grpcConverters) {
GeneratedMessageV3 grpcMessage = converter.fromSpringMessage(request, grpcClass);
if (grpcMessage != null) {
return (T) grpcMessage;
}
}
throw new IllegalStateException("Failed to convert Grpc Message to Spring Message: " + request);
}
@Override
public void start() {
this.running = true;
}
@Override
public void stop() {
this.executor.shutdown();
try {
Assert.isTrue(this.executor.awaitTermination(5000, TimeUnit.MILLISECONDS), "gRPC Server executor timed out while stopping, "
+ "since there are currently executing tasks");
}
catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
this.running = false;
}
@Override
public boolean isRunning() {
return this.running;
}
@SuppressWarnings({ "rawtypes", "unchecked" })
private Message<byte[]> toSpringMessage(GeneratedMessageV3 request) {
for (GrpcMessageConverter converter : this.grpcConverters) {
Message<byte[]> springMessage = converter.toSpringMessage(request);
if (springMessage != null) {
return springMessage;
}
}
throw new IllegalStateException("Failed to convert Grpc Message to Spring Message: " + request);
}
private FunctionInvocationWrapper resolveFunction(Map<String, Object> headers) {
String functionDefinition = funcProperties.getDefinition();
if (!CollectionUtils.isEmpty(headers) && headers.containsKey(FunctionProperties.FUNCTION_DEFINITION)) {
functionDefinition = (String) headers.get(FunctionProperties.FUNCTION_DEFINITION);
}
FunctionInvocationWrapper function = this.functionCatalog.lookup(functionDefinition, "application/json");
Assert.notNull(function, "Failed to lookup function " + funcProperties.getDefinition());
return function;
}
}

View File

@@ -2,17 +2,17 @@ syntax = "proto3";
option java_multiple_files = true;
package org.springframework.cloud.function.grpc;
message GrpcMessage {
message GrpcSpringMessage {
bytes payload = 1;
map<string, string> headers = 2;
}
service MessagingService {
rpc biStream(stream GrpcMessage) returns (stream GrpcMessage);
rpc biStream(stream GrpcSpringMessage) returns (stream GrpcSpringMessage);
rpc clientStream(stream GrpcMessage) returns (GrpcMessage);
rpc clientStream(stream GrpcSpringMessage) returns (GrpcSpringMessage);
rpc serverStream(GrpcMessage) returns (stream GrpcMessage);
rpc serverStream(GrpcSpringMessage) returns (stream GrpcSpringMessage);
rpc requestReply(GrpcMessage) returns (GrpcMessage);
rpc requestReply(GrpcSpringMessage) returns (GrpcSpringMessage);
}