Skip to content
Open
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
427 changes: 251 additions & 176 deletions src/main/java/org/rumbledb/compiler/InferTypeVisitor.java

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,9 @@ public static boolean canBeFunctionCoercedTo(SequenceType sourceType, SequenceTy
}

public static boolean canItemTypeBeFunctionCoercedTo(ItemType sourceItemType, ItemType targetItemType) {
if (sourceItemType.isUnionType()) {
return sourceItemType.allMemberTypesMatch(member -> canItemTypeBeFunctionCoercedTo(member, targetItemType));
}
if (!targetItemType.isFunctionItemType() || targetItemType.getSignature() == null) {
return false;
}
Expand Down
3 changes: 3 additions & 0 deletions src/main/java/org/rumbledb/types/FunctionItemType.java
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,9 @@ public void resolve(DynamicContext context, ExceptionMetadata metadata) {

@Override
public boolean isSubtypeOf(ItemType superType) {
if (superType.isUnionType()) {
return superType.getTypes().stream().anyMatch(this::isSubtypeOf);
}
if (this.equals(superType)
|| superType.equals(anyFunctionItem)
|| superType.equals(BuiltinTypesCatalogue.item)) {
Expand Down
5 changes: 5 additions & 0 deletions src/main/java/org/rumbledb/types/SequenceType.java
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,11 @@ public boolean isSubtypeOfOrCanBePromotedTo(SequenceType superType) {
if (this.itemType.equals(BuiltinTypesCatalogue.errorItem)) {
return !this.cardinality.allowsZero() || superType.cardinality.allowsZero();
}
if (this.itemType.isUnionType()) {
// A member can fit directly while another needs promotion or function coercion.
return this.itemType.allMemberTypesMatch(
member -> new SequenceType(member, this.cardinality).isSubtypeOfOrCanBePromotedTo(superType));
}
return this.cardinality.isSubtypeOf(superType.cardinality)
&& (this.itemType.isSubtypeOf(superType.itemType)
|| this.itemType.canBePromotedTo(superType.itemType)
Expand Down
20 changes: 8 additions & 12 deletions src/main/java/org/rumbledb/types/UnionItemType.java
Original file line number Diff line number Diff line change
Expand Up @@ -105,18 +105,14 @@ public boolean isNumeric() {

@Override
public boolean isStaticallyCastableAs(ItemType other) {
if (other.equals(this)) {
return true;
}
if (other.isNumeric()) {
return true;
}
for (ItemType member : this.types) {
if (other.isSubtypeOf(member)) {
return true;
}
}
return false;
return other.equals(this)
|| allMemberTypesMatch(member ->
member.isSubtypeOf(BuiltinTypesCatalogue.atomicItem) && member.isStaticallyCastableAs(other));
}

@Override
public boolean canBePromotedTo(ItemType other) {
return allMemberTypesMatch(member -> member.isSubtypeOf(other) || member.canBePromotedTo(other));
}

@Override
Expand Down
7 changes: 5 additions & 2 deletions src/test/java/iq/StaticTypeTests.java
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,11 @@
package iq;

import java.io.File;
import java.io.IOException;
import java.util.List;

import iq.base.SparkAnnotationsTestsBase;
import iq.base.TestFileDiscovery;

import org.rumbledb.config.RumbleConfiguration;

Expand All @@ -41,7 +44,7 @@ protected File testDirectory() {
}

@Override
protected boolean checkOutput() {
return false;
protected List<File> testFiles() throws IOException {
return TestFileDiscovery.files(testDirectory(), ".jq", ".xq");
}
}
3 changes: 2 additions & 1 deletion src/test/java/iq/base/AnnotationTestExecutor.java
Original file line number Diff line number Diff line change
Expand Up @@ -211,7 +211,8 @@ private static void checkExpectedOutput(
boolean checkOutput,
boolean applyUpdates,
int resultSizeCap) {
if (!checkOutput) {
// A fixture without an Output annotation only checks that the query runs.
if (!checkOutput || expectedOutput == null) {
if (applyUpdates && sequence.availableAsPUL()) {
sequence.applyPUL();
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,6 @@
import org.rumbledb.config.CompilationConfiguration;
import org.rumbledb.config.RumbleConfiguration;
import org.rumbledb.exceptions.IsStaticallyUnexpectedTypeException;
import org.rumbledb.exceptions.UnexpectedStaticTypeException;
import org.rumbledb.types.BuiltinTypesCatalogue;
import org.rumbledb.types.ItemType;
import org.rumbledb.types.SequenceCardinality;
Expand Down Expand Up @@ -171,35 +170,6 @@ void assertedMultipleValuesStillBecomeOnlyAnArrayInObjectFields() {
assertEquals(BuiltinTypesCatalogue.integerItem, type.getItemType().getArrayContentFacet());
}

@Test
void arithmeticPromotesEachUnionMember() {
SequenceType type = infer("for $x in (1, 2.5e0) return $x + 1", "jq");
assertEquals(
List.of(BuiltinTypesCatalogue.integerItem, BuiltinTypesCatalogue.doubleItem),
type.getItemType().getTypes());
}

@Test
void unionsAreComparableIfEveryMemberPairIs() {
String flags = "declare variable $c as xs:boolean external; declare variable $d as xs:boolean external; ";
assertEquals(
BuiltinTypesCatalogue.booleanItem,
infer(flags + "(if ($c) then \"a\" else xs:anyURI(\"b\")) eq \"a\"", "jq")
.getItemType());
assertThrows(
UnexpectedStaticTypeException.class,
() -> infer(flags + "(if ($c) then 1 else \"a\") eq (if ($d) then 2 else \"b\")", "jq"));
}

@Test
void unionHasAnEffectiveBooleanValueIfEveryMemberHasOne() {
String flag = "declare variable $c as xs:boolean external; ";
assertDoesNotThrow(() -> infer(flag + "if (if ($c) then 1 else \"a\") then 1 else 2", "jq"));
assertThrows(
UnexpectedStaticTypeException.class,
() -> infer(flag + "if (if ($c) then 1 else current-date()) then 1 else 2", "jq"));
}

@ParameterizedTest
@ValueSource(
strings = {
Expand Down
181 changes: 181 additions & 0 deletions src/test/java/org/rumbledb/compiler/UnionTypeInferenceTest.java
Original file line number Diff line number Diff line change
@@ -0,0 +1,181 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*
* Contributor acknowledgements are maintained in the CONTRIBUTORS file at the project root.
*/
package org.rumbledb.compiler;

import java.net.URI;
import java.util.ArrayList;
import java.util.List;
import java.util.stream.Stream;

import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
import org.junit.jupiter.params.provider.MethodSource;
import org.junit.jupiter.params.provider.ValueSource;

import static org.junit.jupiter.api.Assertions.*;

import org.rumbledb.api.Item;
import org.rumbledb.api.Rumble;
import org.rumbledb.bindings.ExternalBindings;
import org.rumbledb.config.CompilationConfiguration;
import org.rumbledb.config.RumbleConfiguration;
import org.rumbledb.exceptions.UnexpectedStaticTypeException;
import org.rumbledb.types.BuiltinTypesCatalogue;
import org.rumbledb.types.ItemType;
import org.rumbledb.types.ItemTypeFactory;
import org.rumbledb.types.SequenceCardinality;
import org.rumbledb.types.SequenceType;

class UnionTypeInferenceTest {
private static final RumbleConfiguration CONFIGURATION = RumbleConfiguration.builder()
.configureAnalysis(analysis -> analysis.enableStaticTyping(true))
.build();

private record Case(String query, List<ItemType> expectedTypes) {}

static Stream<Arguments> unionQueries() {
List<ItemType> floating = List.of(BuiltinTypesCatalogue.floatItem, BuiltinTypesCatalogue.doubleItem);
List<ItemType> integers = List.of(BuiltinTypesCatalogue.integerItem);
List<ItemType> strings = List.of(BuiltinTypesCatalogue.stringItem);
List<Case> cases = new ArrayList<>();
for (String operator : List.of("+", "-", "*", "div", "mod")) {
cases.add(new Case("for $x in (xs:float(1), xs:double(2)) return $x " + operator + " 1", floating));
}
cases.addAll(List.of(
new Case("for $x in (xs:float(1), xs:double(2)) for $y in (xs:decimal(3), 4) return $x + $y", floating),
new Case("for $x in (xs:float(1), xs:double(2)) return $x idiv 1", integers),
new Case(
"for $x in (xs:date(\"2000-01-01\"), xs:dateTime(\"2000-01-01T00:00:00\")) return $x + xs:dayTimeDuration(\"P1D\")",
List.of(BuiltinTypesCatalogue.dateItem, BuiltinTypesCatalogue.dateTimeItem)),
new Case(
"for $x in (xs:float(1), xs:decimal(2)) return math:pow($x, 2)",
List.of(BuiltinTypesCatalogue.doubleItem)),
new Case("for $x in (xs:anyURI(\"a\"), \"b\") return string-length($x)", integers),
new Case(
"for $x in (xs:anyURI(\"a\"), \"b\") return $x eq \"a\"",
List.of(BuiltinTypesCatalogue.booleanItem)),
new Case("for $x in (1, \"2\") return $x cast as xs:string", strings),
new Case("for $x in (xs:boolean(\"true\"), 1) return $x cast as xs:string", strings),
new Case("sum((xs:positiveInteger(1), xs:negativeInteger(-1)))", integers),
new Case(
"sum((xs:untypedAtomic(\"1\"), 2))",
List.of(BuiltinTypesCatalogue.integerItem, BuiltinTypesCatalogue.doubleItem)),
new Case("avg((xs:untypedAtomic(\"1\"), 2))", List.of(BuiltinTypesCatalogue.numericItem)),
new Case(
"min((xs:untypedAtomic(\"1\"), 2))",
List.of(BuiltinTypesCatalogue.integerItem, BuiltinTypesCatalogue.doubleItem)),
new Case(
"max((xs:untypedAtomic(\"1\"), 2))",
List.of(BuiltinTypesCatalogue.integerItem, BuiltinTypesCatalogue.doubleItem)),
new Case(
"min((xs:int(1), xs:float(2)))",
List.of(BuiltinTypesCatalogue.integerItem, BuiltinTypesCatalogue.floatItem)),
new Case(
"max((xs:anyURI(\"a\"), \"b\"))",
List.of(BuiltinTypesCatalogue.anyURIItem, BuiltinTypesCatalogue.stringItem)),
new Case(
"sum((1, 2)[xs:boolean(\"false\")], \"empty\")",
List.of(BuiltinTypesCatalogue.integerItem, BuiltinTypesCatalogue.stringItem))));
return cases.stream().flatMap(test -> Stream.of("jq", "xq").map(extension -> Arguments.of(test, extension)));
}

@ParameterizedTest
@MethodSource("unionQueries")
void inferredTypesCoverRuntimeResults(Case test, String extension) {
URI uri = URI.create("file:///union-inference." + extension);
SequenceType inferred = infer(test.query(), uri);
ItemType expected = ItemTypeFactory.createInferredUnionType(test.expectedTypes());
assertTrue(inferred.getItemType().isSubtypeOf(expected), () -> "Unexpected inferred type: " + inferred);
List<Item> values =
new Rumble(CONFIGURATION).runQuery(test.query(), uri).getAsList();
assertFalse(values.isEmpty());
for (Item value : values) {
assertTrue(
value.getDynamicType().isSubtypeOf(inferred.getItemType()),
() -> value.getDynamicType() + " is excluded by " + inferred + " for " + test.query());
}
}

@Test
void sumEmptyInputUsesItsZeroArgument() {
URI uri = URI.create("file:///sum-empty.xq");
assertEquals(SequenceCardinality.ONE, infer("sum(())", uri).getCardinality());
assertEquals(BuiltinTypesCatalogue.integerItem, infer("sum(())", uri).getItemType());
assertEquals(
BuiltinTypesCatalogue.stringItem,
infer("sum((), \"zero\")", uri).getItemType());
assertTrue(infer("sum((), ())", uri).isEmptySequence());
}

@ParameterizedTest
@ValueSource(booleans = {true, false})
void callableFieldAlternativesCanBeInvoked(boolean singleton) {
String query = "let $f := function($x as xs:integer) as xs:integer {$x} "
+ "let $o := {\"f\": if ("
+ singleton
+ ") then $f else ($f, $f)} return ($o.f)(1)";
List<Item> values = new Rumble(CONFIGURATION).runQuery(query).getAsList();
assertEquals(1, values.size());
if (singleton) {
assertEquals(1, values.get(0).getIntValue());
} else {
assertTrue(values.get(0).isFunction());
}
}

@ParameterizedTest
@ValueSource(
strings = {
"for $x in (1, \"a\") return $x + 1",
"for $x in (xs:float(1), \"a\") return math:pow($x, 2)",
"for $x in (xs:date(\"2000-01-01\"), xs:dateTime(\"2000-01-01T00:00:00\")) return $x cast as xs:double"
})
void incompatibleUnionMembersAreStillRejected(String query) {
for (String extension : List.of("jq", "xq")) {
assertThrows(
UnexpectedStaticTypeException.class,
() -> infer(query, URI.create("file:///invalid-union." + extension)));
}
}

@Test
void nullableSumWithEmptyZeroRemainsOptional() {
URI uri = URI.create("file:///sum-optional.xq");
assertEquals(
SequenceCardinality.ZERO_OR_ONE,
infer("sum((1, 2)[xs:boolean(\"false\")], ())", uri).getCardinality());
}

@Test
void nonAtomicCastOperandsRemainRuntimeChecksWhenStaticTypingIsDisabled() {
RumbleConfiguration configuration = RumbleConfiguration.builder()
.configureAnalysis(analysis -> analysis.enableStaticTyping(false))
.build();
assertDoesNotThrow(() -> CompilationPipeline.compileMainModule(
"[1] cast as xs:string",
URI.create("file:///array-cast.jq"),
new CompilationConfiguration(configuration),
ExternalBindings.empty()));
}

private static SequenceType infer(String query, URI uri) {
return CompilationPipeline.compileMainModule(
query, uri, new CompilationConfiguration(CONFIGURATION), ExternalBindings.empty())
.getExpression()
.getStaticSequenceType();
}
}
Loading
Loading