Skip to content

Commit b865296

Browse files
authored
feat(core): expose metadata from YAML extension files (#691)
1 parent 824fc58 commit b865296

3 files changed

Lines changed: 192 additions & 6 deletions

File tree

core/src/main/java/io/substrait/extension/SimpleExtension.java

Lines changed: 51 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,8 @@
2222
import java.io.IOException;
2323
import java.io.InputStream;
2424
import java.io.UncheckedIOException;
25+
import java.util.Collections;
26+
import java.util.HashMap;
2527
import java.util.List;
2628
import java.util.Map;
2729
import java.util.Optional;
@@ -276,6 +278,8 @@ public String description() {
276278

277279
public abstract Map<String, Option> options();
278280

281+
public abstract Optional<Map<String, Object>> metadata();
282+
279283
public List<Argument> requiredArguments() {
280284
return requiredArgsSupplier.get();
281285
}
@@ -381,25 +385,29 @@ public abstract static class ScalarFunction {
381385
@Nullable
382386
public abstract String description();
383387

388+
public abstract Optional<Map<String, Object>> metadata();
389+
384390
public abstract List<ScalarFunctionVariant> impls();
385391

386392
public Stream<ScalarFunctionVariant> resolve(String urn) {
387-
return impls().stream().map(f -> f.resolve(urn, name(), description()));
393+
return impls().stream().map(f -> f.resolve(urn, name(), description(), metadata()));
388394
}
389395
}
390396

391397
@JsonDeserialize(as = ImmutableSimpleExtension.ScalarFunctionVariant.class)
392398
@JsonSerialize(as = ImmutableSimpleExtension.ScalarFunctionVariant.class)
393399
@Value.Immutable
394400
public abstract static class ScalarFunctionVariant extends Function {
395-
public ScalarFunctionVariant resolve(String urn, String name, String description) {
401+
public ScalarFunctionVariant resolve(
402+
String urn, String name, String description, Optional<Map<String, Object>> metadata) {
396403
return ImmutableSimpleExtension.ScalarFunctionVariant.builder()
397404
.urn(urn)
398405
.name(name)
399406
.description(description)
400407
.nullability(nullability())
401408
.args(args())
402409
.options(options())
410+
.metadata(metadata)
403411
.ordered(ordered())
404412
.variadic(variadic())
405413
.returnType(returnType())
@@ -417,10 +425,12 @@ public abstract static class AggregateFunction {
417425
@Nullable
418426
public abstract String description();
419427

428+
public abstract Optional<Map<String, Object>> metadata();
429+
420430
public abstract List<AggregateFunctionVariant> impls();
421431

422432
public Stream<AggregateFunctionVariant> resolve(String urn) {
423-
return impls().stream().map(f -> f.resolve(urn, name(), description()));
433+
return impls().stream().map(f -> f.resolve(urn, name(), description(), metadata()));
424434
}
425435
}
426436

@@ -434,10 +444,12 @@ public abstract static class WindowFunction {
434444
@Nullable
435445
public abstract String description();
436446

447+
public abstract Optional<Map<String, Object>> metadata();
448+
437449
public abstract List<WindowFunctionVariant> impls();
438450

439451
public Stream<WindowFunctionVariant> resolve(String urn) {
440-
return impls().stream().map(f -> f.resolve(urn, name(), description()));
452+
return impls().stream().map(f -> f.resolve(urn, name(), description(), metadata()));
441453
}
442454

443455
public static ImmutableSimpleExtension.WindowFunction.Builder builder() {
@@ -463,14 +475,16 @@ public String toString() {
463475
@Nullable
464476
public abstract TypeExpression intermediate();
465477

466-
AggregateFunctionVariant resolve(String urn, String name, String description) {
478+
AggregateFunctionVariant resolve(
479+
String urn, String name, String description, Optional<Map<String, Object>> metadata) {
467480
return ImmutableSimpleExtension.AggregateFunctionVariant.builder()
468481
.urn(urn)
469482
.name(name)
470483
.description(description)
471484
.nullability(nullability())
472485
.args(args())
473486
.options(options())
487+
.metadata(metadata)
474488
.ordered(ordered())
475489
.variadic(variadic())
476490
.decomposability(decomposability())
@@ -505,14 +519,16 @@ public String toString() {
505519
return super.toString();
506520
}
507521

508-
WindowFunctionVariant resolve(String urn, String name, String description) {
522+
WindowFunctionVariant resolve(
523+
String urn, String name, String description, Optional<Map<String, Object>> metadata) {
509524
return ImmutableSimpleExtension.WindowFunctionVariant.builder()
510525
.urn(urn)
511526
.name(name)
512527
.description(description)
513528
.nullability(nullability())
514529
.args(args())
515530
.options(options())
531+
.metadata(metadata)
516532
.ordered(ordered())
517533
.variadic(variadic())
518534
.decomposability(decomposability())
@@ -549,6 +565,8 @@ public abstract static class Type {
549565

550566
protected abstract Optional<Boolean> variadic();
551567

568+
public abstract Optional<Map<String, Object>> metadata();
569+
552570
public TypeAnchor getAnchor() {
553571
return anchorSupplier.get();
554572
}
@@ -574,6 +592,9 @@ public abstract static class ExtensionSignatures {
574592
@JsonProperty("window_functions")
575593
public abstract List<WindowFunction> windows();
576594

595+
@JsonProperty("metadata")
596+
public abstract Optional<Map<String, Object>> metadata();
597+
577598
public int size() {
578599
return (types() == null ? 0 : types().size())
579600
+ (scalars() == null ? 0 : scalars().size())
@@ -643,6 +664,11 @@ BidiMap<String, String> uriUrnMap() {
643664
return new BidiMap<>();
644665
}
645666

667+
@Value.Default
668+
public Map<String, Map<String, Object>> extensionMetadata() {
669+
return Collections.emptyMap();
670+
}
671+
646672
public abstract List<Type> types();
647673

648674
public abstract List<ScalarFunctionVariant> scalarFunctions();
@@ -655,6 +681,16 @@ public static ImmutableSimpleExtension.ExtensionCollection.Builder builder() {
655681
return ImmutableSimpleExtension.ExtensionCollection.builder();
656682
}
657683

684+
/**
685+
* Gets the top-level metadata for a specific extension by URN.
686+
*
687+
* @param urn The URN of the extension
688+
* @return The metadata map if present, empty Optional otherwise
689+
*/
690+
public Optional<Map<String, Object>> getExtensionMetadata(String urn) {
691+
return Optional.ofNullable(extensionMetadata().get(urn));
692+
}
693+
658694
public Type getType(TypeAnchor anchor) {
659695
Type type = typeLookup.get().get(anchor);
660696
if (type != null) {
@@ -744,6 +780,10 @@ public ExtensionCollection merge(ExtensionCollection extensionCollection) {
744780
mergedUriUrnMap.merge(uriUrnMap());
745781
mergedUriUrnMap.merge(extensionCollection.uriUrnMap());
746782

783+
Map<String, Map<String, Object>> mergedExtensionMetadata = new HashMap<>();
784+
mergedExtensionMetadata.putAll(extensionMetadata());
785+
mergedExtensionMetadata.putAll(extensionCollection.extensionMetadata());
786+
747787
return ImmutableSimpleExtension.ExtensionCollection.builder()
748788
.addAllAggregateFunctions(aggregateFunctions())
749789
.addAllAggregateFunctions(extensionCollection.aggregateFunctions())
@@ -754,6 +794,7 @@ public ExtensionCollection merge(ExtensionCollection extensionCollection) {
754794
.addAllTypes(types())
755795
.addAllTypes(extensionCollection.types())
756796
.uriUrnMap(mergedUriUrnMap)
797+
.extensionMetadata(mergedExtensionMetadata)
757798
.build();
758799
}
759800
}
@@ -859,13 +900,17 @@ public static ExtensionCollection buildExtensionCollection(
859900
BidiMap<String, String> uriUrnMap = new BidiMap<>();
860901
uriUrnMap.put(uri, urn);
861902

903+
Map<String, Map<String, Object>> extMetadata = new HashMap<>();
904+
extensionSignatures.metadata().ifPresent(m -> extMetadata.put(urn, m));
905+
862906
ImmutableSimpleExtension.ExtensionCollection collection =
863907
ImmutableSimpleExtension.ExtensionCollection.builder()
864908
.scalarFunctions(scalarFunctionVariants)
865909
.aggregateFunctions(aggregateFunctionVariants)
866910
.windowFunctions(allWindowFunctionVariants)
867911
.addAllTypes(extensionSignatures.types())
868912
.uriUrnMap(uriUrnMap)
913+
.extensionMetadata(extMetadata)
869914
.build();
870915

871916
LOGGER.atDebug().log(
Lines changed: 102 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,102 @@
1+
package io.substrait.extension;
2+
3+
import static org.junit.jupiter.api.Assertions.assertEquals;
4+
import static org.junit.jupiter.api.Assertions.assertTrue;
5+
6+
import io.substrait.TestBase;
7+
import java.io.IOException;
8+
import java.io.UncheckedIOException;
9+
import java.util.Map;
10+
import org.junit.jupiter.api.Test;
11+
12+
/**
13+
* Verifies that metadata can be read from extension YAML files at multiple levels:
14+
*
15+
* <ul>
16+
* <li>Extension-level metadata (top-level)
17+
* <li>Type-level metadata
18+
* <li>Function-level metadata (scalar, aggregate, window)
19+
* </ul>
20+
*/
21+
class MetadataExtensionTest extends TestBase {
22+
23+
static final String URN = "extension:test:metadata_extensions";
24+
static final SimpleExtension.ExtensionCollection METADATA_EXTENSION;
25+
26+
static {
27+
try {
28+
String extensionStr = asString("extensions/metadata_extensions.yaml");
29+
METADATA_EXTENSION = SimpleExtension.load(URN, extensionStr);
30+
} catch (IOException e) {
31+
throw new UncheckedIOException(e);
32+
}
33+
}
34+
35+
MetadataExtensionTest() {
36+
super(METADATA_EXTENSION);
37+
}
38+
39+
@Test
40+
void testExtensionLevelMetadata() {
41+
Map<String, Object> metadata = extensions.getExtensionMetadata(URN).orElseThrow();
42+
assertEquals("1.0", metadata.get("version"));
43+
assertEquals("test-team", metadata.get("author"));
44+
45+
@SuppressWarnings("unchecked")
46+
Map<String, Object> customData = (Map<String, Object>) metadata.get("custom_data");
47+
assertEquals(true, customData.get("nested_value"));
48+
assertEquals(42, customData.get("numeric_value"));
49+
}
50+
51+
@Test
52+
void testExtensionLevelMetadataMissing() {
53+
assertTrue(extensions.getExtensionMetadata("extension:nonexistent:urn").isEmpty());
54+
}
55+
56+
@Test
57+
void testTypeMetadata() {
58+
SimpleExtension.TypeAnchor anchor = SimpleExtension.TypeAnchor.of(URN, "metadataType");
59+
Map<String, Object> metadata = extensions.getType(anchor).metadata().orElseThrow();
60+
assertEquals("custom-type-metadata", metadata.get("type_info"));
61+
assertEquals("user-defined", metadata.get("category"));
62+
}
63+
64+
@Test
65+
void testScalarFunctionMetadata() {
66+
SimpleExtension.FunctionAnchor anchor =
67+
SimpleExtension.FunctionAnchor.of(URN, "metadataScalar:i64");
68+
Map<String, Object> metadata = extensions.getScalarFunction(anchor).metadata().orElseThrow();
69+
assertEquals("vectorized", metadata.get("perf_hint"));
70+
assertEquals(1, metadata.get("cost"));
71+
}
72+
73+
@Test
74+
void testAggregateFunctionMetadata() {
75+
SimpleExtension.FunctionAnchor anchor =
76+
SimpleExtension.FunctionAnchor.of(URN, "metadataAggregate:i64");
77+
assertEquals(
78+
"incremental",
79+
extensions.getAggregateFunction(anchor).metadata().orElseThrow().get("agg_info"));
80+
}
81+
82+
@Test
83+
void testWindowFunctionMetadata() {
84+
SimpleExtension.FunctionAnchor anchor =
85+
SimpleExtension.FunctionAnchor.of(URN, "metadataWindow:i64");
86+
assertEquals(
87+
"partitioned",
88+
extensions.getWindowFunction(anchor).metadata().orElseThrow().get("window_info"));
89+
}
90+
91+
@Test
92+
void testMergePreservesMetadata() throws IOException {
93+
String customExtensionStr = asString("extensions/custom_extensions.yaml");
94+
SimpleExtension.ExtensionCollection customExtension =
95+
SimpleExtension.load("extension:test:custom_extensions", customExtensionStr);
96+
97+
SimpleExtension.ExtensionCollection merged = METADATA_EXTENSION.merge(customExtension);
98+
99+
assertEquals("1.0", merged.getExtensionMetadata(URN).orElseThrow().get("version"));
100+
assertTrue(merged.getExtensionMetadata("extension:test:custom_extensions").isEmpty());
101+
}
102+
}
Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
%YAML 1.2
2+
---
3+
urn: extension:test:metadata_extensions
4+
metadata:
5+
version: "1.0"
6+
author: "test-team"
7+
custom_data:
8+
nested_value: true
9+
numeric_value: 42
10+
types:
11+
- name: "metadataType"
12+
metadata:
13+
type_info: "custom-type-metadata"
14+
category: "user-defined"
15+
scalar_functions:
16+
- name: "metadataScalar"
17+
metadata:
18+
perf_hint: "vectorized"
19+
cost: 1
20+
impls:
21+
- args:
22+
- value: i64
23+
return: i64
24+
aggregate_functions:
25+
- name: "metadataAggregate"
26+
metadata:
27+
agg_info: "incremental"
28+
impls:
29+
- args:
30+
- value: i64
31+
return: i64
32+
window_functions:
33+
- name: "metadataWindow"
34+
metadata:
35+
window_info: "partitioned"
36+
impls:
37+
- args:
38+
- value: i64
39+
return: i64

0 commit comments

Comments
 (0)