55
66package software .amazon .smithy .java .aws .events ;
77
8+ import java .nio .ByteBuffer ;
9+ import java .time .Instant ;
10+ import java .util .Map ;
811import java .util .Objects ;
912import java .util .concurrent .Flow ;
1013import java .util .function .Supplier ;
14+ import software .amazon .eventstream .HeaderValue ;
1115import software .amazon .eventstream .Message ;
1216import software .amazon .smithy .java .core .schema .Schema ;
1317import software .amazon .smithy .java .core .schema .SerializableStruct ;
@@ -52,7 +56,7 @@ public final class AwsEventShapeDecoder<E extends SerializableStruct, IR extends
5256 public SerializableStruct decode (AwsEventFrame frame ) {
5357 var message = frame .unwrap ();
5458 var eventType = getEventType (message );
55- if (initialEventType .getName ().equals (eventType )) {
59+ if (initialEventType .value ().equals (eventType )) {
5660 return decodeInitialResponse (frame );
5761 }
5862 return decodeEvent (frame );
@@ -71,9 +75,14 @@ private E decodeEvent(AwsEventFrame frame) {
7175 throw new IllegalArgumentException ("Unsupported event type: " + eventType );
7276 }
7377 var codecDeserializer = codec .createDeserializer (message .getPayload ());
74- var eventDeserializer = new AwsEventDeserializer (memberSchema , codecDeserializer );
78+ var headers = message .getHeaders ();
79+ var deserializer = new EventStreamDeserializer (codecDeserializer , new HeadersDeserializer (headers ));
80+ var memberTarget = memberSchema .memberTarget ();
81+ var shapeBuilder = memberTarget .shapeBuilder ();
82+ shapeBuilder .deserialize (deserializer );
7583 var builder = eventBuilder .get ();
76- return builder .deserialize (eventDeserializer ).build ();
84+ builder .setMemberValue (memberSchema , shapeBuilder .build ());
85+ return builder .build ();
7786 }
7887
7988 private IR decodeInitialResponse (AwsEventFrame frame ) {
@@ -82,8 +91,13 @@ private IR decodeInitialResponse(AwsEventFrame frame) {
8291 var builder = initialEventBuilder .get ();
8392 builder .deserialize (codecDeserializer );
8493 var publisherMember = getPublisherMember (builder .schema ());
85- var responseDeserializer = new EventStreamDeserializer (publisherMember , publisher );
94+ // Set the publisher member
95+ var responseDeserializer = new InitialResponseDeserializer (publisherMember , publisher );
8696 builder .deserialize (responseDeserializer );
97+ // Deserialize the rest of the members if any
98+ var headers = message .getHeaders ();
99+ var deserializer = new EventStreamDeserializer (codecDeserializer , new HeadersDeserializer (headers ));
100+ builder .deserialize (deserializer );
87101 return builder .build ();
88102 }
89103
@@ -100,11 +114,11 @@ private String getEventType(Message message) {
100114 return message .getHeaders ().get (":event-type" ).getString ();
101115 }
102116
103- static class EventStreamDeserializer extends SpecificShapeDeserializer {
117+ static class InitialResponseDeserializer extends SpecificShapeDeserializer {
104118 private final Schema publisherMember ;
105119 private final Flow .Publisher <? extends SerializableStruct > publisher ;
106120
107- EventStreamDeserializer (Schema publisherMember , Flow .Publisher <? extends SerializableStruct > publisher ) {
121+ InitialResponseDeserializer (Schema publisherMember , Flow .Publisher <? extends SerializableStruct > publisher ) {
108122 this .publisherMember = publisherMember ;
109123 this .publisher = publisher ;
110124 }
@@ -119,4 +133,99 @@ public <T> void readStruct(Schema schema, T state, ShapeDeserializer.StructMembe
119133 consumer .accept (state , publisherMember , this );
120134 }
121135 }
136+
137+ static class EventStreamDeserializer extends SpecificShapeDeserializer {
138+ private final ShapeDeserializer codecDeserializer ;
139+ private final HeadersDeserializer headersDeserializer ;
140+
141+ EventStreamDeserializer (ShapeDeserializer codecDeserializer , HeadersDeserializer headersDeserializer ) {
142+ this .codecDeserializer = codecDeserializer ;
143+ this .headersDeserializer = headersDeserializer ;
144+ }
145+
146+ @ Override
147+ public <T > void readStruct (Schema schema , T builder , ShapeDeserializer .StructMemberConsumer <T > consumer ) {
148+ var payloadWritten = false ;
149+ for (Schema member : schema .members ()) {
150+ if (member .hasTrait (TraitKey .EVENT_HEADER_TRAIT )) {
151+ consumer .accept (builder , member , headersDeserializer );
152+ } else if (member .hasTrait (TraitKey .EVENT_PAYLOAD_TRAIT )) {
153+ consumer .accept (builder , member , codecDeserializer );
154+ payloadWritten = true ;
155+ }
156+ }
157+ // Deserialize from the payload if still needed.
158+ if (!payloadWritten ) {
159+ codecDeserializer .readStruct (schema , builder , consumer );
160+ }
161+ }
162+ }
163+
164+ static class HeadersDeserializer extends SpecificShapeDeserializer {
165+ private final Map <String , HeaderValue > headers ;
166+
167+ HeadersDeserializer (Map <String , HeaderValue > headers ) {
168+ this .headers = headers ;
169+ }
170+
171+ @ Override
172+ public ByteBuffer readBlob (Schema schema ) {
173+ return getValueForShapeType (schema );
174+ }
175+
176+ @ Override
177+
178+ public byte readByte (Schema schema ) {
179+ return getValueForShapeType (schema );
180+ }
181+
182+ @ Override
183+ public short readShort (Schema schema ) {
184+ return getValueForShapeType (schema );
185+ }
186+
187+ @ Override
188+ public int readInteger (Schema schema ) {
189+ return getValueForShapeType (schema );
190+ }
191+
192+ @ Override
193+ public long readLong (Schema schema ) {
194+ return getValueForShapeType (schema );
195+ }
196+
197+ @ Override
198+ public String readString (Schema schema ) {
199+ return getValueForShapeType (schema );
200+ }
201+
202+ @ Override
203+ public boolean readBoolean (Schema schema ) {
204+ return getValueForShapeType (schema );
205+ }
206+
207+ @ Override
208+ public Instant readTimestamp (Schema schema ) {
209+ return getValueForShapeType (schema );
210+ }
211+
212+ @ SuppressWarnings ("unchecked" )
213+ private <T > T getValueForShapeType (Schema member ) {
214+ HeaderValue value = headers .get (member .memberName ());
215+ if (value == null ) {
216+ return null ;
217+ }
218+ return (T ) switch (member .type ()) {
219+ case BLOB -> value .getByteBuffer ();
220+ case BOOLEAN -> value .getBoolean ();
221+ case BYTE -> value .getByte ();
222+ case SHORT -> value .getShort ();
223+ case INTEGER , INT_ENUM -> value .getInteger ();
224+ case LONG -> value .getLong ();
225+ case TIMESTAMP -> value .getTimestamp ();
226+ case STRING -> value .getString ();
227+ default -> throw new IllegalArgumentException ("Unsupported shape type: " + member .type ());
228+ };
229+ }
230+ }
122231}
0 commit comments