1
2
3
4
5
6
7
8
9
10
11
12
13
14
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
72
73
74
75
76
77
78
79 public class ReactiveJaxbDecoder extends AbstractDecoder<Object> {
80
81
82
83
84
85
86
87
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
103
104
105
106 public ReactiveJaxbDecoder(JaxbContextBuilder jaxbContextBuilder) {
107 this(jaxbContextBuilder, null);
108 }
109
110
111
112
113
114
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
129
130
131
132
133
134
135
136
137 public void setMaxInMemorySize(int byteCount) {
138 this.maxInMemorySize = byteCount;
139 this.xmlEventDecoder.setMaxInMemorySize(byteCount);
140 }
141
142
143
144
145
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
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
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
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
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
351
352
353
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 }