Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -897,25 +897,98 @@ private static MultiValueExpandSpec explicitMultiValueExpand(Correlate correlate
);
}

/**
* Wire-level description of a multi-value expand. {@code outputType} is the Calcite-derived
* record type when the spec is built from a {@link RelNode}; it is {@code null} when the spec
* was decoded from serialized bytes, in which case the record type is derived from the input.
*/
private record MultiValueExpandSpec(int fieldIndex, Integer limit, boolean append, boolean distinct, Type.Struct outputType) {
}

private static final class MultiValueExpandDetail implements Extension.SingleRelDetail {
private static final String TYPE_URL = "opensearch://analytics/multi_value_expand/v1";
/**
* {@link ExtensionSingle} detail for the multi-value expand, mirrored by
* {@code substrait_consumer.rs} on the Rust side. The 16-byte payload carries
* {@code fieldIndex, limit, append, distinct} as big-endian i32s; the output schema is not on
* the wire and is re-derived from the input on decode ({@link #deriveRecordType(Rel)}).
*/
static final class MultiValueExpandDetail implements Extension.SingleRelDetail {
static final String TYPE_URL = "opensearch://analytics/multi_value_expand/v1";
private static final int PAYLOAD_BYTES = 16;
private final MultiValueExpandSpec spec;

private MultiValueExpandDetail(MultiValueExpandSpec spec) {
this.spec = spec;
}

/** Whether {@code any} carries a multi-value expand payload this class can decode. */
static boolean matches(Any any) {
return any != null && TYPE_URL.equals(any.getTypeUrl());
}

/** Inverse of {@link #toProto}; the record type is derived lazily from the input. */
static MultiValueExpandDetail fromProto(Any any) {
if (!matches(any)) {
throw new IllegalArgumentException("Not a multi-value expand extension: " + (any == null ? null : any.getTypeUrl()));
}
ByteString value = any.getValue();
if (value.size() != PAYLOAD_BYTES) {
throw new IllegalArgumentException(
"Malformed multi-value expand payload: expected " + PAYLOAD_BYTES + " bytes, got " + value.size()
);
}
ByteBuffer payload = value.asReadOnlyByteBuffer();
int fieldIndex = payload.getInt();
int limit = payload.getInt();
int append = payload.getInt();
int distinct = payload.getInt();
if (fieldIndex < 0 || limit < -1 || (append != 0 && append != 1) || (distinct != 0 && distinct != 1)) {
throw new IllegalArgumentException(
"Malformed multi-value expand payload: fieldIndex="
+ fieldIndex
+ " limit="
+ limit
+ " append="
+ append
+ " distinct="
+ distinct
);
}
return new MultiValueExpandDetail(
new MultiValueExpandSpec(fieldIndex, limit < 0 ? null : limit, append == 1, distinct == 1, null)
);
}

@Override
public Type.Struct deriveRecordType(Rel input) {
return spec.outputType();
if (spec.outputType() != null) {
return spec.outputType();
}
// Decoded from the wire: rebuild the same shape substrait_consumer.rs produces. The
// expanded column is the LIST's element type, nullable (unnest preserves nulls); in
// append mode it is added after the input columns, otherwise it replaces the LIST column.
List<Type> inputFields = input.getRecordType().fields();
if (spec.fieldIndex() >= inputFields.size()) {
throw new IllegalArgumentException(
"multi-value expand field index " + spec.fieldIndex() + " is outside " + inputFields.size() + " input columns"
);
}
Type source = inputFields.get(spec.fieldIndex());
if (!(source instanceof Type.ListType list)) {
throw new IllegalArgumentException("multi-value expand field " + spec.fieldIndex() + " is not a LIST: " + source);
}
Type element = list.elementType().withNullable(true);
List<Type> outputFields = new ArrayList<>(inputFields);
if (spec.append()) {
outputFields.add(element);
} else {
outputFields.set(spec.fieldIndex(), element);
}
return Type.Struct.builder().nullable(input.getRecordType().nullable()).addAllFields(outputFields).build();
}

@Override
public Any toProto(io.substrait.relation.RelProtoConverter converter) {
ByteBuffer payload = ByteBuffer.allocate(16);
ByteBuffer payload = ByteBuffer.allocate(PAYLOAD_BYTES);
payload.putInt(spec.fieldIndex());
payload.putInt(spec.limit() == null ? -1 : spec.limit());
payload.putInt(spec.append() ? 1 : 0);
Expand All @@ -924,6 +997,32 @@ public Any toProto(io.substrait.relation.RelProtoConverter converter) {
}
}

/**
* {@link ProtoPlanConverter} that understands this backend's own {@link ExtensionSingle}
* details. The stock converter maps every unknown extension to an {@code EmptyDetail} whose
* record type is an empty struct, so any field reference above it (e.g. the GROUP BY key of a
* PARTIAL aggregate over the expanded column) fails with
* "Field reference offset (N) must be less than number of fields in struct (0)" when a
* shard-stage plan is decoded on the coordinator to splice the reduce fragment on top.
*/
private static final class OpenSearchProtoPlanConverter extends ProtoPlanConverter {
OpenSearchProtoPlanConverter(SimpleExtension.ExtensionCollection extensions) {
super(extensions);
}

@Override
protected io.substrait.relation.ProtoRelConverter getProtoRelConverter(io.substrait.extension.ExtensionLookup functionLookup) {
return new io.substrait.relation.ProtoRelConverter(functionLookup, extensionCollection, protoExtensionConverter) {
@Override
protected Extension.SingleRelDetail detailFromExtensionSingleRel(Any any) {
return MultiValueExpandDetail.matches(any)
? MultiValueExpandDetail.fromProto(any)
: super.detailFromExtensionSingleRel(any);
}
};
}
}

/**
* Adds an {@code is_not_null} {@code preMeasureFilter} to each measure whose {@link LocalAggOp}
* declares {@link LocalAggOp#filtersNullArgs} — so the converter stays generic and only the op
Expand Down Expand Up @@ -1102,11 +1201,11 @@ public boolean filtersNullArgs(AggregateCall call) {

// ── Plan serde helpers ──────────────────────────────────────────────────────

/** Decodes serialized Substrait bytes into a model-level {@link Plan}. */
/** Decodes serialized Substrait bytes into a model-level {@link Plan}, including this backend's own extension rels. */
private Plan decodePlan(byte[] bytes) {
try {
io.substrait.proto.Plan proto = io.substrait.proto.Plan.parseFrom(bytes);
return new ProtoPlanConverter(extensions).from(proto);
return new OpenSearchProtoPlanConverter(extensions).from(proto);
} catch (InvalidProtocolBufferException e) {
throw new IllegalArgumentException("Failed to decode Substrait plan bytes", e);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,12 @@
import org.apache.calcite.rel.RelNode;
import org.apache.calcite.rel.core.Aggregate;
import org.apache.calcite.rel.logical.LogicalProject;
import org.apache.calcite.rel.type.RelDataType;
import org.apache.calcite.rel.type.RelDataTypeFactory;
import org.apache.calcite.rel.type.RelDataTypeField;
import org.apache.calcite.rex.RexNode;
import org.apache.calcite.util.ImmutableBitSet;
import org.opensearch.analytics.planner.rel.OpenSearchAggregate;

import java.util.ArrayList;
import java.util.HashMap;
Expand All @@ -25,20 +29,21 @@
*
* <p>This rewriter appends a scalar unnested column per LIST GROUP BY key and remaps
* the grouping set indices so that Substrait serialisation sees scalar types throughout.
* It currently runs at Substrait-conversion time (inside
* It runs at Substrait-conversion time (inside
* {@code DataFusionFragmentConvertor.preprocessForSubstrait}), which is <em>after</em>
* {@code PlannerImpl.decomposeAggregates} has already split the aggregate into
* PARTIAL/FINAL halves. In a multi-shard query the FINAL fragment therefore receives
* a {@code List(Utf8)} GROUP BY key where Calcite expects the element type
* ({@code Utf8View}), causing a type mismatch 500.
* PARTIAL/FINAL halves. The PARTIAL half is expanded here; the FINAL half reads the
* already-expanded rows from its stage input, so its LIST keys are retyped to the element
* type instead of being expanded again (see {@link #retypeExpandedStageInputKeys}).
*
* <p><b>Known issue</b>: the correct fix is to move this expansion into
* <p><b>Known issue</b>: the cleaner fix is to move this expansion into
* {@code PlannerImpl.runAllOptimizations} <em>before</em> the
* {@code decomposeAggregates} call so that Calcite propagates element types into both
* the PARTIAL and FINAL fragments. That requires relocating
* {@link MultiValueExpandRel} into a module that {@code analytics-engine} can depend
* on (the dependency currently flows the other way: {@code analytics-backend-datafusion}
* extends {@code analytics-engine}). This is tracked for a follow-up PR.
* the PARTIAL and FINAL fragments and the FINAL-side retyping becomes unnecessary. That
* requires relocating {@link MultiValueExpandRel} into a module that {@code analytics-engine}
* can depend on (the dependency currently flows the other way:
* {@code analytics-backend-datafusion} extends {@code analytics-engine}). This is tracked
* for a follow-up PR.
*/
final class MultiValueRelRewriter {

Expand All @@ -56,6 +61,9 @@ public RelNode visit(RelNode other) {

private static RelNode rewriteAggregate(Aggregate aggregate) {
RelNode input = aggregate.getInput();
if (OpenSearchAggregate.isFinalMode(aggregate) && input instanceof DataFusionFragmentConvertor.StageInputTableScan stageInput) {
return retypeExpandedStageInputKeys(aggregate, stageInput);
}
Map<Integer, Integer> expandedGroupFields = new HashMap<>();
for (int fieldIndex : aggregate.getGroupSet()) {
if (input.getRowType().getFieldList().get(fieldIndex).getType().getComponentType() != null) {
Expand Down Expand Up @@ -88,6 +96,50 @@ private static RelNode rewriteAggregate(Aggregate aggregate) {
return LogicalProject.create(rewritten, List.of(), projects, aggregate.getRowType().getFieldNames());
}

/**
* FINAL-side counterpart of the expansion above. The PARTIAL fragment already replaced each
* LIST GROUP BY key with its expanded scalar element (see {@link #rewriteAggregate}), so the
* rows arriving on the reduce stage's input partition carry the element type. Calcite,
* however, still types the stage input from the pre-expansion aggregate row type, so the
* FINAL aggregate would (a) expand an already-scalar column a second time and (b) declare
* the partition's Substrait {@code base_schema} as {@code List(Utf8)} where the registered
* stream is {@code Utf8View}, which DataFusion rejects at the ReadRel. Retype the affected
* keys on the stage input to the nullable element type instead; the aggregate's row type
* re-derives from the input, so no expansion or reordering Project is needed.
*/
private static RelNode retypeExpandedStageInputKeys(Aggregate aggregate, DataFusionFragmentConvertor.StageInputTableScan stageInput) {
RelDataTypeFactory typeFactory = aggregate.getCluster().getTypeFactory();
List<RelDataTypeField> fields = stageInput.getRowType().getFieldList();
RelDataTypeFactory.Builder builder = typeFactory.builder();
boolean changed = false;
for (int fieldIndex = 0; fieldIndex < fields.size(); fieldIndex++) {
RelDataTypeField field = fields.get(fieldIndex);
RelDataType elementType = field.getType().getComponentType();
if (elementType != null && aggregate.getGroupSet().get(fieldIndex)) {
builder.add(field.getName(), typeFactory.createTypeWithNullability(elementType, true));
changed = true;
} else {
builder.add(field.getName(), field.getType());
}
}
if (!changed) {
return aggregate;
}
RelNode retyped = new DataFusionFragmentConvertor.StageInputTableScan(
stageInput.getCluster(),
stageInput.getTraitSet(),
stageInput.getTable().getQualifiedName().getFirst(),
builder.build()
);
return aggregate.copy(
aggregate.getTraitSet(),
retyped,
aggregate.getGroupSet(),
aggregate.getGroupSets(),
aggregate.getAggCallList()
);
}

private static ImmutableBitSet remap(ImmutableBitSet fields, Map<Integer, Integer> replacements) {
ImmutableBitSet.Builder builder = ImmutableBitSet.builder();
for (int field : fields) {
Expand Down
Loading
Loading