Skip to content

Commit bb58d49

Browse files
authored
handle union of multiple arrays (#1282)
1 parent 2d3045a commit bb58d49

3 files changed

Lines changed: 123 additions & 4 deletions

File tree

java-client/src/main/java/co/elastic/clients/json/JsonpDeserializerBase.java

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -286,6 +286,10 @@ protected ArrayDeserializer(JsonpDeserializer<T> itemDeserializer) {
286286
this.itemDeserializer = itemDeserializer;
287287
}
288288

289+
JsonpDeserializer<T> itemDeserializer() {
290+
return itemDeserializer;
291+
}
292+
289293
@Override
290294
public EnumSet<Event> nativeEvents() {
291295
return nativeEvents;

java-client/src/main/java/co/elastic/clients/json/UnionDeserializer.java

Lines changed: 83 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,8 @@
2020
package co.elastic.clients.json;
2121

2222
import co.elastic.clients.util.ObjectBuilder;
23+
import jakarta.json.JsonArray;
24+
import jakarta.json.JsonValue;
2325
import jakarta.json.stream.JsonLocation;
2426
import jakarta.json.stream.JsonParser;
2527
import jakarta.json.stream.JsonParser.Event;
@@ -117,11 +119,76 @@ Union deserialize(JsonParser parser, JsonpMapper mapper, Event event, BiFunction
117119
}
118120
}
119121

122+
/**
123+
* An event handler for arrays that disambiguates multiple array members by buffering the array and inspecting
124+
* the JSON event of its first element.
125+
*/
126+
private static class ArrayMemberHandler<Union, Kind, Member> extends EventHandler<Union, Kind, Member> {
127+
// Element JSON event -> array member accepting that element type
128+
private final Map<Event, SingleMemberHandler<Union, Kind, Member>> byElementEvent = new HashMap<>();
129+
// Fallback for empty arrays or unrecognized element types: first declared member wins
130+
private final SingleMemberHandler<Union, Kind, Member> defaultMember;
131+
132+
ArrayMemberHandler(List<SingleMemberHandler<Union, Kind, Member>> members) {
133+
this.defaultMember = members.get(0);
134+
for (SingleMemberHandler<Union, Kind, Member> m: members) {
135+
JsonpDeserializer<?> unwrapped = DelegatingDeserializer.unwrap(m.deserializer);
136+
JsonpDeserializer<?> item = ((JsonpDeserializerBase.ArrayDeserializer<?>) unwrapped).itemDeserializer();
137+
// Key on the element's native events (its canonical output), first writer wins on conflict
138+
for (Event e: item.nativeEvents()) {
139+
byElementEvent.putIfAbsent(e, m);
140+
}
141+
}
142+
}
143+
144+
@Override
145+
EnumSet<Event> nativeEvents() {
146+
return EnumSet.of(Event.START_ARRAY);
147+
}
148+
149+
@Override
150+
Union deserialize(JsonParser parser, JsonpMapper mapper, Event event, BiFunction<Kind, Member, Union> buildFn) {
151+
// event == START_ARRAY. Buffer the whole array so we can inspect it and then replay it.
152+
JsonArray array = parser.getArray();
153+
154+
SingleMemberHandler<Union, Kind, Member> member = defaultMember;
155+
// Find the first non-null element: leading nulls don't identify a variant
156+
for (JsonValue element: array) {
157+
if (element.getValueType() != JsonValue.ValueType.NULL) {
158+
SingleMemberHandler<Union, Kind, Member> found =
159+
byElementEvent.get(valueTypeToEvent(element.getValueType()));
160+
if (found != null) {
161+
member = found;
162+
}
163+
break;
164+
}
165+
}
166+
167+
// Replay the buffered array into the chosen member's array deserializer
168+
JsonParser arrayParser = JsonpUtils.jsonValueParser(array, mapper);
169+
return member.deserialize(arrayParser, mapper, arrayParser.next(), buildFn);
170+
}
171+
}
172+
173+
private static Event valueTypeToEvent(JsonValue.ValueType type) {
174+
switch (type) {
175+
case OBJECT: return Event.START_OBJECT;
176+
case ARRAY: return Event.START_ARRAY;
177+
case STRING: return Event.VALUE_STRING;
178+
case NUMBER: return Event.VALUE_NUMBER;
179+
case TRUE: return Event.VALUE_TRUE;
180+
case FALSE: return Event.VALUE_FALSE;
181+
case NULL: return Event.VALUE_NULL;
182+
default: throw new IllegalArgumentException("Unknown JSON value type: " + type);
183+
}
184+
}
185+
120186
public static class Builder<Union, Kind, Member> implements ObjectBuilder<JsonpDeserializer<Union>> {
121187

122188
private final BiFunction<Kind, Member, Union> buildFn;
123189

124190
private final List<SingleMemberHandler<Union, Kind, Member>> objectMembers = new ArrayList<>();
191+
private final List<SingleMemberHandler<Union, Kind, Member>> arrayMembers = new ArrayList<>();
125192
private final Map<Event, EventHandler<Union, Kind, Member>> otherMembers = new HashMap<>();
126193
private final boolean allowAmbiguousPrimitive;
127194

@@ -185,6 +252,10 @@ public Builder<Union, Kind, Member> addMember(Kind tag, JsonpDeserializer<? exte
185252
// also add it as a string
186253
addMember(Event.VALUE_STRING, tag, member);
187254
}
255+
} else if (unwrapped instanceof JsonpDeserializerBase.ArrayDeserializer) {
256+
// All arrays produce START_ARRAY, so they can't be keyed directly in `otherMembers`.
257+
// They are disambiguated later by the JSON event of their first element (see build()).
258+
arrayMembers.add(new SingleMemberHandler<>(tag, deserializer));
188259
} else {
189260
SingleMemberHandler<Union, Kind, Member> member = new SingleMemberHandler<>(tag, deserializer);
190261
for (Event e: deserializer.nativeEvents()) {
@@ -204,15 +275,23 @@ public JsonpDeserializer<Union> build() {
204275
}
205276
}
206277

278+
if (arrayMembers.size() == 1) {
279+
// A single array member can be keyed directly on START_ARRAY, no disambiguation needed
280+
addMember(Event.START_ARRAY, arrayMembers.get(0).tag, arrayMembers.get(0));
281+
} else if (arrayMembers.size() > 1) {
282+
if (otherMembers.containsKey(Event.START_ARRAY)) {
283+
throw new AmbiguousUnionException(
284+
"Array member '" + arrayMembers.get(0).tag + "' conflicts with another START_ARRAY member");
285+
}
286+
// Multiple array members are disambiguated by looking ahead at their first element
287+
otherMembers.put(Event.START_ARRAY, new ArrayMemberHandler<>(arrayMembers));
288+
}
289+
207290
if (objectMembers.size() == 1 && !otherMembers.containsKey(Event.START_OBJECT)) {
208291
// A single deserializer handles objects: promote it to otherMembers as we don't need property-based disambiguation
209292
otherMembers.put(Event.START_OBJECT, objectMembers.remove(0));
210293
}
211294

212-
// if (objectMembers.size() > 1) {
213-
// System.out.println("multiple objects in " + buildFn);
214-
// }
215-
216295
return new UnionDeserializer<>(objectMembers, otherMembers, buildFn);
217296
}
218297
}

java-client/src/test/java/co/elastic/clients/json/WithJsonTest.java

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
import co.elastic.clients.elasticsearch.core.IndexRequest;
3030
import co.elastic.clients.elasticsearch.core.SearchResponse;
3131
import co.elastic.clients.elasticsearch.indices.PutIndicesSettingsRequest;
32+
import co.elastic.clients.elasticsearch.inference.RerankRequest;
3233
import co.elastic.clients.elasticsearch.security.RoleTemplateScript;
3334
import co.elastic.clients.testkit.ModelTestCase;
3435
import org.junit.jupiter.api.Test;
@@ -250,5 +251,40 @@ public void testBooleanEnum() {
250251
// Asserting that both gets deserialized in the same way
251252
assertEquals(respBool.toString(), respString.toString());
252253
}
254+
255+
@Test
256+
public void testArrayUnion() {
257+
258+
String stringForm = """
259+
{
260+
"input": ["luke", "like", "leia", "chewy","r2d2", "star", "wars"],
261+
"query": "star wars main character"
262+
}
263+
""";
264+
265+
String objectForm = """
266+
{
267+
"input": [
268+
{
269+
"type": "text",
270+
"format": "text",
271+
"value": "some document text"
272+
},
273+
{
274+
"type": "image",
275+
"format": "base64",
276+
"value": "data:image/jpeg;base64,..."
277+
}
278+
],
279+
"query": "star wars main character"
280+
}
281+
""";
282+
283+
RerankRequest r1 = RerankRequest.of(r -> r.inferenceId("cohere_rerank").withJson(new StringReader(stringForm)));
284+
RerankRequest r2 = RerankRequest.of(r -> r.inferenceId("cohere_rerank").withJson(new StringReader(objectForm)));
285+
286+
assertTrue(r1.input().isString());
287+
assertTrue(r2.input().isObject());
288+
}
253289
}
254290

0 commit comments

Comments
 (0)