Skip to content

Commit 7fd825a

Browse files
committed
feat(core): graceful uri <-> urn handling
Implement comprehensive fallback strategy to resolve extensions from either URN or legacy URI references with conflict detection for protobuf ambiguity. Introduce round-trip tests to ensure output always has uri + urn regardless of combination of uri or urn in input plans. BREAKING CHANGE: This PR alters the extension loading API to require both URI explicitly, and URN implicitly in the plan. Extensions which lack URN will throw an error, and parsed plans which contain a uri/urn without a matching urn/uri pre-loaded into the context will throw an error.
1 parent c7639b3 commit 7fd825a

30 files changed

Lines changed: 2539 additions & 66 deletions

core/build.gradle.kts

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,8 @@ configurations[JavaPlugin.TEST_IMPLEMENTATION_CONFIGURATION_NAME].extendsFrom(sh
108108

109109
dependencies {
110110
testImplementation(platform(libs.junit.bom))
111+
testImplementation(libs.protobuf.java.util)
112+
111113
testImplementation(libs.junit.jupiter)
112114
testRuntimeOnly(libs.junit.platform.launcher)
113115

core/src/main/java/io/substrait/extendedexpression/ProtoExtendedExpressionConverter.java

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,14 +28,17 @@ public ProtoExtendedExpressionConverter() {
2828
}
2929

3030
public ProtoExtendedExpressionConverter(SimpleExtension.ExtensionCollection extensionCollection) {
31+
if (extensionCollection == null) {
32+
throw new IllegalArgumentException("ExtensionCollection is required");
33+
}
3134
this.extensionCollection = extensionCollection;
3235
}
3336

3437
public ExtendedExpression from(io.substrait.proto.ExtendedExpression extendedExpression) {
3538
// fill in simple extension information through a discovery in the current proto-extended
3639
// expression
3740
ExtensionLookup functionLookup =
38-
ImmutableExtensionLookup.builder().from(extendedExpression).build();
41+
ImmutableExtensionLookup.builder(extensionCollection).from(extendedExpression).build();
3942

4043
NamedStruct baseSchemaProto = extendedExpression.getBaseSchema();
4144

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

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,26 @@ public AbstractExtensionLookup(
1313
this.typeAnchorMap = typeAnchorMap;
1414
}
1515

16+
/**
17+
* Gets the function anchor for a given reference (primarily for testing).
18+
*
19+
* @param reference The function reference
20+
* @return The function anchor, or null if not found
21+
*/
22+
public SimpleExtension.FunctionAnchor getFunctionAnchor(int reference) {
23+
return functionAnchorMap.get(reference);
24+
}
25+
26+
/**
27+
* Gets the type anchor for a given reference (primarily for testing).
28+
*
29+
* @param reference The type reference
30+
* @return The type anchor, or null if not found
31+
*/
32+
public SimpleExtension.TypeAnchor getTypeAnchor(int reference) {
33+
return typeAnchorMap.get(reference);
34+
}
35+
1636
@Override
1737
public SimpleExtension.ScalarFunctionVariant getScalarFunction(
1838
int reference, SimpleExtension.ExtensionCollection extensions) {

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import java.util.Set;
66

77
/** We don't depend on guava... */
8-
class BidiMap<T1, T2> {
8+
public class BidiMap<T1, T2> {
99
private final Map<T1, T2> forwardMap;
1010
private final Map<T2, T1> reverseMap;
1111

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

Lines changed: 87 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import io.substrait.proto.ExtendedExpression;
44
import io.substrait.proto.Plan;
55
import io.substrait.proto.SimpleExtensionDeclaration;
6+
import io.substrait.proto.SimpleExtensionURI;
67
import io.substrait.proto.SimpleExtensionURN;
78
import java.util.ArrayList;
89
import java.util.HashMap;
@@ -19,14 +20,27 @@
1920
public class ExtensionCollector extends AbstractExtensionLookup {
2021
private final BidiMap<Integer, SimpleExtension.FunctionAnchor> funcMap;
2122
private final BidiMap<Integer, SimpleExtension.TypeAnchor> typeMap;
23+
private final SimpleExtension.ExtensionCollection extensionCollection;
2224

2325
// start at 0 to make sure functionAnchors start with 1 according to spec
2426
private int counter = 0;
2527

28+
private String getUriFromUrn(String urn) {
29+
return extensionCollection.getUriFromUrn(urn);
30+
}
31+
2632
public ExtensionCollector() {
33+
this(SimpleExtension.loadDefaults());
34+
}
35+
36+
public ExtensionCollector(SimpleExtension.ExtensionCollection extensionCollection) {
2737
super(new HashMap<>(), new HashMap<>());
38+
if (extensionCollection == null) {
39+
throw new IllegalArgumentException("ExtensionCollection is required");
40+
}
2841
funcMap = new BidiMap<>(functionAnchorMap);
2942
typeMap = new BidiMap<>(typeAnchorMap);
43+
this.extensionCollection = extensionCollection;
3044
}
3145

3246
public int getFunctionReference(SimpleExtension.Function declaration) {
@@ -53,70 +67,124 @@ public void addExtensionsToPlan(Plan.Builder builder) {
5367
SimpleExtensions simpleExtensions = getExtensions();
5468

5569
builder.addAllExtensionUrns(simpleExtensions.urns.values());
70+
builder.addAllExtensionUris(simpleExtensions.uris.values());
5671
builder.addAllExtensions(simpleExtensions.extensionList);
5772
}
5873

5974
public void addExtensionsToExtendedExpression(ExtendedExpression.Builder builder) {
6075
SimpleExtensions simpleExtensions = getExtensions();
6176

6277
builder.addAllExtensionUrns(simpleExtensions.urns.values());
78+
builder.addAllExtensionUris(simpleExtensions.uris.values());
6379
builder.addAllExtensions(simpleExtensions.extensionList);
6480
}
6581

6682
private SimpleExtensions getExtensions() {
6783
AtomicInteger urnPos = new AtomicInteger(1);
84+
AtomicInteger uriPos = new AtomicInteger(1);
6885
HashMap<String, SimpleExtensionURN> urns = new HashMap<>();
86+
HashMap<String, SimpleExtensionURI> uris = new HashMap<>();
6987

7088
ArrayList<SimpleExtensionDeclaration> extensionList = new ArrayList<>();
7189
for (Map.Entry<Integer, SimpleExtension.FunctionAnchor> e : funcMap.forwardEntrySet()) {
72-
SimpleExtensionURN urn =
90+
String urn = e.getValue().urn();
91+
String uri = getUriFromUrn(urn);
92+
93+
// Create URN entry
94+
SimpleExtensionURN urnObj =
7395
urns.computeIfAbsent(
74-
e.getValue().urn(),
96+
urn,
7597
k ->
7698
SimpleExtensionURN.newBuilder()
7799
.setExtensionUrnAnchor(urnPos.getAndIncrement())
78100
.setUrn(k)
79101
.build());
102+
103+
// Create URI entry if mapping exists
104+
SimpleExtensionURI uriObj = null;
105+
if (uri != null) {
106+
uriObj =
107+
uris.computeIfAbsent(
108+
uri,
109+
k ->
110+
SimpleExtensionURI.newBuilder()
111+
.setExtensionUriAnchor(uriPos.getAndIncrement())
112+
.setUri(k)
113+
.build());
114+
}
115+
116+
// Create function declaration with both URN and URI references
117+
SimpleExtensionDeclaration.ExtensionFunction.Builder funcBuilder =
118+
SimpleExtensionDeclaration.ExtensionFunction.newBuilder()
119+
.setFunctionAnchor(e.getKey())
120+
.setName(e.getValue().key())
121+
.setExtensionUrnReference(urnObj.getExtensionUrnAnchor());
122+
123+
if (uriObj != null) {
124+
funcBuilder.setExtensionUriReference(uriObj.getExtensionUriAnchor());
125+
}
126+
80127
SimpleExtensionDeclaration decl =
81-
SimpleExtensionDeclaration.newBuilder()
82-
.setExtensionFunction(
83-
SimpleExtensionDeclaration.ExtensionFunction.newBuilder()
84-
.setFunctionAnchor(e.getKey())
85-
.setName(e.getValue().key())
86-
.setExtensionUrnReference(urn.getExtensionUrnAnchor()))
87-
.build();
128+
SimpleExtensionDeclaration.newBuilder().setExtensionFunction(funcBuilder).build();
88129
extensionList.add(decl);
89130
}
131+
90132
for (Map.Entry<Integer, SimpleExtension.TypeAnchor> e : typeMap.forwardEntrySet()) {
91-
SimpleExtensionURN urn =
133+
String urn = e.getValue().urn();
134+
String uri = getUriFromUrn(urn);
135+
136+
// Create URN entry
137+
SimpleExtensionURN urnObj =
92138
urns.computeIfAbsent(
93-
e.getValue().urn(),
139+
urn,
94140
k ->
95141
SimpleExtensionURN.newBuilder()
96142
.setExtensionUrnAnchor(urnPos.getAndIncrement())
97143
.setUrn(k)
98144
.build());
145+
146+
// Create URI entry if mapping exists
147+
SimpleExtensionURI uriObj = null;
148+
if (uri != null) {
149+
uriObj =
150+
uris.computeIfAbsent(
151+
uri,
152+
k ->
153+
SimpleExtensionURI.newBuilder()
154+
.setExtensionUriAnchor(uriPos.getAndIncrement())
155+
.setUri(k)
156+
.build());
157+
}
158+
159+
// Create type declaration with both URN and URI references
160+
SimpleExtensionDeclaration.ExtensionType.Builder typeBuilder =
161+
SimpleExtensionDeclaration.ExtensionType.newBuilder()
162+
.setTypeAnchor(e.getKey())
163+
.setName(e.getValue().key())
164+
.setExtensionUrnReference(urnObj.getExtensionUrnAnchor());
165+
166+
if (uriObj != null) {
167+
typeBuilder.setExtensionUriReference(uriObj.getExtensionUriAnchor());
168+
}
169+
99170
SimpleExtensionDeclaration decl =
100-
SimpleExtensionDeclaration.newBuilder()
101-
.setExtensionType(
102-
SimpleExtensionDeclaration.ExtensionType.newBuilder()
103-
.setTypeAnchor(e.getKey())
104-
.setName(e.getValue().key())
105-
.setExtensionUrnReference(urn.getExtensionUrnAnchor()))
106-
.build();
171+
SimpleExtensionDeclaration.newBuilder().setExtensionType(typeBuilder).build();
107172
extensionList.add(decl);
108173
}
109-
return new SimpleExtensions(urns, extensionList);
174+
return new SimpleExtensions(urns, uris, extensionList);
110175
}
111176

112177
private static final class SimpleExtensions {
113178
final HashMap<String, SimpleExtensionURN> urns;
179+
final HashMap<String, SimpleExtensionURI> uris;
114180
final ArrayList<SimpleExtensionDeclaration> extensionList;
115181

116182
SimpleExtensions(
117183
HashMap<String, SimpleExtensionURN> urns,
184+
HashMap<String, SimpleExtensionURI> uris,
118185
ArrayList<SimpleExtensionDeclaration> extensionList) {
119186
this.urns = urns;
187+
this.uris = uris;
120188
this.extensionList = extensionList;
121189
}
122190
}

0 commit comments

Comments
 (0)