View Javadoc
1   /*
2   * Copyright 2019-2026 the original author or authors.
3    *
4    * Licensed under the Apache License, Version 2.0 (the "License");
5    * you may not use this file except in compliance with the License.
6    * You may obtain a copy of the License at
7    *
8    *      http://www.apache.org/licenses/LICENSE-2.0
9    *
10   * Unless required by applicable law or agreed to in writing, software
11   * distributed under the License is distributed on an "AS IS" BASIS,
12   * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13   * See the License for the specific language governing permissions and
14   * limitations under the License.
15   */
16  
17  package org.bremersee.xml.http.codec;
18  
19  import static java.util.Objects.isNull;
20  import static java.util.Objects.nonNull;
21  
22  import jakarta.xml.bind.JAXBElement;
23  import jakarta.xml.bind.JAXBException;
24  import jakarta.xml.bind.UnmarshalException;
25  import jakarta.xml.bind.Unmarshaller;
26  import jakarta.xml.bind.annotation.XmlRootElement;
27  import jakarta.xml.bind.annotation.XmlSchema;
28  import jakarta.xml.bind.annotation.XmlSeeAlso;
29  import jakarta.xml.bind.annotation.XmlType;
30  import java.util.ArrayList;
31  import java.util.HashSet;
32  import java.util.Iterator;
33  import java.util.List;
34  import java.util.Map;
35  import java.util.Optional;
36  import java.util.Set;
37  import java.util.function.BiConsumer;
38  import javax.xml.XMLConstants;
39  import javax.xml.namespace.QName;
40  import javax.xml.stream.XMLEventReader;
41  import javax.xml.stream.XMLInputFactory;
42  import javax.xml.stream.XMLStreamException;
43  import javax.xml.stream.events.XMLEvent;
44  import org.bremersee.xml.JaxbContextBuilder;
45  import org.jspecify.annotations.NonNull;
46  import org.jspecify.annotations.Nullable;
47  import org.reactivestreams.Publisher;
48  import org.springframework.core.ResolvableType;
49  import org.springframework.core.codec.AbstractDecoder;
50  import org.springframework.core.codec.CodecException;
51  import org.springframework.core.codec.DecodingException;
52  import org.springframework.core.codec.Hints;
53  import org.springframework.core.io.buffer.DataBuffer;
54  import org.springframework.core.io.buffer.DataBufferLimitException;
55  import org.springframework.core.io.buffer.DataBufferUtils;
56  import org.springframework.core.log.LogFormatUtils;
57  import org.springframework.http.MediaType;
58  import org.springframework.http.codec.xml.XmlEventDecoder;
59  import org.springframework.http.codec.xml.XmlEventDecoder.ReceivedByteTracker;
60  import org.springframework.util.Assert;
61  import org.springframework.util.ClassUtils;
62  import org.springframework.util.MimeType;
63  import org.springframework.util.MimeTypeUtils;
64  import org.springframework.util.xml.StaxUtils;
65  import reactor.core.Exceptions;
66  import reactor.core.publisher.Flux;
67  import reactor.core.publisher.Mono;
68  import reactor.core.publisher.SynchronousSink;
69  
70  /**
71   * Decode from a bytes stream containing XML elements to a stream of {@code Object}s (POJOs).
72   *
73   * <p>The decoding parts are taken from {@link org.springframework.http.codec.xml.Jaxb2XmlDecoder}.
74   *
75   * @author Sebastien Deleuze
76   * @author Arjen Poutsma
77   * @author Christian Bremer
78   */
79  public class ReactiveJaxbDecoder extends AbstractDecoder<Object> {
80  
81    /**
82     * The default value for JAXB annotations.
83     *
84     * @see XmlRootElement#name()
85     * @see XmlRootElement#namespace()
86     * @see XmlType#name()
87     * @see XmlType#namespace()
88     */
89    private static final String JAXB_DEFAULT_ANNOTATION_VALUE = "##default";
90  
91    private static final XMLInputFactory inputFactory = StaxUtils.createDefensiveInputFactory();
92  
93    private final XmlEventDecoder xmlEventDecoder = new XmlEventDecoder();
94  
95    private final JaxbContextBuilder jaxbContextBuilder;
96  
97    private final Set<Class<?>> ignoreReadingClasses;
98  
99    private int maxInMemorySize = 256 * 1024;
100 
101   /**
102    * Instantiates a new reactive jaxb decoder.
103    *
104    * @param jaxbContextBuilder the jaxb context builder
105    */
106   public ReactiveJaxbDecoder(JaxbContextBuilder jaxbContextBuilder) {
107     this(jaxbContextBuilder, null);
108   }
109 
110   /**
111    * Instantiates a new reactive jaxb decoder.
112    *
113    * @param jaxbContextBuilder the jaxb context builder
114    * @param ignoreReadingClasses the ignore reading classes
115    */
116   public ReactiveJaxbDecoder(
117       JaxbContextBuilder jaxbContextBuilder,
118       Set<Class<?>> ignoreReadingClasses) {
119 
120     super(MimeTypeUtils.APPLICATION_XML, MimeTypeUtils.TEXT_XML,
121         new MediaType("application", "*+xml"));
122     Assert.notNull(jaxbContextBuilder, "JaxbContextBuilder must be present.");
123     this.jaxbContextBuilder = jaxbContextBuilder;
124     this.ignoreReadingClasses = isNull(ignoreReadingClasses) ? Set.of() : ignoreReadingClasses;
125   }
126 
127   /**
128    * Set the max number of bytes that can be buffered by this decoder. This is either the size of
129    * the entire input when decoding as a whole, or when using async parsing with Aalto XML, it is
130    * the size of one top-level XML tree. When the limit is exceeded,
131    * {@link org.springframework.core.io.buffer.DataBufferLimitException} is raised.
132    *
133    * <p>By default, this is set to 256K.
134    *
135    * @param byteCount the max number of bytes to buffer, or -1 for unlimited
136    */
137   public void setMaxInMemorySize(int byteCount) {
138     this.maxInMemorySize = byteCount;
139     this.xmlEventDecoder.setMaxInMemorySize(byteCount);
140   }
141 
142   /**
143    * Return the {@link #setMaxInMemorySize configured} byte count limit.
144    *
145    * @return the max in memory size
146    */
147   public int getMaxInMemorySize() {
148     return this.maxInMemorySize;
149   }
150 
151   @Override
152   public boolean canDecode(
153       @NonNull ResolvableType elementType,
154       @Nullable MimeType mimeType) {
155 
156     if (super.canDecode(elementType, mimeType)) {
157       final Class<?> outputClass = elementType.getRawClass();
158       return !ignoreReadingClasses.contains(outputClass)
159           && jaxbContextBuilder.canUnmarshal(outputClass);
160     } else {
161       return false;
162     }
163   }
164 
165   @NonNull
166   @Override
167   public Flux<Object> decode(
168       @NonNull Publisher<DataBuffer> inputStream,
169       ResolvableType elementType,
170       @Nullable MimeType mimeType,
171       @Nullable Map<String, Object> hints) {
172 
173     ReceivedByteTracker byteTracker = new ReceivedByteTracker(this.maxInMemorySize);
174 
175     Flux<XMLEvent> xmlEventFlux = this.xmlEventDecoder.decode(
176         inputStream, ResolvableType.forClass(XMLEvent.class), mimeType, hints);
177 
178     Class<?> outputClass = elementType.toClass();
179     Set<QName> names = toQualifiedNames(outputClass);
180     Flux<List<XMLEvent>> splitEvents = split(xmlEventFlux, names, byteTracker);
181 
182     return splitEvents.map(events -> {
183       Object value = unmarshal(events, outputClass);
184       LogFormatUtils.traceDebug(logger, traceOn -> {
185         String formatted = LogFormatUtils.formatValue(value, !traceOn);
186         return Hints.getLogPrefix(hints) + "Decoded [" + formatted + "]";
187       });
188       return value;
189     });
190   }
191 
192   @NonNull
193   @Override
194   @SuppressWarnings({"rawtypes", "unchecked", "cast"})
195   // XMLEventReader is Iterator<Object> on JDK 9
196   public Object decode(
197       DataBuffer dataBuffer,
198       ResolvableType targetType,
199       @Nullable MimeType mimeType,
200       @Nullable Map<String, Object> hints) throws DecodingException {
201 
202     try {
203       Iterator eventReader = inputFactory.createXMLEventReader(dataBuffer.asInputStream());
204       List<XMLEvent> events = new ArrayList<>();
205       eventReader.forEachRemaining(event -> events.add((XMLEvent) event));
206       return unmarshal(events, targetType.toClass());
207     } catch (XMLStreamException ex) {
208       throw Exceptions.propagate(ex);
209     } finally {
210       DataBufferUtils.release(dataBuffer);
211     }
212   }
213 
214   @NonNull
215   @Override
216   public Mono<Object> decodeToMono(
217       @NonNull Publisher<DataBuffer> input,
218       @NonNull ResolvableType elementType,
219       @Nullable MimeType mimeType,
220       @Nullable Map<String, Object> hints) {
221 
222     return DataBufferUtils.join(input, this.maxInMemorySize)
223         .map(dataBuffer -> decode(dataBuffer, elementType, mimeType, hints));
224   }
225 
226   private Object unmarshal(List<XMLEvent> events, Class<?> outputClass) {
227     try {
228       Unmarshaller unmarshaller = jaxbContextBuilder.buildUnmarshaller(outputClass);
229       XMLEventReader eventReader = StaxUtils.createXMLEventReader(events);
230       if (outputClass.isAnnotationPresent(XmlRootElement.class)) {
231         return unmarshaller.unmarshal(eventReader);
232       } else {
233         JAXBElement<?> jaxbElement = unmarshaller.unmarshal(eventReader, outputClass);
234         return jaxbElement.getValue();
235       }
236     } catch (UnmarshalException ex) {
237       throw new DecodingException("Could not unmarshal XML to " + outputClass, ex);
238     } catch (JAXBException ex) {
239       throw new CodecException("Invalid JAXB configuration", ex);
240     }
241   }
242 
243   /**
244    * Split a flux of {@link XMLEvent XMLEvents} into a flux of XMLEvent lists, one list for each
245    * branch of the tree that starts with the given qualified name. That is, given the XMLEvents
246    * shown {@linkplain XmlEventDecoder here}, and the {@code desiredName} "{@code child}", this
247    * method returns a flux of two lists, each of which containing the events of a particular branch
248    * of the tree that starts with "{@code child}".
249    * <ol>
250    * <li>The first list, dealing with the first branch of the tree:
251    * <ol>
252    * <li>{@link javax.xml.stream.events.StartElement} {@code child}</li>
253    * <li>{@link javax.xml.stream.events.Characters} {@code foo}</li>
254    * <li>{@link javax.xml.stream.events.EndElement} {@code child}</li>
255    * </ol>
256    * <li>The second list, dealing with the second branch of the tree:
257    * <ol>
258    * <li>{@link javax.xml.stream.events.StartElement} {@code child}</li>
259    * <li>{@link javax.xml.stream.events.Characters} {@code bar}</li>
260    * <li>{@link javax.xml.stream.events.EndElement} {@code child}</li>
261    * </ol>
262    * </li>
263    * </ol>
264    *
265    * @param xmlEventFlux the xml event as flux
266    * @param desiredNames the desired names
267    * @param byteTracker the byte tracker
268    * @return the list of xml events as flux
269    */
270   private static Flux<List<XMLEvent>> split(
271       Flux<XMLEvent> xmlEventFlux,
272       Set<QName> desiredNames,
273       XmlEventDecoder.ReceivedByteTracker byteTracker) {
274     return xmlEventFlux.handle(new SplitHandler(desiredNames, byteTracker));
275   }
276 
277   private static Set<QName> toQualifiedNames(Class<?> outputClass) {
278     Set<QName> result = HashSet.newHashSet(1);
279     findQNames(outputClass, result, new HashSet<>());
280     return result;
281   }
282 
283   private static void findQNames(Class<?> clazz, Set<QName> qNames, Set<Class<?>> completedClasses) {
284     // safety against circular XmlSeeAlso references
285     if (completedClasses.contains(clazz)) {
286       return;
287     }
288     if (clazz.isAnnotationPresent(XmlRootElement.class)) {
289       XmlRootElement annotation = clazz.getAnnotation(XmlRootElement.class);
290       qNames.add(new QName(namespace(annotation.namespace(), clazz),
291           localPart(annotation.name(), clazz)));
292     }
293     else if (clazz.isAnnotationPresent(XmlType.class)) {
294       XmlType annotation = clazz.getAnnotation(XmlType.class);
295       qNames.add(new QName(namespace(annotation.namespace(), clazz),
296           localPart(annotation.name(), clazz)));
297     }
298     else {
299       throw new IllegalArgumentException("Output class [" + clazz.getName() +
300           "] is neither annotated with @XmlRootElement nor @XmlType");
301     }
302     completedClasses.add(clazz);
303     if (clazz.isAnnotationPresent(XmlSeeAlso.class)) {
304       XmlSeeAlso annotation = clazz.getAnnotation(XmlSeeAlso.class);
305       for (Class<?> seeAlso : annotation.value()) {
306         findQNames(seeAlso, qNames, completedClasses);
307       }
308     }
309   }
310 
311   private static String localPart(String value, Class<?> outputClass) {
312     if (JAXB_DEFAULT_ANNOTATION_VALUE.equals(value)) {
313       return ClassUtils.getShortNameAsProperty(outputClass);
314     }
315     else {
316       return value;
317     }
318   }
319 
320   private static String namespace(String value, Class<?> outputClass) {
321     if (JAXB_DEFAULT_ANNOTATION_VALUE.equals(value)) {
322       Package outputClassPackage = outputClass.getPackage();
323       if (nonNull(outputClassPackage) && outputClassPackage.isAnnotationPresent(XmlSchema.class)) {
324         XmlSchema annotation = outputClassPackage.getAnnotation(XmlSchema.class);
325         return annotation.namespace();
326       }
327       else {
328         return XMLConstants.NULL_NS_URI;
329       }
330     }
331     else {
332       return value;
333     }
334   }
335 
336   private static class SplitHandler implements
337       BiConsumer<XMLEvent, SynchronousSink<List<XMLEvent>>> {
338 
339     private final Set<QName> names;
340 
341     private final ReceivedByteTracker byteTracker;
342 
343     private List<XMLEvent> events;
344 
345     private int elementDepth = 0;
346 
347     private int barrier = Integer.MAX_VALUE;
348 
349     /**
350      * Instantiates a new split handler.
351      *
352      * @param names the names
353      * @param byteTracker the byte tracker
354      */
355     public SplitHandler(Set<QName> names, ReceivedByteTracker byteTracker) {
356       this.names = names;
357       this.byteTracker = Optional.ofNullable(byteTracker).orElse(ReceivedByteTracker.NO_OP);
358     }
359 
360     @Override
361     public void accept(XMLEvent event, SynchronousSink<List<XMLEvent>> sink) {
362       if (event.isStartElement()) {
363         if (this.barrier == Integer.MAX_VALUE) {
364           QName startElementName = event.asStartElement().getName();
365           if (this.names.contains(startElementName)) {
366             this.events = new ArrayList<>();
367             this.barrier = this.elementDepth;
368           }
369         }
370         this.elementDepth++;
371       }
372       if (this.elementDepth > this.barrier) {
373         Assert.state(this.events != null, "No XMLEvent List");
374         this.events.add(event);
375       }
376       if (event.isEndElement()) {
377         this.elementDepth--;
378         if (this.elementDepth == this.barrier) {
379           Assert.state(this.events != null, "No XMLEvent List");
380           sink.next(this.events);
381           this.barrier = Integer.MAX_VALUE;
382           this.events = null;
383         }
384       }
385       if (isNull(this.events)) {
386         this.byteTracker.reset();
387       } else if (this.byteTracker.isMaxInMemorySizeExceeded()) {
388         throw new DataBufferLimitException(
389             "Exceeded limit on max bytes per XML node: " + this.byteTracker.getMaxInMemorySize());
390       }
391     }
392   }
393 
394 }