diff --git a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java index cb35f272dbe..8869abfb2f6 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java +++ b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java @@ -129,13 +129,22 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements Decoder decode(Publisher inputStream, ResolvableType elementType, @Nullable MimeType mimeType, @Nullable Map hints) { - MessageDecoderFunction decoderFunction = new MessageDecoderFunction(elementType, this.maxMessageSize); + MessageDecoderFunction decoderFunction = + new MessageDecoderFunction(elementType, this.maxMessageSize, initMessageSizeReader()); return Flux.from(inputStream) .flatMapIterable(decoderFunction) .doOnTerminate(decoderFunction::discard); } + /** + * Return a reader for message size information encoded in the input stream. + * @since 7.0 + */ + protected MessageSizeReader initMessageSizeReader() { + return new DefaultMessageSizeReader(); + } + @Override public Mono decodeToMono(Publisher inputStream, ResolvableType elementType, @Nullable MimeType mimeType, @Nullable Map hints) { @@ -150,9 +159,7 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements Decoder apply(DataBuffer input) { try { @@ -214,9 +231,11 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements Decoder 0 && this.messageBytesToRead > this.maxMessageSize) { throw new DataBufferLimitException( "The number of bytes to read for message " + @@ -262,56 +281,6 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements DecoderBase 128 Varints - */ - private boolean readMessageSize(DataBuffer input) { - if (this.offset == 0) { - if (input.readableByteCount() == 0) { - return false; - } - int firstByte = input.read(); - if ((firstByte & 0x80) == 0) { - this.messageBytesToRead = firstByte; - return true; - } - this.messageBytesToRead = firstByte & 0x7f; - this.offset = 7; - } - - if (this.offset < 32) { - for (; this.offset < 32; this.offset += 7) { - if (input.readableByteCount() == 0) { - return false; - } - final int b = input.read(); - this.messageBytesToRead |= (b & 0x7f) << this.offset; - if ((b & 0x80) == 0) { - this.offset = 0; - return true; - } - } - } - // Keep reading up to 64 bits. - for (; this.offset < 64; this.offset += 7) { - if (input.readableByteCount() == 0) { - return false; - } - final int b = input.read(); - if ((b & 0x80) == 0) { - this.offset = 0; - return true; - } - } - this.offset = 0; - throw new DecodingException("Cannot parse message size: malformed varint"); - } - public void discard() { if (this.output != null) { DataBufferUtils.release(this.output); @@ -319,4 +288,83 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements DecoderParses the message size as a varint from the input stream. + * Inspired by {@link CodedInputStream#readRawVarint32(int, java.io.InputStream)}, + * @see Base 128 Varints + */ + private static class DefaultMessageSizeReader implements MessageSizeReader { + + private int offset; + + private int messageSize; + + @Override + public @Nullable Integer readMessageSize(DataBuffer input) { + if (this.offset == 0) { + if (input.readableByteCount() == 0) { + return null; + } + int firstByte = input.read(); + if ((firstByte & 0x80) == 0) { + this.messageSize = firstByte; + return getAndReset(); + } + this.messageSize = firstByte & 0x7f; + this.offset = 7; + } + + if (this.offset < 32) { + for (; this.offset < 32; this.offset += 7) { + if (input.readableByteCount() == 0) { + return null; + } + final int b = input.read(); + this.messageSize |= (b & 0x7f) << this.offset; + if ((b & 0x80) == 0) { + return getAndReset(); + } + } + } + // Keep reading up to 64 bits. + for (; this.offset < 64; this.offset += 7) { + if (input.readableByteCount() == 0) { + return null; + } + final int b = input.read(); + if ((b & 0x80) == 0) { + return getAndReset(); + } + } + getAndReset(); + throw new DecodingException("Cannot parse message size: malformed varint"); + } + + private @Nullable Integer getAndReset() { + Integer result = (this.messageSize != 0 ? this.messageSize : null); + this.offset = 0; + this.messageSize = 0; + return result; + } + } + } diff --git a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufEncoder.java b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufEncoder.java index 5f7f5443299..6230591ba2a 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufEncoder.java +++ b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufEncoder.java @@ -107,15 +107,10 @@ public class ProtobufEncoder extends ProtobufCodecSupport implements HttpMessage } private DataBuffer encodeValue(Message message, DataBufferFactory bufferFactory, boolean delimited) { - FastByteArrayOutputStream bos = new FastByteArrayOutputStream(); + FastByteArrayOutputStream outputStream = new FastByteArrayOutputStream(); try { - if (delimited) { - message.writeDelimitedTo((OutputStream) bos); - } - else { - message.writeTo((OutputStream) bos); - } - byte[] bytes = bos.toByteArrayUnsafe(); + writeMessage(message, delimited, outputStream); + byte[] bytes = outputStream.toByteArrayUnsafe(); return bufferFactory.wrap(bytes); } catch (IOException ex) { @@ -123,4 +118,19 @@ public class ProtobufEncoder extends ProtobufCodecSupport implements HttpMessage } } + /** + * Use write methods on {@link Message} to write to the given {@code OutputStream}. + * @since 7.0 + */ + protected void writeMessage( + Message message, boolean delimited, OutputStream outputStream) throws IOException { + + if (delimited) { + message.writeDelimitedTo(outputStream); + } + else { + message.writeTo(outputStream); + } + } + } diff --git a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufHttpMessageWriter.java b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufHttpMessageWriter.java index 9c4cb07f51f..103c5dea6fc 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufHttpMessageWriter.java +++ b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufHttpMessageWriter.java @@ -77,30 +77,47 @@ public class ProtobufHttpMessageWriter extends EncoderHttpMessageWriter @SuppressWarnings("unchecked") @Override public Mono write(Publisher inputStream, ResolvableType elementType, - @Nullable MediaType mediaType, ReactiveHttpOutputMessage message, Map hints) { + @Nullable MediaType mediaType, ReactiveHttpOutputMessage outputMessage, Map hints) { try { Message.Builder builder = getMessageBuilder(elementType.toClass()); Descriptors.Descriptor descriptor = builder.getDescriptorForType(); - message.getHeaders().add(X_PROTOBUF_SCHEMA_HEADER, descriptor.getFile().getName()); - message.getHeaders().add(X_PROTOBUF_MESSAGE_HEADER, descriptor.getFullName()); + outputMessage.getHeaders().add(X_PROTOBUF_SCHEMA_HEADER, descriptor.getFile().getName()); + outputMessage.getHeaders().add(X_PROTOBUF_MESSAGE_HEADER, descriptor.getFullName()); if (inputStream instanceof Flux) { - if (mediaType == null) { - message.getHeaders().setContentType(((HttpMessageEncoder)getEncoder()).getStreamingMediaTypes().get(0)); - } - else if (!ProtobufEncoder.DELIMITED_VALUE.equals(mediaType.getParameters().get(ProtobufEncoder.DELIMITED_KEY))) { - Map parameters = new HashMap<>(mediaType.getParameters()); - parameters.put(ProtobufEncoder.DELIMITED_KEY, ProtobufEncoder.DELIMITED_VALUE); - message.getHeaders().setContentType(new MediaType(mediaType.getType(), mediaType.getSubtype(), parameters)); - } + outputMessage.getHeaders().setContentType(getStreamingContentType(mediaType)); } - return super.write(inputStream, elementType, mediaType, message, hints); + extendHeaders(outputMessage, hints); + return super.write(inputStream, elementType, mediaType, outputMessage, hints); } catch (Exception ex) { return Mono.error(new EncodingException("Could not write Protobuf message: " + ex.getMessage(), ex)); } } + /** + * Return the {@code MediaType} to use when the input Publisher is multivalued. + * @since 7.0 + */ + protected MediaType getStreamingContentType(@Nullable MediaType mediaType) { + if (mediaType == null) { + return ((HttpMessageEncoder) getEncoder()).getStreamingMediaTypes().get(0); + } + Map params = new HashMap<>(mediaType.getParameters()); + if (!ProtobufEncoder.DELIMITED_VALUE.equals(params.get(ProtobufEncoder.DELIMITED_KEY))) { + params.put(ProtobufEncoder.DELIMITED_KEY, ProtobufEncoder.DELIMITED_VALUE); + mediaType = new MediaType(mediaType, params); + } + return mediaType; + } + + /** + * Make further updates to headers. + * @since 7.0 + */ + protected void extendHeaders(ReactiveHttpOutputMessage message, Map hints) { + } + /** * Create a new {@code Message.Builder} instance for the given class. *

This method uses a ConcurrentHashMap for caching method lookups.