diff --git a/.travis.yml b/.travis.yml
index a0217279..19e25c84 100644
--- a/.travis.yml
+++ b/.travis.yml
@@ -20,6 +20,7 @@ branches:
script:
# Print output every minute to avoid travis timeout
- while sleep 1m; do echo "=====[ $SECONDS seconds elapsed -- still running ]====="; done &
- - ./gradlew build -s && ./gradlew -p transportable-udfs-examples build -s && ./gradlew ciPerformRelease -s
+ # With the exception of release commands, all build logic goes in travis-build.sh
+ - ./travis-build.sh && ./gradlew ciPerformRelease -s
# Killing background sleep loop
- kill %1
diff --git a/build.gradle b/build.gradle
index c737741e..eb5d6ea7 100644
--- a/build.gradle
+++ b/build.gradle
@@ -18,7 +18,8 @@ buildscript {
}
plugins {
- id "org.shipkit.java" version "2.0.31"
+ id "org.shipkit.java" version "2.3.4"
+ id "checkstyle"
}
allprojects {
@@ -48,11 +49,24 @@ subprojects {
strictCheck true
}
+ configurations {
+ all {
+ // Transport as a library should only expose slf4j-api to its consumers and should not keep any SLF4J bindings
+ // in its dependency graph
+ // Quote from http://www.slf4j.org/faq.html#maven2
+ // "Thus, as far as your users are concerned you are exporting slf4j-api as a transitive dependency of your
+ // library, but not any SLF4J-binding or any underlying logging system."
+ // Our dependencies (e.g. hadoop-common) do bring in slf4j-log4j12 in their dependency graph so we exclude them
+ exclude group: 'org.slf4j', module: 'slf4j-log4j12'
+ }
+ }
+
plugins.withType(JavaPlugin) {
project.apply(plugin: 'checkstyle')
dependencies {
testCompile 'org.testng:testng:6.11'
+ testCompile 'org.slf4j:slf4j-simple:1.7.25'
}
test {
@@ -60,7 +74,9 @@ subprojects {
}
checkstyle {
- configFile = file("${rootDir}/gradle/checkstyle/linkedin-checkstyle.xml")
+ configFile = file("${rootDir}/gradle/checkstyle/checkstyle.xml")
+ configProperties = ['config_loc' : "${rootDir}/gradle/checkstyle/"]
+ toolVersion '8.23'
}
}
diff --git a/defaultEnvironment.gradle b/defaultEnvironment.gradle
index ced4d3d5..c6b83602 100644
--- a/defaultEnvironment.gradle
+++ b/defaultEnvironment.gradle
@@ -10,8 +10,8 @@ subprojects {
url "https://conjars.org/repo"
}
}
- project.ext.setProperty('presto-version', '319')
- project.ext.setProperty('airlift-slice-version', '0.33')
+ project.ext.setProperty('presto-version', '333')
+ project.ext.setProperty('airlift-slice-version', '0.38')
project.ext.setProperty('spark-group', 'org.apache.spark')
project.ext.setProperty('spark-version', '2.3.0')
}
diff --git a/docs/release-notes.md b/docs/release-notes.md
index cb7b2b3c..11d75db1 100644
--- a/docs/release-notes.md
+++ b/docs/release-notes.md
@@ -1,5 +1,76 @@
*Release notes were automatically generated by [Shipkit](http://shipkit.org/)*
+#### 0.0.61
+ - 2020-11-26 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.60...v0.0.61) by [Carl Steinbach](https://github.com/cwsteinbach) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.61)
+ - Fail build if there are checkstyle violations [(#64)](https://github.com/linkedin/transport/pull/64)
+
+#### 0.0.60
+ - 2020-11-18 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.59...v0.0.60) by [Raymond](https://github.com/raymondlam12) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.60)
+ - Add Avro ENUM read support. [(#62)](https://github.com/linkedin/transport/pull/62)
+
+#### 0.0.59
+ - 2020-11-10 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.58...v0.0.59) by [curtiscwang](https://github.com/curtiscwang) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.59)
+ - Support conversion of String type to Utf8 in AvroWrapper [(#61)](https://github.com/linkedin/transport/pull/61)
+
+#### 0.0.58
+ - 2020-11-04 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.57...v0.0.58) by [Khai Tran](https://github.com/khaitranq) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.58)
+ - Support simple Avro Union schema [(#60)](https://github.com/linkedin/transport/pull/60)
+
+#### 0.0.57
+ - 2020-09-23 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.56...v0.0.57) by [Raymond](https://github.com/raymondlam12) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.57)
+ - Create a MutableMap when accessing keySet and values in SparkMap [(#58)](https://github.com/linkedin/transport/pull/58)
+
+#### 0.0.56
+ - 2020-08-17 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.55...v0.0.56) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.56)
+ - Fix test SQL generation for binary inputs [(#55)](https://github.com/linkedin/transport/pull/55)
+
+#### 0.0.55
+ - 2020-08-13 - [3 commits](https://github.com/linkedin/transport/compare/v0.0.54...v0.0.55) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.55)
+ - Bump shipkit version [(#54)](https://github.com/linkedin/transport/pull/54)
+ - Remove slf4j-log4j12 from Transport dependency graph [(#51)](https://github.com/linkedin/transport/pull/51)
+ - Hive: Struct data should not be converted to object array during StdStruct creation [(#50)](https://github.com/linkedin/transport/pull/50)
+
+#### 0.0.54
+ - 2020-08-10 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.53...v0.0.54) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.54)
+ - Publish Presto thin jar which allows consumers to control dependency graph [(#49)](https://github.com/linkedin/transport/pull/49)
+
+#### 0.0.53
+ - 2020-06-30 - [2 commits](https://github.com/linkedin/transport/compare/v0.0.52...v0.0.53) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.53)
+ - No pull requests referenced in commit messages.
+
+#### 0.0.52
+ - 2020-06-26 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.51...v0.0.52) by [John Joyce](https://github.com/jjoyce0510) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.52)
+ - No pull requests referenced in commit messages.
+
+#### 0.0.51
+ - 2020-06-23 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.50...v0.0.51) by [Khai Tran](https://github.com/khaitranq) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.51)
+ - Add support for StdFloat, StdDouble, and StdBinary [(#46)](https://github.com/linkedin/transport/pull/46)
+
+#### 0.0.50
+ - 2020-05-08 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.49...v0.0.50) by [Xingyuan Lin](https://github.com/lxynov) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.50)
+ - Upgrade to PrestoSQL 333 [(#45)](https://github.com/linkedin/transport/pull/45)
+
+#### 0.0.49
+ - 2020-04-30 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.48...v0.0.49) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.49)
+ - Presto: Make ScalarFunctionImplementation state independent of StdUdfWrapper [(#44)](https://github.com/linkedin/transport/pull/44)
+
+#### 0.0.48
+ - 2020-04-14 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.47...v0.0.48) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.48)
+ - Presto: Pass custom configuration object when using FileSystemUtils [(#43)](https://github.com/linkedin/transport/pull/43)
+
+#### 0.0.47
+ - 2020-04-01 - [2 commits](https://github.com/linkedin/transport/compare/v0.0.46...v0.0.47) by [Suren Nihalani](https://github.com/SurenNihalani) (1), [Sushant Raikar](https://github.com/HotSushi) (1) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.47)
+ - FileSystemUtils: remove an unreliable check for unit testing [(#42)](https://github.com/linkedin/transport/pull/42)
+ - Fixed Presto UDF patch's broken link [(#41)](https://github.com/linkedin/transport/pull/41)
+
+#### 0.0.46
+ - 2020-02-03 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.45...v0.0.46) by [Shardul Mahadik](https://github.com/shardulm94) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.46)
+ - Disable Jacoco for platform tests [(#37)](https://github.com/linkedin/transport/pull/37)
+
+#### 0.0.45
+ - 2020-01-14 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.44...v0.0.45) by [Suren Nihalani](https://github.com/SurenNihalani) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.45)
+ - Allow for absolute paths instead of assuming defaultFS in UDF's required files [(#33)](https://github.com/linkedin/transport/pull/33)
+
#### 0.0.44
- 2019-12-13 - [1 commit](https://github.com/linkedin/transport/compare/v0.0.43...v0.0.44) by [Raymond](https://github.com/raymondlam12) - published to [](https://bintray.com/linkedin-transport/maven/transport/0.0.44)
- Remove log4j properties from main sources [(#32)](https://github.com/linkedin/transport/pull/32)
diff --git a/docs/transport-udfs-api.md b/docs/transport-udfs-api.md
index fadaa75f..deecd1cb 100644
--- a/docs/transport-udfs-api.md
+++ b/docs/transport-udfs-api.md
@@ -9,7 +9,8 @@ The `StdType` interface is the parent class of all type objects that
are used to describe the schema of the data objects that can be
manipulated by `StdUDFs`. Sub-interfaces of this interface include
`StdIntegerType`, `StdBooleanType`, `StdLongType`, `StdStringType`,
-`StdArrayType`, `StdMapType`, `StdStructType`. Each sub-interface is
+`StdDoubleType`, `StdFloatType`, `StdBinaryType`, `StdArrayType`,
+`StdMapType`, and `StdStructType`. Each sub-interface is
defined by methods that are specific to the corresponding type. For
example, `StdMapType` interface is defined by the two methods shown
below. The `keyType()` and `valueType()` methods can be used to obtain
@@ -39,9 +40,10 @@ public interface StdStructType extends StdType {
manipulated by Transport UDFs. As a top-level interface, `StdData`
itself does not contain any methods. A number of type-specific
interfaces extend `StdData`, such as `StdInteger`, `StdLong`,
-`StdBoolean`, `StdString`, `StdArray`, `StdMap`, `StdStruct` to
-represent `INTEGER`, `LONG`, `BOOLEAN`, `VARCHAR`, `ARRAY`, `MAP`,
-`STRUCT` SQL types respectively. Each of those interfaces exposes
+`StdBoolean`, `StdString`, `StdDouble`, `StdFloat`, `StdBinary`,
+`StdArray`, `StdMap`, and `StdStruct` to represent `INTEGER`,
+`LONG`, `BOOLEAN`, `VARCHAR`, `DOUBLE`, `REAL`, `VARBINARY`, `ARRAY`, `MAP`,
+and `STRUCT` SQL types respectively. Each of those interfaces exposes
operations that can manipulate that type of data. For example,
`StdMap` interface is defined by the following methods:
@@ -108,6 +110,12 @@ definition:
is StdInteger.
* `"boolean"`: to represent SQL Boolean type. The respective Standard
Type is StdBoolean.
+* `"double"`: to represent SQL Double type. The respective Standard
+ Type is StdDouble.
+* `"real"`: to represent SQL Real type. The respective Standard
+ Type is StdFloat.
+* `"varbinary"`: to represent SQL Binary type. The respective Standard
+ Type is StdBinary.
* `"array(T)"`: to represent SQL Array type, with elements of type
T. The respective Standard Type is StdArray.
* `"map(K,V)"`: to represent SQL Map type, with keys of type K and
@@ -132,6 +140,9 @@ public interface StdFactory {
StdLong createLong(long value);
StdBoolean createBoolean(boolean value);
StdString createString(String value);
+ StdDouble createDouble(double value);
+ StdFloat createFloat(float value);
+ StdBinary createBinary(ByteBuffer value);
StdArray createArray(StdType stdType, int expectedSize);
StdArray createArray(StdType stdType);
StdMap createMap(StdType stdType);
diff --git a/transportable-udfs-documentation/transport-udfs-presto.patch b/docs/transport-udfs-presto.patch
similarity index 58%
rename from transportable-udfs-documentation/transport-udfs-presto.patch
rename to docs/transport-udfs-presto.patch
index 3fa1a9b8..c29b1dd6 100644
--- a/transportable-udfs-documentation/transport-udfs-presto.patch
+++ b/docs/transport-udfs-presto.patch
@@ -1,28 +1,5 @@
-diff --git a/presto-main/pom.xml b/presto-main/pom.xml
-index 83a20a23fe..02afa5310f 100644
---- a/presto-main/pom.xml
-+++ b/presto-main/pom.xml
-@@ -424,6 +424,18 @@
- test
-
-
-+
-+ org.apache.maven.plugins
-+ maven-jar-plugin
-+ 2.2
-+
-+
-+
-+ test-jar
-+
-+
-+
-+
-
-
-
diff --git a/presto-main/src/main/java/io/prestosql/server/PluginManager.java b/presto-main/src/main/java/io/prestosql/server/PluginManager.java
-index f02ceeab03..88de943bed 100644
+index abcd001031..053c17aeed 100644
--- a/presto-main/src/main/java/io/prestosql/server/PluginManager.java
+++ b/presto-main/src/main/java/io/prestosql/server/PluginManager.java
@@ -23,6 +23,7 @@ import io.prestosql.connector.ConnectorManager;
@@ -31,17 +8,17 @@ index f02ceeab03..88de943bed 100644
import io.prestosql.metadata.MetadataManager;
+import io.prestosql.metadata.SqlScalarFunction;
import io.prestosql.security.AccessControlManager;
+ import io.prestosql.security.GroupProviderManager;
import io.prestosql.server.security.PasswordAuthenticatorManager;
- import io.prestosql.spi.Plugin;
-@@ -52,6 +53,7 @@ import java.util.List;
- import java.util.ServiceLoader;
+@@ -54,6 +55,7 @@ import java.util.ServiceLoader;
import java.util.Set;
import java.util.concurrent.atomic.AtomicBoolean;
+ import java.util.function.Supplier;
+import java.util.stream.Collectors;
import static com.google.common.base.Preconditions.checkState;
import static io.prestosql.metadata.FunctionExtractor.extractFunctions;
-@@ -62,8 +64,22 @@ import static java.util.Objects.requireNonNull;
+@@ -64,8 +66,22 @@ import static java.util.Objects.requireNonNull;
@ThreadSafe
public class PluginManager
{
@@ -54,8 +31,8 @@ index f02ceeab03..88de943bed 100644
private static final ImmutableList SPI_PACKAGES = ImmutableList.builder()
+ // io.prestosql.metadata is required for SqlScalarFunction and FunctionRegistry classes
+ .add("io.prestosql.metadata.")
-+ // io.prestosql.operator.scalar is required for ScalarFunctionImplementation
-+ .add("io.prestosql.operator.scalar.")
++ // io.prestosql.operator. is required for ScalarFunctionImplementation and TypeSignatureParser
++ .add("io.prestosql.operator.")
.add("io.prestosql.spi.")
+ // io.prestosql.type is required for TypeManager, and all supported types
+ .add("io.prestosql.type.")
@@ -63,8 +40,8 @@ index f02ceeab03..88de943bed 100644
+ .add("io.prestosql.util.")
.add("com.fasterxml.jackson.annotation.")
.add("io.airlift.slice.")
- .add("io.airlift.units.")
-@@ -155,7 +171,21 @@ public class PluginManager
+ .add("org.openjdk.jol.")
+@@ -159,11 +175,22 @@ public class PluginManager
{
ServiceLoader serviceLoader = ServiceLoader.load(Plugin.class, pluginClassLoader);
List plugins = ImmutableList.copyOf(serviceLoader);
@@ -74,32 +51,21 @@ index f02ceeab03..88de943bed 100644
+ pluginClassLoader);
+ List sqlScalarFunctions = ImmutableList.copyOf(sqlScalarFunctionsServiceLoader);
+
-+ checkState(!plugins.isEmpty() || !sqlScalarFunctions.isEmpty(), "No service providers of type %s or %s",
-+ Plugin.class.getName(), SqlScalarFunction.class.getName());
++ checkState(!plugins.isEmpty() || !sqlScalarFunctions.isEmpty(), "No service providers of type %s or %s", Plugin.class.getName(), SqlScalarFunction.class.getName());
+
-+ installPlugins(plugins);
-+
-+ registerSqlScalarFunctions(sqlScalarFunctions);
-+ }
-+
-+ private void installPlugins(List plugins)
-+ {
for (Plugin plugin : plugins) {
log.info("Installing %s", plugin.getClass().getName());
- installPlugin(plugin);
-@@ -215,6 +245,15 @@ public class PluginManager
+ installPlugin(plugin, pluginClassLoader::duplicate);
}
- }
-
-+ public void registerSqlScalarFunctions(List sqlScalarFunctions)
-+ {
++
+ for (SqlScalarFunction sqlScalarFunction : sqlScalarFunctions) {
-+ log.info("Registering function %s(%s)", sqlScalarFunction.getSignature().getName(), sqlScalarFunction.getSignature().getArgumentTypes().stream().map(e -> e.toString()).collect(
-+ Collectors.joining(", ")));
++ log.info("Registering function %s(%s)",
++ sqlScalarFunction.getFunctionMetadata().getSignature().getName(),
++ sqlScalarFunction.getFunctionMetadata().getSignature().getArgumentTypes().stream()
++ .map(e -> e.toString())
++ .collect(Collectors.joining(", ")));
+ metadataManager.addFunctions(ImmutableList.of(sqlScalarFunction));
+ }
-+ }
-+
- private URLClassLoader buildClassLoader(String plugin)
- throws Exception
- {
+ }
+
+ public void installPlugin(Plugin plugin, Supplier duplicatePluginClassLoaderFactory)
diff --git a/docs/using-transport-udfs.md b/docs/using-transport-udfs.md
index 70f06edc..9b4e5f97 100644
--- a/docs/using-transport-udfs.md
+++ b/docs/using-transport-udfs.md
@@ -86,7 +86,7 @@ If the UDF class is `com.linkedin.transport.example.ExampleUDF` then the platfor
Unlike Hive and Spark, Presto currently does not allow dynamically loading jar files once the Presto server has started.
In Presto, the jar is deployed to the `plugin` directory.
However, a small patch is required for the Presto engine to recognize the jar as a plugin, since the generated Presto UDFs implement the `SqlScalarFunction` API, which is currently not part of Presto's SPI architecture.
-You can find the patch [here](transportable-udfs-documentation/transport-udfs-presto.patch) and apply it before deploying your UDFs jar to the Presto engine.
+You can find the patch [here](transport-udfs-presto.patch) and apply it before deploying your UDFs jar to the Presto engine.
2. Call the UDF in a query
To call the UDF, you will need to use the function name defined in the Transport UDF definition.
diff --git a/gradle/checkstyle/linkedin-checkstyle.xml b/gradle/checkstyle/checkstyle.xml
similarity index 98%
rename from gradle/checkstyle/linkedin-checkstyle.xml
rename to gradle/checkstyle/checkstyle.xml
index 6af59c2e..a205c87e 100644
--- a/gradle/checkstyle/linkedin-checkstyle.xml
+++ b/gradle/checkstyle/checkstyle.xml
@@ -8,7 +8,7 @@
LinkedIn Java style.
-->
-
+
@@ -189,11 +189,11 @@ LinkedIn Java style.
- -->
+
diff --git a/gradle/checkstyle/suppressions.xml b/gradle/checkstyle/suppressions.xml
new file mode 100644
index 00000000..d4d3fa1a
--- /dev/null
+++ b/gradle/checkstyle/suppressions.xml
@@ -0,0 +1,9 @@
+
+
+
+
+
+
+
diff --git a/transportable-udfs-annotation-processor/src/main/java/com/linkedin/transport/processor/TransportProcessor.java b/transportable-udfs-annotation-processor/src/main/java/com/linkedin/transport/processor/TransportProcessor.java
index 12e25002..30a110a8 100644
--- a/transportable-udfs-annotation-processor/src/main/java/com/linkedin/transport/processor/TransportProcessor.java
+++ b/transportable-udfs-annotation-processor/src/main/java/com/linkedin/transport/processor/TransportProcessor.java
@@ -130,10 +130,19 @@ private void processUDFClass(TypeElement udfClassElement) {
udfClassElement
);
} else {
- String topLevelStdUdfClassName =
- elementsOverridingTopLevelStdUDFMethods.iterator().next().getQualifiedName().toString();
+ TypeElement topLevelStdUdfTypeElement = elementsOverridingTopLevelStdUDFMethods.iterator().next();
+ String topLevelStdUdfClassName = topLevelStdUdfTypeElement.getQualifiedName().toString();
debug(String.format("TopLevelStdUDF class found: %s", topLevelStdUdfClassName));
+ String udfClassName = udfClassElement.getQualifiedName().toString();
_transportUdfMetadata.addUDF(topLevelStdUdfClassName, udfClassElement.getQualifiedName().toString());
+ _transportUdfMetadata.setClassNumberOfTypeParameters(
+ topLevelStdUdfClassName,
+ topLevelStdUdfTypeElement.getTypeParameters().size()
+ );
+ _transportUdfMetadata.setClassNumberOfTypeParameters(
+ udfClassName,
+ udfClassElement.getTypeParameters().size()
+ );
}
}
diff --git a/transportable-udfs-annotation-processor/src/test/java/com/linkedin/transport/processor/TransportProcessorTest.java b/transportable-udfs-annotation-processor/src/test/java/com/linkedin/transport/processor/TransportProcessorTest.java
index 0322f306..dd3c25c5 100644
--- a/transportable-udfs-annotation-processor/src/test/java/com/linkedin/transport/processor/TransportProcessorTest.java
+++ b/transportable-udfs-annotation-processor/src/test/java/com/linkedin/transport/processor/TransportProcessorTest.java
@@ -81,7 +81,7 @@ public void shouldNotContainMultipleOverridingsOfTopLevelStdUDFMethods1() throws
.withErrorCount(1)
.withErrorContaining(Constants.MORE_THAN_ONE_TYPE_OVERRIDING_ERROR)
.in(forResource("udfs/UDFWithMultipleInterfaces1.java"))
- .onLine(14)
+ .onLine(13)
.atColumn(8);
}
@@ -96,7 +96,7 @@ public void shouldNotContainMultipleOverridingsOfTopLevelStdUDFMethods2() throws
.withErrorCount(1)
.withErrorContaining(Constants.MORE_THAN_ONE_TYPE_OVERRIDING_ERROR)
.in(forResource("udfs/UDFWithMultipleInterfaces2.java"))
- .onLine(13)
+ .onLine(12)
.atColumn(8);
}
@@ -110,7 +110,7 @@ public void udfShouldNotOverrideInterfaceMethods() throws IOException {
.withErrorCount(1)
.withErrorContaining(Constants.MORE_THAN_ONE_TYPE_OVERRIDING_ERROR)
.in(forResource("udfs/UDFOverridingInterfaceMethod.java"))
- .onLine(14)
+ .onLine(13)
.atColumn(8);
}
@@ -123,7 +123,7 @@ public void udfShouldImplementTopLevelStdUDF() throws IOException {
.withErrorCount(1)
.withErrorContaining(Constants.INTERFACE_NOT_IMPLEMENTED_ERROR)
.in(forResource("udfs/UDFNotImplementingTopLevelStdUDF.java"))
- .onLine(14)
+ .onLine(13)
.atColumn(8);
}
diff --git a/transportable-udfs-annotation-processor/src/test/resources/outputs/empty.json b/transportable-udfs-annotation-processor/src/test/resources/outputs/empty.json
index 7836d3b6..4e9cfcb3 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/outputs/empty.json
+++ b/transportable-udfs-annotation-processor/src/test/resources/outputs/empty.json
@@ -1,3 +1,4 @@
{
- "udfs": []
+ "udfs": {},
+ "classToNumberOfTypeParameters": {}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/outputs/overloadedUDF.json b/transportable-udfs-annotation-processor/src/test/resources/outputs/overloadedUDF.json
index 9f6a2450..4f0526f0 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/outputs/overloadedUDF.json
+++ b/transportable-udfs-annotation-processor/src/test/resources/outputs/overloadedUDF.json
@@ -1,11 +1,13 @@
{
- "udfs": [
- {
- "topLevelClass": "udfs.OverloadedUDF1",
- "stdUDFImplementations": [
- "udfs.OverloadedUDFInt",
- "udfs.OverloadedUDFString"
- ]
- }
- ]
+ "udfs": {
+ "udfs.OverloadedUDF1": [
+ "udfs.OverloadedUDFInt",
+ "udfs.OverloadedUDFString"
+ ]
+ },
+ "classToNumberOfTypeParameters": {
+ "udfs.OverloadedUDFString": 0,
+ "udfs.OverloadedUDF1": 0,
+ "udfs.OverloadedUDFInt": 0
+ }
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/outputs/simpleUDF.json b/transportable-udfs-annotation-processor/src/test/resources/outputs/simpleUDF.json
index 9323ddbd..34c7cee2 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/outputs/simpleUDF.json
+++ b/transportable-udfs-annotation-processor/src/test/resources/outputs/simpleUDF.json
@@ -1,10 +1,10 @@
{
- "udfs": [
- {
- "topLevelClass": "udfs.SimpleUDF",
- "stdUDFImplementations": [
- "udfs.SimpleUDF"
- ]
- }
- ]
+ "udfs": {
+ "udfs.SimpleUDF": [
+ "udfs.SimpleUDF"
+ ]
+ },
+ "classToNumberOfTypeParameters": {
+ "udfs.SimpleUDF": 0
+ }
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDF.json b/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDF.json
index ab58d4d8..5b72a274 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDF.json
+++ b/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDF.json
@@ -1,10 +1,10 @@
{
- "udfs": [
- {
- "topLevelClass": "udfs.UDFExtendingAbstractUDF",
- "stdUDFImplementations": [
- "udfs.UDFExtendingAbstractUDF"
- ]
- }
- ]
+ "udfs": {
+ "udfs.UDFExtendingAbstractUDF": [
+ "udfs.UDFExtendingAbstractUDF"
+ ]
+ },
+ "classToNumberOfTypeParameters": {
+ "udfs.UDFExtendingAbstractUDF": 0
+ }
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDFImplementingInterface.json b/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDFImplementingInterface.json
index d2551a77..b75531e1 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDFImplementingInterface.json
+++ b/transportable-udfs-annotation-processor/src/test/resources/outputs/udfExtendingAbstractUDFImplementingInterface.json
@@ -1,10 +1,11 @@
{
- "udfs": [
- {
- "topLevelClass": "udfs.AbstractUDFImplementingInterface",
- "stdUDFImplementations": [
- "udfs.UDFExtendingAbstractUDFImplementingInterface"
- ]
- }
- ]
+ "udfs": {
+ "udfs.AbstractUDFImplementingInterface": [
+ "udfs.UDFExtendingAbstractUDFImplementingInterface"
+ ]
+ },
+ "classToNumberOfTypeParameters": {
+ "udfs.UDFExtendingAbstractUDFImplementingInterface": 0,
+ "udfs.AbstractUDFImplementingInterface": 0
+ }
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDF.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDF.java
index 4a482115..06536aa2 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDF.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDF.java
@@ -5,11 +5,10 @@
*/
package udfs;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF0;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
-public abstract class AbstractUDF extends StdUDF0 implements TopLevelStdUDF {
+public abstract class AbstractUDF extends StdUDF0 implements TopLevelStdUDF {
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDFImplementingInterface.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDFImplementingInterface.java
index 4078d7bc..7d85fb36 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDFImplementingInterface.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/AbstractUDFImplementingInterface.java
@@ -5,12 +5,11 @@
*/
package udfs;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF0;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
-public abstract class AbstractUDFImplementingInterface extends StdUDF0 implements TopLevelStdUDF {
+public abstract class AbstractUDFImplementingInterface extends StdUDF0 implements TopLevelStdUDF {
@Override
public String getFunctionName() {
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/OuterClassForInnerUDF.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/OuterClassForInnerUDF.java
index d5d2551d..a84ee746 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/OuterClassForInnerUDF.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/OuterClassForInnerUDF.java
@@ -6,14 +6,13 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF0;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
public class OuterClassForInnerUDF {
- public class InnerUDF extends StdUDF0 implements TopLevelStdUDF {
+ public class InnerUDF extends StdUDF0 implements TopLevelStdUDF {
@Override
public String getFunctionName() {
@@ -36,7 +35,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFInt.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFInt.java
index 3f130d9d..292c8606 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFInt.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFInt.java
@@ -6,12 +6,11 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdInteger;
import com.linkedin.transport.api.udf.StdUDF0;
import java.util.List;
-public class OverloadedUDFInt extends StdUDF0 implements OverloadedUDF1 {
+public class OverloadedUDFInt extends StdUDF0 implements OverloadedUDF1 {
@Override
public List getInputParameterSignatures() {
@@ -24,7 +23,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdInteger eval() {
+ public Integer eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFString.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFString.java
index 9782d683..d0855f55 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFString.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/OverloadedUDFString.java
@@ -6,12 +6,11 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF0;
import java.util.List;
-public class OverloadedUDFString extends StdUDF0 implements OverloadedUDF1 {
+public class OverloadedUDFString extends StdUDF0 implements OverloadedUDF1 {
@Override
public List getInputParameterSignatures() {
@@ -24,7 +23,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/SimpleUDF.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/SimpleUDF.java
index 15e749ec..46231c63 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/SimpleUDF.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/SimpleUDF.java
@@ -6,13 +6,12 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF0;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class SimpleUDF extends StdUDF0 implements TopLevelStdUDF {
+public class SimpleUDF extends StdUDF0 implements TopLevelStdUDF {
@Override
public String getFunctionName() {
@@ -35,7 +34,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDF.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDF.java
index 99fb068c..564fcd7e 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDF.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDF.java
@@ -6,7 +6,6 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
@@ -34,7 +33,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDFImplementingInterface.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDFImplementingInterface.java
index 2fc6d5ce..db8e3bc7 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDFImplementingInterface.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFExtendingAbstractUDFImplementingInterface.java
@@ -6,7 +6,6 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import java.util.List;
@@ -23,7 +22,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFNotImplementingTopLevelStdUDF.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFNotImplementingTopLevelStdUDF.java
index 43e8ffa9..862403e4 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFNotImplementingTopLevelStdUDF.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFNotImplementingTopLevelStdUDF.java
@@ -6,12 +6,11 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF0;
import java.util.List;
-public class UDFNotImplementingTopLevelStdUDF extends StdUDF0 {
+public class UDFNotImplementingTopLevelStdUDF extends StdUDF0 {
@Override
public List getInputParameterSignatures() {
@@ -24,7 +23,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFOverridingInterfaceMethod.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFOverridingInterfaceMethod.java
index 346ff553..2d97547e 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFOverridingInterfaceMethod.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFOverridingInterfaceMethod.java
@@ -6,12 +6,11 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdBoolean;
import com.linkedin.transport.api.udf.StdUDF0;
import java.util.List;
-public class UDFOverridingInterfaceMethod extends StdUDF0 implements OverloadedUDF1 {
+public class UDFOverridingInterfaceMethod extends StdUDF0 implements OverloadedUDF1 {
@Override
public String getFunctionName() {
@@ -29,7 +28,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdBoolean eval() {
+ public Boolean eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces1.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces1.java
index 54e57c16..83a38c4d 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces1.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces1.java
@@ -6,12 +6,11 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdBoolean;
import com.linkedin.transport.api.udf.StdUDF0;
import java.util.List;
-public class UDFWithMultipleInterfaces1 extends StdUDF0 implements OverloadedUDF1, OverloadedUDF2 {
+public class UDFWithMultipleInterfaces1 extends StdUDF0 implements OverloadedUDF1, OverloadedUDF2 {
@Override
public String getFunctionName() {
@@ -34,7 +33,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdBoolean eval() {
+ public Boolean eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces2.java b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces2.java
index e8e62bb1..f2fbe270 100644
--- a/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces2.java
+++ b/transportable-udfs-annotation-processor/src/test/resources/udfs/UDFWithMultipleInterfaces2.java
@@ -6,7 +6,6 @@
package udfs;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdString;
import java.util.List;
@@ -33,7 +32,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdString eval() {
+ public String eval() {
return null;
}
}
\ No newline at end of file
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/StdFactory.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/StdFactory.java
index 3b9490af..d04d0ff5 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/StdFactory.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/StdFactory.java
@@ -5,17 +5,11 @@
*/
package com.linkedin.transport.api;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdBoolean;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdInteger;
-import com.linkedin.transport.api.data.StdLong;
-import com.linkedin.transport.api.data.StdMap;
-import com.linkedin.transport.api.data.StdString;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.ArrayData;
+import com.linkedin.transport.api.data.MapData;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.api.types.StdArrayType;
import com.linkedin.transport.api.types.StdMapType;
-import com.linkedin.transport.api.types.StdStructType;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF;
import java.io.Serializable;
@@ -23,7 +17,8 @@
/**
- * {@link StdFactory} is used to create {@link StdData} and {@link StdType} objects inside Standard UDFs.
+ * {@link StdFactory} is used to create containter types (e.g., {@link ArrayData}, {@link MapData}, {@link RowData})
+ * and {@link StdType} objects inside Standard UDFs.
*
* Specific APIs of {@link StdFactory} are implemented by each target platform (e.g., Spark, Presto, Hive) individually.
* A {@link StdFactory} object is available inside Standard UDFs using {@link StdUDF#getStdFactory()}.
@@ -32,111 +27,79 @@
public interface StdFactory extends Serializable {
/**
- * Creates a {@link StdInteger} representing a given integer value.
- *
- * @param value the input integer value
- * @return {@link StdInteger} with the given integer value
- */
- StdInteger createInteger(int value);
-
- /**
- * Creates a {@link StdLong} representing a given long value.
- *
- * @param value the input long value
- * @return {@link StdLong} with the given long value
- */
- StdLong createLong(long value);
-
- /**
- * Creates a {@link StdBoolean} representing a given boolean value.
- *
- * @param value the input boolean value
- * @return {@link StdBoolean} with the given boolean value
- */
- StdBoolean createBoolean(boolean value);
-
- /**
- * Creates a {@link StdString} representing a given {@link String} value.
- *
- * @param value the input {@link String} value
- * @return {@link StdString} with the given {@link String} value
- */
- StdString createString(String value);
-
- /**
- * Creates an empty {@link StdArray} whose type is given by the given {@link StdType}.
+ * Creates an empty {@link ArrayData} whose type is given by the given {@link StdType}.
*
* It is expected that the top-level {@link StdType} is a {@link StdArrayType}.
*
* @param stdType type of the array to be created
* @param expectedSize expected number of entries in the array
- * @return an empty {@link StdArray}
+ * @return an empty {@link ArrayData}
*/
- StdArray createArray(StdType stdType, int expectedSize);
+ ArrayData createArray(StdType stdType, int expectedSize);
/**
- * Creates an empty {@link StdArray} whose type is given by the given {@link StdType}.
+ * Creates an empty {@link ArrayData} whose type is given by the given {@link StdType}.
*
* It is expected that the top-level {@link StdType} is a {@link StdArrayType}.
*
* @param stdType type of the array to be created
- * @return an empty {@link StdArray}
+ * @return an empty {@link ArrayData}
*/
- StdArray createArray(StdType stdType);
+ ArrayData createArray(StdType stdType);
/**
- * Creates an empty {@link StdMap} whose type is given by the given {@link StdType}.
+ * Creates an empty {@link MapData} whose type is given by the given {@link StdType}.
*
* It is expected that the top-level {@link StdType} is a {@link StdMapType}.
*
* @param stdType type of the map to be created
- * @return an empty {@link StdMap}
+ * @return an empty {@link MapData}
*/
- StdMap createMap(StdType stdType);
+ MapData createMap(StdType stdType);
/**
- * Creates a {@link StdStruct} with the given field names and types.
+ * Creates a {@link RowData} with the given field names and types.
*
* @param fieldNames names of the struct fields
* @param fieldTypes types of the struct fields
- * @return a {@link StdStruct} with all fields initialized to null
+ * @return a {@link RowData} with all fields initialized to null
*/
- StdStruct createStruct(List fieldNames, List fieldTypes);
+ RowData createStruct(List fieldNames, List fieldTypes);
/**
- * Creates a {@link StdStruct} with the given field types. Field names will be field0, field1, field2...
+ * Creates a {@link RowData} with the given field types. Field names will be field0, field1, field2...
*
* @param fieldTypes types of the struct fields
- * @return a {@link StdStruct} with all fields initialized to null
+ * @return a {@link RowData} with all fields initialized to null
*/
- StdStruct createStruct(List fieldTypes);
+ RowData createStruct(List fieldTypes);
/**
- * Creates a {@link StdStruct} whose type is given by the given {@link StdType}.
+ * Creates a {@link RowData} whose type is given by the given {@link StdType}.
*
- * It is expected that the top-level {@link StdType} is a {@link StdStructType}.
+ * It is expected that the top-level {@link StdType} is a {@link com.linkedin.transport.api.types.RowType}.
*
* @param stdType type of the struct to be created
- * @return a {@link StdStruct} with all fields initialized to null
+ * @return a {@link RowData} with all fields initialized to null
*/
- StdStruct createStruct(StdType stdType);
+ RowData createStruct(StdType stdType);
/**
* Creates a {@link StdType} representing the given type signature.
*
* The following are considered valid type signatures:
*
- * {@code "varchar"} - Represents SQL varchar type. Corresponding standard type is {@link StdString}
- * {@code "integer"} - Represents SQL int type. Corresponding standard type is {@link StdInteger}
- * {@code "bigint"} - Represents SQL bigint/long type. Corresponding standard type is {@link StdLong}
- * {@code "boolean"} - Represents SQL boolean type. Corresponding standard type is {@link StdBoolean}
+ * {@code "varchar"} - Represents SQL varchar type. Corresponding Transport type is {@link String}
+ * {@code "integer"} - Represents SQL int type. Corresponding Transport type is {@link Integer}
+ * {@code "bigint"} - Represents SQL bigint/long type. Corresponding Transport type is {@link Long}
+ * {@code "boolean"} - Represents SQL boolean type. Corresponding Transport type is {@link Boolean}
* {@code "array(T)"} - Represents SQL array type, where {@code T} is type signature of array element.
- * Corresponding standard type is {@link StdArray}
+ * Corresponding Transport type is {@link ArrayData}
* {@code "map(K,V)"} - Represents SQL map type, where {@code K} and {@code V} are type signatures of the map
- * keys and values respectively. array element. Corresponding standard type is {@link StdMap}
+ * keys and values respectively. Corresponding Transport type is {@link MapData}
* {@code "row(f0 T0, f1 T1,... fn Tn)"} - Represents SQL struct type, where {@code f0}...{@code fn} are field
* names and {@code T0}...{@code Tn} are type signatures for the fields. Field names are optional; if not
- * specified they default to {@code field0}...{@code fieldn}. Corresponding standard type is {@link StdStruct}
+ * specified they default to {@code field0}...{@code fieldn}. Corresponding Transport type is {@link RowData}
*
*
* Generic type parameters can also be used as part of the type signatures; e.g., The type signature {@code "map(K,V)"}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdArray.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/ArrayData.java
similarity index 76%
rename from transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdArray.java
rename to transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/ArrayData.java
index ac698ae8..65a6b9a6 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdArray.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/ArrayData.java
@@ -5,8 +5,8 @@
*/
package com.linkedin.transport.api.data;
-/** A Standard UDF data type for representing arrays. */
-public interface StdArray extends StdData, Iterable {
+/** A Transport UDF data type for representing arrays. */
+public interface ArrayData extends Iterable {
/** Returns the number of elements in the array. */
int size();
@@ -16,12 +16,12 @@ public interface StdArray extends StdData, Iterable {
*
* @param idx the index of the element to be retrieved
*/
- StdData get(int idx);
+ E get(int idx);
/**
* Adds an element to the end of the array.
*
* @param e the element to append to the array
*/
- void add(StdData e);
+ void add(E e);
}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdMap.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/MapData.java
similarity index 78%
rename from transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdMap.java
rename to transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/MapData.java
index 8e67500e..39bd6965 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdMap.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/MapData.java
@@ -9,8 +9,8 @@
import java.util.Set;
-/** A Standard UDF data type for representing maps. */
-public interface StdMap extends StdData {
+/** A Transport UDF data type for representing maps. */
+public interface MapData {
/** Returns the number of key-value pairs in the map. */
int size();
@@ -20,7 +20,7 @@ public interface StdMap extends StdData {
*
* @param key the key whose value is to be returned
*/
- StdData get(StdData key);
+ V get(K key);
/**
* Adds the given value to the map against the given key.
@@ -28,18 +28,18 @@ public interface StdMap extends StdData {
* @param key the key to which the value is to be associated
* @param value the value to be associated with the key
*/
- void put(StdData key, StdData value);
+ void put(K key, V value);
/** Returns a {@link Set} of all the keys in the map. */
- Set keySet();
+ Set keySet();
/** Returns a {@link Collection} of all the values in the map. */
- Collection values();
+ Collection values();
/**
* Returns true if the map contains the given key, false otherwise.
*
* @param key the key to be checked
*/
- boolean containsKey(StdData key);
+ boolean containsKey(K key);
}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/PlatformData.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/PlatformData.java
index 75df0518..41fc7bae 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/PlatformData.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/PlatformData.java
@@ -5,7 +5,7 @@
*/
package com.linkedin.transport.api.data;
-/** An interface for all platform-specific implementations of {@link StdData}. */
+/** An interface to handle platform-specific container types. */
public interface PlatformData {
/** Returns the underlying platform-specific object holding the data. */
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdStruct.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/RowData.java
similarity index 50%
rename from transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdStruct.java
rename to transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/RowData.java
index 14ccff80..2d8f1ce0 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdStruct.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/RowData.java
@@ -8,39 +8,39 @@
import java.util.List;
-/** A Standard UDF data type for representing structs. */
-public interface StdStruct extends StdData {
+/** A Transport UDF data type for representing SQL ROW/STRUCT data type. */
+public interface RowData {
/**
- * Returns the value of the field at the given position in the struct.
+ * Returns the value of the field at the given position in the row.
*
- * @param index the position of the field in the struct
+ * @param index the position of the field in the row
*/
- StdData getField(int index);
+ Object getField(int index);
/**
- * Returns the value of the given field from the struct.
+ * Returns the value of the given field from the row.
*
* @param name the name of the field
*/
- StdData getField(String name);
+ Object getField(String name);
/**
- * Sets the value of the field at the given position in the struct.
+ * Sets the value of the field at the given position in the row.
*
- * @param index the position of the field in the struct
+ * @param index the position of the field in the row
* @param value the value to be set
*/
- void setField(int index, StdData value);
+ void setField(int index, Object value);
/**
- * Sets the value of the given field in the struct.
+ * Sets the value of the given field in the row.
*
* @param name the name of the field
* @param value the value to be set
*/
- void setField(String name, StdData value);
+ void setField(String name, Object value);
/** Returns a {@link List} of all fields in the struct. */
- List fields();
+ List fields();
}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdBoolean.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdBoolean.java
deleted file mode 100644
index ff230bc1..00000000
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdBoolean.java
+++ /dev/null
@@ -1,13 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.api.data;
-
-/** A Standard UDF data type for representing booleans. */
-public interface StdBoolean extends StdData {
-
- /** Returns the underlying boolean value. */
- boolean get();
-}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdData.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdData.java
deleted file mode 100644
index 77b3d1d7..00000000
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdData.java
+++ /dev/null
@@ -1,19 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.api.data;
-
-import com.linkedin.transport.api.StdFactory;
-
-
-/**
- * An interface for all data types in Standard UDFs.
- *
- * {@link StdData} is the main interface through which StdUDFs receive input data and return output data. All Standard
- * UDF data types (e.g., {@link StdInteger}, {@link StdArray}, {@link StdMap}) must extend {@link StdData}. Methods
- * inside {@link StdFactory} can be used to create {@link StdData} objects.
- */
-public interface StdData {
-}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdInteger.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdInteger.java
deleted file mode 100644
index c74a92dd..00000000
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdInteger.java
+++ /dev/null
@@ -1,13 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.api.data;
-
-/** A Standard UDF data type for representing integers. */
-public interface StdInteger extends StdData {
-
- /** Returns the underlying int value. */
- int get();
-}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdLong.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdLong.java
deleted file mode 100644
index 84f322f7..00000000
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdLong.java
+++ /dev/null
@@ -1,13 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.api.data;
-
-/** A Standard UDF data type for representing longs. */
-public interface StdLong extends StdData {
-
- /** Returns the underlying long value. */
- long get();
-}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdString.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdString.java
deleted file mode 100644
index 7ccd8385..00000000
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdString.java
+++ /dev/null
@@ -1,13 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.api.data;
-
-/** A Standard UDF data type for representing strings. */
-public interface StdString extends StdData {
-
- /** Returns the underlying {@link String} value. */
- String get();
-}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdTimestamp.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdTimestamp.java
deleted file mode 100644
index 18ff9bc1..00000000
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/data/StdTimestamp.java
+++ /dev/null
@@ -1,13 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.api.data;
-
-/** A Standard UDF data type for representing timestamps. */
-public interface StdTimestamp extends StdData {
-
- /** Returns the number of milliseconds elapsed from epoch for the {@link StdTimestamp}. */
- long toEpoch();
-}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdStructType.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/RowType.java
similarity index 89%
rename from transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdStructType.java
rename to transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/RowType.java
index 521ec2d7..59938208 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdStructType.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/RowType.java
@@ -9,7 +9,7 @@
/** A {@link StdType} representing a struct type. */
-public interface StdStructType extends StdType {
+public interface RowType extends StdType {
/** Returns a {@link List} of the types of all the struct fields. */
List extends StdType> fieldTypes();
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdBinaryType.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdBinaryType.java
new file mode 100644
index 00000000..5fbe53e1
--- /dev/null
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdBinaryType.java
@@ -0,0 +1,10 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.api.types;
+
+/** A {@link StdType} representing a {@link java.nio.ByteBuffer} type. */
+public interface StdBinaryType extends StdType {
+}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdDoubleType.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdDoubleType.java
new file mode 100644
index 00000000..1179729a
--- /dev/null
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdDoubleType.java
@@ -0,0 +1,10 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.api.types;
+
+/** A {@link StdType} representing a double type. */
+public interface StdDoubleType extends StdType {
+}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdFloatType.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdFloatType.java
new file mode 100644
index 00000000..d1ff9952
--- /dev/null
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/types/StdFloatType.java
@@ -0,0 +1,10 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.api.types;
+
+/** A {@link StdType} representing a float type. */
+public interface StdFloatType extends StdType {
+}
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF.java
index a4e29220..f90b91ca 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF.java
@@ -6,8 +6,6 @@
package com.linkedin.transport.api.udf;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.types.StdType;
import java.util.List;
@@ -19,8 +17,7 @@
* abstract class for UDFs expecting {@code i} arguments. Similar to lambda expressions, StdUDF(i) abstract classes are
* type-parameterized by the input types and output type of the eval function. Each class is type-parameterized by
* {@code (i+1)} type parameters; {@code i} type parameters for the UDF input types, and one type parameter for the
- * output type. All types (both input and output types) must extend the {@link StdData}
- * interface.
+ * output type.
*/
public abstract class StdUDF {
private StdFactory _stdFactory;
@@ -40,7 +37,7 @@ public abstract class StdUDF {
* of contained UDF.
*
* @param stdFactory a {@link StdFactory} object which can be used to create
- * {@link StdData} and {@link StdType} objects
+ * data and type objects
*/
public void init(StdFactory stdFactory) {
_stdFactory = stdFactory;
@@ -85,8 +82,8 @@ public final boolean[] getAndCheckNullableArguments() {
protected abstract int numberOfArguments();
/**
- * Returns a {@link StdFactory} object which can be used to create {@link StdData} and
- * {@link StdType} objects
+ * Returns a {@link StdFactory} object which can be used to create data and
+ * type objects
*/
public StdFactory getStdFactory() {
return _stdFactory;
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF0.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF0.java
index b62fe95f..d3558fc3 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF0.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF0.java
@@ -5,15 +5,13 @@
*/
package com.linkedin.transport.api.udf;
-import com.linkedin.transport.api.data.StdData;
-
/**
* A Standard UDF with zero input arguments.
*
* @param the type of the return value of the {@link StdUDF}
*/
-public abstract class StdUDF0 extends StdUDF {
+public abstract class StdUDF0 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF1.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF1.java
index 28d0ad71..18e8e769 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF1.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF1.java
@@ -5,8 +5,6 @@
*/
package com.linkedin.transport.api.udf;
-import com.linkedin.transport.api.data.StdData;
-
/**
* A Standard UDF with one input argument.
@@ -17,7 +15,7 @@
// Suppressing class parameter type parameter name and arg naming style checks since this naming convention is more
// suitable to Standard UDFs, and the code is more readable this way.
@SuppressWarnings({"checkstyle:classtypeparametername", "checkstyle:regexpsinglelinejava"})
-public abstract class StdUDF1 extends StdUDF {
+public abstract class StdUDF1 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
@@ -38,7 +36,7 @@ public abstract class StdUDF1 extends Std
* hence obtaining the most recent version of a file.
* Example: 'hdfs:///data/derived/dwh/prop/testMemberId/#LATEST/testMemberId.txt'
*
- * The arguments passed to {@link #eval(StdData)} are passed to this method as well to allow users to construct
+ * The arguments passed to {@link #eval(Object)} are passed to this method as well to allow users to construct
* required file paths from arguments passed to the UDF. Since this method is called before any rows are processed,
* only constant UDF arguments should be used to construct the file paths. Values of non-constant arguments are not
* deterministic, and are null for most platforms. (Constant arguments are arguments whose literal values are given
diff --git a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF2.java b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF2.java
index 3e020ae1..3eb293ba 100644
--- a/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF2.java
+++ b/transportable-udfs-api/src/main/java/com/linkedin/transport/api/udf/StdUDF2.java
@@ -5,8 +5,6 @@
*/
package com.linkedin.transport.api.udf;
-import com.linkedin.transport.api.data.StdData;
-
/**
* A Standard UDF with three input arguments.
@@ -18,7 +16,7 @@
// Suppressing class parameter type parameter name and arg naming style checks since this naming convention is more
// suitable to Standard UDFs, and the code is more readable this way.
@SuppressWarnings({"checkstyle:classtypeparametername", "checkstyle:regexpsinglelinejava"})
-public abstract class StdUDF2 extends StdUDF {
+public abstract class StdUDF2 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
@@ -40,7 +38,7 @@ public abstract class StdUDF2
+public abstract class StdUDF3
extends StdUDF {
/**
@@ -43,7 +41,7 @@ public abstract class StdUDF3
+public abstract class StdUDF4
extends StdUDF {
/**
@@ -45,7 +43,7 @@ public abstract class StdUDF4 extends StdUDF {
+public abstract class StdUDF5 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
@@ -47,7 +45,7 @@ public abstract class StdUDF5 extends StdUDF {
+public abstract class StdUDF6 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
@@ -49,7 +47,7 @@ public abstract class StdUDF6 extends StdUDF {
+public abstract class StdUDF7 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
@@ -51,7 +49,7 @@ public abstract class StdUDF7 extends StdUDF {
+public abstract class StdUDF8 extends StdUDF {
/**
* Returns the output of the {@link StdUDF} given the input arguments.
@@ -53,7 +51,7 @@ public abstract class StdUDF8 boundVariables) {
}
@Override
- public StdInteger createInteger(int value) {
- return new AvroInteger(value);
+ public ArrayData createArray(StdType stdType, int size) {
+ return new AvroArrayData((Schema) stdType.underlyingType(), size);
}
@Override
- public StdLong createLong(long value) {
- return new AvroLong(value);
- }
-
- @Override
- public StdBoolean createBoolean(boolean value) {
- return new AvroBoolean(value);
- }
-
- @Override
- public StdString createString(String value) {
- return new AvroString(new Utf8(value));
- }
-
- @Override
- public StdArray createArray(StdType stdType, int size) {
- return new AvroArray((Schema) stdType.underlyingType(), size);
- }
-
- @Override
- public StdArray createArray(StdType stdType) {
+ public ArrayData createArray(StdType stdType) {
return createArray(stdType, 0);
}
@Override
- public StdMap createMap(StdType stdType) {
- return new AvroMap((Schema) stdType.underlyingType());
+ public MapData createMap(StdType stdType) {
+ return new AvroMapData((Schema) stdType.underlyingType());
}
@Override
- public StdStruct createStruct(List fieldNames, List fieldTypes) {
+ public RowData createStruct(List fieldNames, List fieldTypes) {
if (fieldNames.size() != fieldTypes.size()) {
throw new RuntimeException(
"Field names and types are of different lengths: " + "Field names length is " + fieldNames.size() + ". "
@@ -90,18 +62,18 @@ public StdStruct createStruct(List fieldNames, List fieldTypes)
for (int i = 0; i < fieldTypes.size(); i++) {
fields.add(new Field(fieldNames.get(i), (Schema) fieldTypes.get(i).underlyingType(), null, null));
}
- return new AvroStruct(Schema.createRecord(fields));
+ return new AvroRowData(Schema.createRecord(fields));
}
@Override
- public StdStruct createStruct(List fieldTypes) {
+ public RowData createStruct(List fieldTypes) {
return createStruct(IntStream.range(0, fieldTypes.size()).mapToObj(i -> "field" + i).collect(Collectors.toList()),
fieldTypes);
}
@Override
- public StdStruct createStruct(StdType stdType) {
- return new AvroStruct((Schema) stdType.underlyingType());
+ public RowData createStruct(StdType stdType) {
+ return new AvroRowData((Schema) stdType.underlyingType());
}
@Override
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/AvroWrapper.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/AvroWrapper.java
index a4c65904..b493f680 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/AvroWrapper.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/AvroWrapper.java
@@ -5,22 +5,23 @@
*/
package com.linkedin.transport.avro;
-import com.linkedin.transport.api.data.StdData;
+import com.linkedin.transport.api.data.PlatformData;
import com.linkedin.transport.api.types.StdType;
-import com.linkedin.transport.avro.data.AvroArray;
-import com.linkedin.transport.avro.data.AvroBoolean;
-import com.linkedin.transport.avro.data.AvroInteger;
-import com.linkedin.transport.avro.data.AvroLong;
-import com.linkedin.transport.avro.data.AvroMap;
-import com.linkedin.transport.avro.data.AvroString;
-import com.linkedin.transport.avro.data.AvroStruct;
+import com.linkedin.transport.avro.data.AvroArrayData;
+import com.linkedin.transport.avro.data.AvroMapData;
+import com.linkedin.transport.avro.data.AvroRowData;
import com.linkedin.transport.avro.types.AvroArrayType;
+import com.linkedin.transport.avro.types.AvroBinaryType;
import com.linkedin.transport.avro.types.AvroBooleanType;
+import com.linkedin.transport.avro.types.AvroDoubleType;
+import com.linkedin.transport.avro.types.AvroFloatType;
import com.linkedin.transport.avro.types.AvroIntegerType;
import com.linkedin.transport.avro.types.AvroLongType;
import com.linkedin.transport.avro.types.AvroMapType;
import com.linkedin.transport.avro.types.AvroStringType;
-import com.linkedin.transport.avro.types.AvroStructType;
+import com.linkedin.transport.avro.types.AvroRowType;
+import java.nio.ByteBuffer;
+import java.util.List;
import java.util.Map;
import org.apache.avro.Schema;
import org.apache.avro.generic.GenericArray;
@@ -33,22 +34,35 @@ public class AvroWrapper {
private AvroWrapper() {
}
- public static StdData createStdData(Object avroData, Schema avroSchema) {
+ public static Object createStdData(Object avroData, Schema avroSchema) {
switch (avroSchema.getType()) {
case INT:
- return new AvroInteger((Integer) avroData);
case LONG:
- return new AvroLong((Long) avroData);
case BOOLEAN:
- return new AvroBoolean((Boolean) avroData);
+ case FLOAT:
+ case DOUBLE:
+ case BYTES:
+ return avroData;
case STRING:
- return new AvroString((Utf8) avroData);
+ case ENUM:
+ if (avroData == null) {
+ return null;
+ } else {
+ return avroData.toString();
+ }
case ARRAY:
- return new AvroArray((GenericArray) avroData, avroSchema);
+ return new AvroArrayData((GenericArray) avroData, avroSchema);
case MAP:
- return new AvroMap((Map) avroData, avroSchema);
+ return new AvroMapData((Map) avroData, avroSchema);
case RECORD:
- return new AvroStruct((GenericRecord) avroData, avroSchema);
+ return new AvroRowData((GenericRecord) avroData, avroSchema);
+ case UNION: {
+ Schema nonNullableType = getNonNullComponent(avroSchema);
+ if (avroData == null) {
+ return null;
+ }
+ return createStdData(avroData, nonNullableType);
+ }
case NULL:
return null;
default:
@@ -56,6 +70,36 @@ public static StdData createStdData(Object avroData, Schema avroSchema) {
}
}
+ public static Object getPlatformData(Object transportData) {
+ if (transportData instanceof Integer || transportData instanceof Long || transportData instanceof Double ||
+ transportData instanceof Boolean || transportData instanceof ByteBuffer) {
+ return transportData;
+ } else if (transportData instanceof String) {
+ return transportData == null? null : new Utf8((String) transportData);
+ } else {
+ return transportData == null ? null : ((PlatformData) transportData).getUnderlyingData();
+ }
+ }
+
+
+ /**
+ * Returns a non null component of a simple union schema. The supported union schema must have
+ * only two fields where one of them is null type, the other is returned.
+ */
+ private static Schema getNonNullComponent(Schema unionSchema) {
+ List types = unionSchema.getTypes();
+ if (types.size() == 2) {
+ if (types.get(0).getType().equals(Schema.Type.NULL)) {
+ return types.get(1);
+ }
+
+ if (types.get(1).getType().equals(Schema.Type.NULL)) {
+ return types.get(0);
+ }
+ }
+ throw new RuntimeException("Unsupported union type: " + unionSchema);
+ }
+
public static StdType createStdType(Schema avroSchema) {
switch (avroSchema.getType()) {
case INT:
@@ -66,12 +110,22 @@ public static StdType createStdType(Schema avroSchema) {
return new AvroBooleanType(avroSchema);
case STRING:
return new AvroStringType(avroSchema);
+ case FLOAT:
+ return new AvroFloatType(avroSchema);
+ case DOUBLE:
+ return new AvroDoubleType(avroSchema);
+ case BYTES:
+ return new AvroBinaryType(avroSchema);
case ARRAY:
return new AvroArrayType(avroSchema);
case MAP:
return new AvroMapType(avroSchema);
case RECORD:
- return new AvroStructType(avroSchema);
+ return new AvroRowType(avroSchema);
+ case UNION: {
+ Schema nonNullableType = getNonNullComponent(avroSchema);
+ return createStdType(nonNullableType);
+ }
default:
throw new RuntimeException("Unrecognized Avro Schema: " + avroSchema.getClass());
}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/StdUdfWrapper.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/StdUdfWrapper.java
index a1ea3e0d..41eb59b4 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/StdUdfWrapper.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/StdUdfWrapper.java
@@ -7,7 +7,6 @@
import com.linkedin.transport.api.StdFactory;
import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdData;
import com.linkedin.transport.api.udf.StdUDF;
import com.linkedin.transport.api.udf.StdUDF0;
import com.linkedin.transport.api.udf.StdUDF1;
@@ -36,7 +35,7 @@ public abstract class StdUdfWrapper {
protected boolean _requiredFilesProcessed;
protected StdFactory _stdFactory;
private boolean[] _nullableArguments;
- private StdData[] _args;
+ private Object[] _args;
/**
* Given input schemas, this method matches them to the expected type signatures, and finds bindings to the
@@ -68,12 +67,27 @@ protected boolean containsNullValuedNonNullableArgument(Object[] arguments) {
return false;
}
- protected StdData wrap(Object avroObject, StdData stdData) {
- if (avroObject != null) {
- ((PlatformData) stdData).setUnderlyingData(avroObject);
- return stdData;
- } else {
- return null;
+ protected Object wrap(Object avroObject, Schema inputSchema, Object stdData) {
+ switch (inputSchema.getType()) {
+ case INT:
+ case LONG:
+ case BOOLEAN:
+ return avroObject;
+ case STRING:
+ return avroObject == null? null : avroObject.toString();
+ case ARRAY:
+ case MAP:
+ case RECORD:
+ if (avroObject != null) {
+ ((PlatformData) stdData).setUnderlyingData(avroObject);
+ return stdData;
+ } else {
+ return null;
+ }
+ case NULL:
+ return null;
+ default:
+ throw new RuntimeException("Unrecognized Avro Schema: " + inputSchema.getClass());
}
}
@@ -82,22 +96,24 @@ protected StdData wrap(Object avroObject, StdData stdData) {
protected abstract Class extends TopLevelStdUDF> getTopLevelUdfClass();
protected void createStdData() {
- _args = new StdData[_inputSchemas.length];
+ _args = new Object[_inputSchemas.length];
for (int i = 0; i < _inputSchemas.length; i++) {
_args[i] = AvroWrapper.createStdData(null, _inputSchemas[i]);
}
}
- private StdData[] wrapArguments(Object[] arguments) {
- return IntStream.range(0, _args.length).mapToObj(i -> wrap(arguments[i], _args[i])).toArray(StdData[]::new);
+ private Object[] wrapArguments(Object[] arguments) {
+ return IntStream.range(0, _args.length).mapToObj(
+ i -> wrap(arguments[i], _inputSchemas[i], _args[i])
+ ).toArray(Object[]::new);
}
public Object evaluate(Object[] arguments) {
if (containsNullValuedNonNullableArgument(arguments)) {
return null;
}
- StdData[] args = wrapArguments(arguments);
- StdData result;
+ Object[] args = wrapArguments(arguments);
+ Object result;
switch (args.length) {
case 0:
result = ((StdUDF0) _stdUdf).eval();
@@ -129,6 +145,6 @@ public Object evaluate(Object[] arguments) {
default:
throw new UnsupportedOperationException("eval not yet supported for StdUDF" + args.length);
}
- return result == null ? null : ((PlatformData) result).getUnderlyingData();
+ return AvroWrapper.getPlatformData(result);
}
}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroArray.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroArrayData.java
similarity index 65%
rename from transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroArray.java
rename to transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroArrayData.java
index 1557ed6c..3cb2e6f6 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroArray.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroArrayData.java
@@ -6,8 +6,7 @@
package com.linkedin.transport.avro.data;
import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdData;
+import com.linkedin.transport.api.data.ArrayData;
import com.linkedin.transport.avro.AvroWrapper;
import java.util.Iterator;
import org.apache.avro.Schema;
@@ -15,16 +14,16 @@
import org.apache.avro.generic.GenericData;
-public class AvroArray implements StdArray, PlatformData {
+public class AvroArrayData implements ArrayData, PlatformData {
private final Schema _elementSchema;
private GenericArray _genericArray;
- public AvroArray(GenericArray genericArray, Schema arraySchema) {
+ public AvroArrayData(GenericArray genericArray, Schema arraySchema) {
_genericArray = genericArray;
_elementSchema = arraySchema.getElementType();
}
- public AvroArray(Schema arraySchema, int size) {
+ public AvroArrayData(Schema arraySchema, int size) {
_elementSchema = arraySchema.getElementType();
_genericArray = new GenericData.Array(size, arraySchema);
}
@@ -35,18 +34,18 @@ public int size() {
}
@Override
- public StdData get(int idx) {
- return AvroWrapper.createStdData(_genericArray.get(idx), _elementSchema);
+ public E get(int idx) {
+ return (E) AvroWrapper.createStdData(_genericArray.get(idx), _elementSchema);
}
@Override
- public void add(StdData e) {
- _genericArray.add(((PlatformData) e).getUnderlyingData());
+ public void add(E e) {
+ _genericArray.add(AvroWrapper.getPlatformData(e));
}
@Override
- public Iterator iterator() {
- return new Iterator() {
+ public Iterator iterator() {
+ return new Iterator() {
private final Iterator _iterator = _genericArray.iterator();
@Override
@@ -55,8 +54,8 @@ public boolean hasNext() {
}
@Override
- public StdData next() {
- return AvroWrapper.createStdData(_iterator.next(), _elementSchema);
+ public E next() {
+ return (E) AvroWrapper.createStdData(_iterator.next(), _elementSchema);
}
};
}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroBoolean.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroBoolean.java
deleted file mode 100644
index 99f83738..00000000
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroBoolean.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.avro.data;
-
-import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdBoolean;
-
-
-public class AvroBoolean implements StdBoolean, PlatformData {
- private Boolean _boolean;
-
- public AvroBoolean(Boolean aBoolean) {
- _boolean = aBoolean;
- }
-
- @Override
- public boolean get() {
- return _boolean;
- }
-
- @Override
- public Object getUnderlyingData() {
- return _boolean;
- }
-
- @Override
- public void setUnderlyingData(Object value) {
- _boolean = (Boolean) value;
- }
-}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroInteger.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroInteger.java
deleted file mode 100644
index 5a170f3b..00000000
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroInteger.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.avro.data;
-
-import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdInteger;
-
-
-public class AvroInteger implements StdInteger, PlatformData {
- private Integer _integer;
-
- public AvroInteger(Integer integer) {
- _integer = integer;
- }
-
- @Override
- public int get() {
- return _integer;
- }
-
- @Override
- public Object getUnderlyingData() {
- return _integer;
- }
-
- @Override
- public void setUnderlyingData(Object value) {
- _integer = (Integer) value;
- }
-}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroLong.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroLong.java
deleted file mode 100644
index a56af06c..00000000
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroLong.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.avro.data;
-
-import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdLong;
-
-
-public class AvroLong implements StdLong, PlatformData {
- private Long _long;
-
- public AvroLong(Long aLong) {
- _long = aLong;
- }
-
- @Override
- public long get() {
- return _long;
- }
-
- @Override
- public Object getUnderlyingData() {
- return _long;
- }
-
- @Override
- public void setUnderlyingData(Object value) {
- _long = (Long) value;
- }
-}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroMap.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroMapData.java
similarity index 59%
rename from transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroMap.java
rename to transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroMapData.java
index d0913d53..2b95796c 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroMap.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroMapData.java
@@ -6,8 +6,7 @@
package com.linkedin.transport.avro.data;
import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdMap;
+import com.linkedin.transport.api.data.MapData;
import com.linkedin.transport.avro.AvroWrapper;
import java.util.AbstractSet;
import java.util.Collection;
@@ -21,18 +20,18 @@
import static org.apache.avro.Schema.Type.*;
-public class AvroMap implements StdMap, PlatformData {
+public class AvroMapData implements MapData, PlatformData {
private Map _map;
private final Schema _keySchema;
private final Schema _valueSchema;
- public AvroMap(Map map, Schema mapSchema) {
+ public AvroMapData(Map map, Schema mapSchema) {
_map = map;
_keySchema = Schema.create(STRING);
_valueSchema = mapSchema.getValueType();
}
- public AvroMap(Schema mapSchema) {
+ public AvroMapData(Schema mapSchema) {
_map = new LinkedHashMap<>();
_keySchema = Schema.create(STRING);
_valueSchema = mapSchema.getValueType();
@@ -54,21 +53,21 @@ public int size() {
}
@Override
- public StdData get(StdData key) {
- return AvroWrapper.createStdData(_map.get(((PlatformData) key).getUnderlyingData()), _valueSchema);
+ public V get(K key) {
+ return (V) AvroWrapper.createStdData(_map.get(AvroWrapper.getPlatformData(key)), _valueSchema);
}
@Override
- public void put(StdData key, StdData value) {
- _map.put(((PlatformData) key).getUnderlyingData(), ((PlatformData) value).getUnderlyingData());
+ public void put(K key, V value) {
+ _map.put(AvroWrapper.getPlatformData(key), AvroWrapper.getPlatformData(value));
}
@Override
- public Set keySet() {
- return new AbstractSet() {
+ public Set keySet() {
+ return new AbstractSet() {
@Override
- public Iterator iterator() {
- return new Iterator() {
+ public Iterator iterator() {
+ return new Iterator() {
Iterator keySet = _map.keySet().iterator();
@Override
public boolean hasNext() {
@@ -76,8 +75,8 @@ public boolean hasNext() {
}
@Override
- public StdData next() {
- return AvroWrapper.createStdData(keySet.next(), _keySchema);
+ public K next() {
+ return (K) AvroWrapper.createStdData(keySet.next(), _keySchema);
}
};
}
@@ -90,12 +89,12 @@ public int size() {
}
@Override
- public Collection values() {
- return _map.values().stream().map(v -> AvroWrapper.createStdData(v, _valueSchema)).collect(Collectors.toList());
+ public Collection values() {
+ return _map.values().stream().map(v -> (V) AvroWrapper.createStdData(v, _valueSchema)).collect(Collectors.toList());
}
@Override
- public boolean containsKey(StdData key) {
- return _map.containsKey(((PlatformData) key).getUnderlyingData());
+ public boolean containsKey(K key) {
+ return _map.containsKey(AvroWrapper.getPlatformData(key));
}
}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroStruct.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroRowData.java
similarity index 68%
rename from transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroStruct.java
rename to transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroRowData.java
index f018d5bc..64fa5e4c 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroStruct.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroRowData.java
@@ -6,8 +6,7 @@
package com.linkedin.transport.avro.data;
import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.avro.AvroWrapper;
import java.util.List;
import java.util.stream.Collectors;
@@ -17,17 +16,17 @@
import org.apache.avro.generic.GenericRecord;
-public class AvroStruct implements StdStruct, PlatformData {
+public class AvroRowData implements RowData, PlatformData {
private final Schema _recordSchema;
private GenericRecord _genericRecord;
- public AvroStruct(GenericRecord genericRecord, Schema recordSchema) {
+ public AvroRowData(GenericRecord genericRecord, Schema recordSchema) {
_genericRecord = genericRecord;
_recordSchema = recordSchema;
}
- public AvroStruct(Schema recordSchema) {
+ public AvroRowData(Schema recordSchema) {
_genericRecord = new Record(recordSchema);
_recordSchema = recordSchema;
}
@@ -43,27 +42,27 @@ public void setUnderlyingData(Object value) {
}
@Override
- public StdData getField(int index) {
+ public Object getField(int index) {
return AvroWrapper.createStdData(_genericRecord.get(index), _recordSchema.getFields().get(index).schema());
}
@Override
- public StdData getField(String name) {
+ public Object getField(String name) {
return AvroWrapper.createStdData(_genericRecord.get(name), _recordSchema.getField(name).schema());
}
@Override
- public void setField(int index, StdData value) {
- _genericRecord.put(index, ((PlatformData) value).getUnderlyingData());
+ public void setField(int index, Object value) {
+ _genericRecord.put(index, AvroWrapper.getPlatformData(value));
}
@Override
- public void setField(String name, StdData value) {
- _genericRecord.put(name, ((PlatformData) value).getUnderlyingData());
+ public void setField(String name, Object value) {
+ _genericRecord.put(name, AvroWrapper.getPlatformData(value));
}
@Override
- public List fields() {
+ public List fields() {
return IntStream.range(0, _recordSchema.getFields().size()).mapToObj(i -> getField(i)).collect(Collectors.toList());
}
}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroString.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroString.java
deleted file mode 100644
index 745df05e..00000000
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/data/AvroString.java
+++ /dev/null
@@ -1,34 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.avro.data;
-
-import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdString;
-import org.apache.avro.util.Utf8;
-
-
-public class AvroString implements StdString, PlatformData {
- private Utf8 _string;
-
- public AvroString(Utf8 string) {
- _string = string;
- }
-
- @Override
- public String get() {
- return _string.toString();
- }
-
- @Override
- public Object getUnderlyingData() {
- return _string;
- }
-
- @Override
- public void setUnderlyingData(Object value) {
- _string = (Utf8) value;
- }
-}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroBinaryType.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroBinaryType.java
new file mode 100644
index 00000000..883d37bc
--- /dev/null
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroBinaryType.java
@@ -0,0 +1,23 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.avro.types;
+
+import com.linkedin.transport.api.types.StdBinaryType;
+import org.apache.avro.Schema;
+
+
+public class AvroBinaryType implements StdBinaryType {
+ final private Schema _schema;
+
+ public AvroBinaryType(Schema schema) {
+ _schema = schema;
+ }
+
+ @Override
+ public Object underlyingType() {
+ return _schema;
+ }
+}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroDoubleType.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroDoubleType.java
new file mode 100644
index 00000000..fe9b847d
--- /dev/null
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroDoubleType.java
@@ -0,0 +1,23 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.avro.types;
+
+import com.linkedin.transport.api.types.StdDoubleType;
+import org.apache.avro.Schema;
+
+
+public class AvroDoubleType implements StdDoubleType {
+ final private Schema _schema;
+
+ public AvroDoubleType(Schema schema) {
+ _schema = schema;
+ }
+
+ @Override
+ public Object underlyingType() {
+ return _schema;
+ }
+}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroFloatType.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroFloatType.java
new file mode 100644
index 00000000..c277fd54
--- /dev/null
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroFloatType.java
@@ -0,0 +1,23 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.avro.types;
+
+import com.linkedin.transport.api.types.StdFloatType;
+import org.apache.avro.Schema;
+
+
+public class AvroFloatType implements StdFloatType {
+ final private Schema _schema;
+
+ public AvroFloatType(Schema schema) {
+ _schema = schema;
+ }
+
+ @Override
+ public Object underlyingType() {
+ return _schema;
+ }
+}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroStructType.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroRowType.java
similarity index 82%
rename from transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroStructType.java
rename to transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroRowType.java
index 2c97b39c..3923b3f5 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroStructType.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/types/AvroRowType.java
@@ -5,7 +5,7 @@
*/
package com.linkedin.transport.avro.types;
-import com.linkedin.transport.api.types.StdStructType;
+import com.linkedin.transport.api.types.RowType;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.avro.AvroWrapper;
import java.util.List;
@@ -13,10 +13,10 @@
import org.apache.avro.Schema;
-public class AvroStructType implements StdStructType {
+public class AvroRowType implements RowType {
final private Schema _schema;
- public AvroStructType(Schema schema) {
+ public AvroRowType(Schema schema) {
_schema = schema;
}
diff --git a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/typesystem/AvroTypeSystem.java b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/typesystem/AvroTypeSystem.java
index 881fa709..f85e75d7 100644
--- a/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/typesystem/AvroTypeSystem.java
+++ b/transportable-udfs-avro/src/main/java/com/linkedin/transport/avro/typesystem/AvroTypeSystem.java
@@ -60,6 +60,21 @@ protected boolean isStringType(Schema dataType) {
return dataType.getType() == STRING;
}
+ @Override
+ protected boolean isFloatType(Schema dataType) {
+ return dataType.getType() == FLOAT;
+ }
+
+ @Override
+ protected boolean isDoubleType(Schema dataType) {
+ return dataType.getType() == DOUBLE;
+ }
+
+ @Override
+ protected boolean isBinaryType(Schema dataType) {
+ return dataType.getType() == BYTES;
+ }
+
@Override
protected boolean isArrayType(Schema dataType) {
return dataType.getType() == ARRAY;
@@ -95,6 +110,21 @@ protected Schema createStringType() {
return Schema.create(STRING);
}
+ @Override
+ protected Schema createFloatType() {
+ return Schema.create(FLOAT);
+ }
+
+ @Override
+ protected Schema createDoubleType() {
+ return Schema.create(DOUBLE);
+ }
+
+ @Override
+ protected Schema createBinaryType() {
+ return Schema.create(BYTES);
+ }
+
@Override
protected Schema createUnknownType() {
return Schema.create(NULL);
diff --git a/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/TestAvroWrapper.java b/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/TestAvroWrapper.java
new file mode 100644
index 00000000..d81a452a
--- /dev/null
+++ b/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/TestAvroWrapper.java
@@ -0,0 +1,264 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.avro;
+
+import com.google.common.collect.ImmutableList;
+import com.google.common.collect.ImmutableMap;
+import com.linkedin.transport.api.data.PlatformData;
+import com.linkedin.transport.api.data.StdData;
+import com.linkedin.transport.api.types.StdType;
+import com.linkedin.transport.avro.data.AvroArray;
+import com.linkedin.transport.avro.data.AvroBinary;
+import com.linkedin.transport.avro.data.AvroBoolean;
+import com.linkedin.transport.avro.data.AvroDouble;
+import com.linkedin.transport.avro.data.AvroFloat;
+import com.linkedin.transport.avro.data.AvroInteger;
+import com.linkedin.transport.avro.data.AvroLong;
+import com.linkedin.transport.avro.data.AvroMap;
+import com.linkedin.transport.avro.data.AvroString;
+import com.linkedin.transport.avro.data.AvroStruct;
+import com.linkedin.transport.avro.types.AvroArrayType;
+import com.linkedin.transport.avro.types.AvroBinaryType;
+import com.linkedin.transport.avro.types.AvroBooleanType;
+import com.linkedin.transport.avro.types.AvroDoubleType;
+import com.linkedin.transport.avro.types.AvroFloatType;
+import com.linkedin.transport.avro.types.AvroIntegerType;
+import com.linkedin.transport.avro.types.AvroLongType;
+import com.linkedin.transport.avro.types.AvroMapType;
+import com.linkedin.transport.avro.types.AvroStringType;
+import com.linkedin.transport.avro.types.AvroStructType;
+import java.nio.ByteBuffer;
+import java.util.Arrays;
+import java.util.Map;
+import org.apache.avro.Schema;
+import org.apache.avro.generic.GenericArray;
+import org.apache.avro.generic.GenericData;
+import org.apache.avro.generic.GenericRecord;
+import org.apache.avro.util.Utf8;
+import org.testng.annotations.Test;
+
+import static org.testng.Assert.*;
+
+
+public class TestAvroWrapper {
+
+ private Schema createSchema(String typeName) {
+ return createSchema("testField", typeName);
+ }
+
+ private Schema createSchema(String fieldName, String typeName) {
+ return new Schema.Parser().parse(
+ String.format("{\"name\": \"%s\",\"type\": %s}", fieldName, typeName));
+ }
+
+ private void testSimpleType(String typeName, Class extends StdType> expectedAvroTypeClass,
+ Object testData, Class extends StdData> expectedDataClass) {
+ Schema avroSchema = createSchema(String.format("\"%s\"", typeName));
+
+ StdType stdType = AvroWrapper.createStdType(avroSchema);
+ assertTrue(expectedAvroTypeClass.isAssignableFrom(stdType.getClass()));
+ assertEquals(avroSchema, stdType.underlyingType());
+
+ StdData stdData = AvroWrapper.createStdData(testData, avroSchema);
+ assertNotNull(stdData);
+ assertTrue(expectedDataClass.isAssignableFrom(stdData.getClass()));
+ if ("string".equals(typeName)) {
+ // Use String values for equality assertion as we support both Utf8 and String input types
+ assertEquals(testData.toString(), ((PlatformData) stdData).getUnderlyingData().toString());
+ } else {
+ assertEquals(testData, ((PlatformData) stdData).getUnderlyingData());
+ }
+ }
+
+ @Test
+ public void testBooleanType() {
+ testSimpleType("boolean", AvroBooleanType.class, true, AvroBoolean.class);
+ }
+
+ @Test
+ public void testIntegerType() {
+ testSimpleType("int", AvroIntegerType.class, 1, AvroInteger.class);
+ }
+
+ @Test
+ public void testLongType() {
+ testSimpleType("long", AvroLongType.class, 1L, AvroLong.class);
+ }
+
+ @Test
+ public void testFloatType() {
+ testSimpleType("float", AvroFloatType.class, 1.0f, AvroFloat.class);
+ }
+
+ @Test
+ public void testDoubleType() {
+ testSimpleType("double", AvroDoubleType.class, 1.0, AvroDouble.class);
+ }
+
+ @Test
+ public void testStringType() {
+ testSimpleType("string", AvroStringType.class, new Utf8("foo"), AvroString.class);
+ testSimpleType("string", AvroStringType.class, "foo", AvroString.class);
+ }
+
+ @Test
+ public void testBinaryType() {
+ testSimpleType("bytes", AvroBinaryType.class, ByteBuffer.wrap("bar".getBytes()), AvroBinary.class);
+ }
+
+ @Test
+ public void testEnumType() {
+ Schema field1 = createSchema("field1", ""
+ + "\"enum\","
+ + "\"name\":\"SampleEnum\","
+ + "\"doc\":\"\","
+ + "\"symbols\":[\"A\",\"B\"]");
+ Schema structSchema = Schema.createRecord(ImmutableList.of(
+ new Schema.Field("field1", field1, null, null)
+ ));
+
+ GenericRecord record1 = new GenericData.Record(structSchema);
+ record1.put("field1", "A");
+ StdData stdEnumData1 = AvroWrapper.createStdData(record1.get("field1"),
+ Schema.createEnum("SampleEnum", "", "", Arrays.asList("A", "B")));
+ assertTrue(stdEnumData1 instanceof AvroString);
+ assertEquals("A", ((AvroString) stdEnumData1).get());
+
+ GenericRecord record2 = new GenericData.Record(structSchema);
+ record1.put("field1", new GenericData.EnumSymbol(field1, "A"));
+ StdData stdEnumData2 = AvroWrapper.createStdData(record1.get("field1"),
+ Schema.createEnum("SampleEnum", "", "", Arrays.asList("A", "B")));
+ assertTrue(stdEnumData2 instanceof AvroString);
+ assertEquals("A", ((AvroString) stdEnumData2).get());
+ }
+
+ @Test
+ public void testArrayType() {
+ Schema elementType = createSchema("\"int\"");
+ Schema arraySchema = Schema.createArray(elementType);
+
+ StdType stdArrayType = AvroWrapper.createStdType(arraySchema);
+ assertTrue(stdArrayType instanceof AvroArrayType);
+ assertEquals(arraySchema, stdArrayType.underlyingType());
+ assertEquals(elementType, ((AvroArrayType) stdArrayType).elementType().underlyingType());
+
+ GenericArray value = new GenericData.Array<>(arraySchema, Arrays.asList(1, 2));
+ StdData stdArrayData = AvroWrapper.createStdData(value, arraySchema);
+ assertTrue(stdArrayData instanceof AvroArray);
+ assertEquals(2, ((AvroArray) stdArrayData).size());
+ assertEquals(value, ((AvroArray) stdArrayData).getUnderlyingData());
+ }
+
+ @Test
+ public void testMapType() {
+ Schema valueType = createSchema("\"long\"");
+ Schema mapSchema = Schema.createMap(valueType);
+
+ StdType stdMapType = AvroWrapper.createStdType(mapSchema);
+ assertTrue(stdMapType instanceof AvroMapType);
+ assertEquals(mapSchema, stdMapType.underlyingType());
+ assertEquals(valueType, ((AvroMapType) stdMapType).valueType().underlyingType());
+
+ Map value = ImmutableMap.of("foo", 1L, "bar", 2L);
+ StdData stdMapData = AvroWrapper.createStdData(value, mapSchema);
+ assertTrue(stdMapData instanceof AvroMap);
+ assertEquals(2, ((AvroMap) stdMapData).size());
+ assertEquals(value, ((AvroMap) stdMapData).getUnderlyingData());
+ }
+
+ @Test
+ public void testRecordType() {
+ Schema field1 = createSchema("field1", "\"int\"");
+ Schema field2 = createSchema("field2", "\"double\"");
+ Schema structSchema = Schema.createRecord(ImmutableList.of(
+ new Schema.Field("field1", field1, null, null),
+ new Schema.Field("field2", field2, null, null)
+ ));
+
+ StdType stdStructType = AvroWrapper.createStdType(structSchema);
+ assertTrue(stdStructType instanceof AvroStructType);
+ assertEquals(structSchema, stdStructType.underlyingType());
+ assertEquals(field1, ((AvroStructType) stdStructType).fieldTypes().get(0).underlyingType());
+ assertEquals(field2, ((AvroStructType) stdStructType).fieldTypes().get(1).underlyingType());
+
+ GenericRecord value = new GenericData.Record(structSchema);
+ value.put("field1", 1);
+ value.put("field2", 2.0);
+ StdData stdStructData = AvroWrapper.createStdData(value, structSchema);
+ assertTrue(stdStructData instanceof AvroStruct);
+ AvroStruct avroStruct = (AvroStruct) stdStructData;
+ assertEquals(2, avroStruct.fields().size());
+ assertEquals(value, avroStruct.getUnderlyingData());
+ assertEquals(1, ((PlatformData) avroStruct.getField("field1")).getUnderlyingData());
+ assertEquals(2.0, ((PlatformData) avroStruct.getField("field2")).getUnderlyingData());
+ }
+
+ @Test
+ public void testValidUnionType() {
+ Schema nonNullType = createSchema("\"long\"");
+ Schema unionSchema = Schema.createUnion(Arrays.asList(nonNullType, Schema.create(Schema.Type.NULL)));
+
+ StdType stdLongType = AvroWrapper.createStdType(unionSchema);
+ assertTrue(stdLongType instanceof AvroLongType);
+ assertEquals(nonNullType, stdLongType.underlyingType());
+
+ StdData stdLongData = AvroWrapper.createStdData(1L, unionSchema);
+ assertTrue(stdLongData instanceof AvroLong);
+ assertEquals(1L, ((AvroLong) stdLongData).get());
+
+ StdData stdNullData = AvroWrapper.createStdData(null, unionSchema);
+ assertNull(stdNullData);
+ }
+
+ @Test(expectedExceptions = RuntimeException.class)
+ public void testInvalidUnionType1() {
+ Schema nonNullType = createSchema("\"long\"");
+ Schema unionSchema = Schema.createUnion(Arrays.asList(nonNullType));
+ AvroWrapper.createStdType(unionSchema);
+ }
+
+ @Test(expectedExceptions = RuntimeException.class)
+ public void testInvalidUnionType2() {
+ Schema nonNullType1 = createSchema("\"long\"");
+ Schema nonNullType2 = createSchema("\"int\"");
+ Schema unionSchema = Schema.createUnion(Arrays.asList(nonNullType1, nonNullType2));
+ AvroWrapper.createStdData(1L, unionSchema);
+ }
+
+ @Test
+ public void testStructWithSimpleUnionField() {
+ Schema field1 = createSchema("field1", "\"int\"");
+ Schema nonNullableField2 = createSchema("field2", "\"double\"");
+ Schema field2 = Schema.createUnion(Arrays.asList(Schema.create(Schema.Type.NULL), nonNullableField2));
+
+ Schema structSchema = Schema.createRecord(ImmutableList.of(
+ new Schema.Field("field1", field1, null, null),
+ new Schema.Field("field2", field2, null, null)
+ ));
+
+ GenericRecord record1 = new GenericData.Record(structSchema);
+ record1.put("field1", 1);
+ record1.put("field2", 3.0);
+ AvroStruct avroStruct1 = (AvroStruct) AvroWrapper.createStdData(record1, structSchema);
+ assertEquals(2, avroStruct1.fields().size());
+ assertEquals(3.0, ((PlatformData) avroStruct1.getField("field2")).getUnderlyingData());
+
+ GenericRecord record2 = new GenericData.Record(structSchema);
+ record2.put("field1", 1);
+ record2.put("field2", null);
+ AvroStruct avroStruct2 = (AvroStruct) AvroWrapper.createStdData(record2, structSchema);
+ assertEquals(2, avroStruct2.fields().size());
+ assertNull(avroStruct2.getField("field2"));
+ assertNull(avroStruct2.fields().get(1));
+
+ GenericRecord record3 = new GenericData.Record(structSchema);
+ record3.put("field1", 1);
+ AvroStruct avroStruct3 = (AvroStruct) AvroWrapper.createStdData(record3, structSchema);
+ assertEquals(2, avroStruct3.fields().size());
+ assertNull(avroStruct3.getField("field2"));
+ assertNull(avroStruct3.fields().get(1));
+ }
+}
diff --git a/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/typesystem/TestAvroBoundVariables.java b/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/typesystem/TestAvroBoundVariables.java
index 523d1d2a..9d2e3279 100644
--- a/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/typesystem/TestAvroBoundVariables.java
+++ b/transportable-udfs-avro/src/test/java/com/linkedin/transport/avro/typesystem/TestAvroBoundVariables.java
@@ -9,8 +9,10 @@
import com.linkedin.transport.typesystem.AbstractTestBoundVariables;
import com.linkedin.transport.typesystem.AbstractTypeSystem;
import org.apache.avro.Schema;
+import org.testng.annotations.Test;
+@Test
public class TestAvroBoundVariables extends AbstractTestBoundVariables {
@Override
diff --git a/transportable-udfs-codegen/src/main/java/com/linkedin/transport/codegen/SparkWrapperGenerator.java b/transportable-udfs-codegen/src/main/java/com/linkedin/transport/codegen/SparkWrapperGenerator.java
index bec08d8e..1bc19ded 100644
--- a/transportable-udfs-codegen/src/main/java/com/linkedin/transport/codegen/SparkWrapperGenerator.java
+++ b/transportable-udfs-codegen/src/main/java/com/linkedin/transport/codegen/SparkWrapperGenerator.java
@@ -11,7 +11,9 @@
import java.io.File;
import java.io.IOException;
import java.io.InputStream;
+import java.util.Arrays;
import java.util.Collection;
+import java.util.Map;
import java.util.stream.Collectors;
import org.apache.commons.io.IOUtils;
import org.apache.commons.text.StringSubstitutor;
@@ -23,6 +25,7 @@ public class SparkWrapperGenerator implements WrapperGenerator {
private static final String SPARK_WRAPPER_TEMPLATE_RESOURCE_PATH = "wrapper-templates/spark";
private static final String SUBSTITUTOR_KEY_WRAPPER_PACKAGE = "wrapperPackage";
private static final String SUBSTITUTOR_KEY_WRAPPER_CLASS = "wrapperClass";
+ private static final String SUBSTITUTOR_KEY_WRAPPER_CLASS_PARAMERTERS = "wrapperClassParameters";
private static final String SUBSTITUTOR_KEY_UDF_TOP_LEVEL_CLASS = "udfTopLevelClass";
private static final String SUBSTITUTOR_KEY_UDF_IMPLEMENTATIONS = "udfImplementations";
@@ -30,12 +33,16 @@ public class SparkWrapperGenerator implements WrapperGenerator {
public void generateWrappers(WrapperGeneratorContext context) {
TransportUDFMetadata udfMetadata = context.getTransportUdfMetadata();
for (String topLevelClass : udfMetadata.getTopLevelClasses()) {
- generateWrapper(topLevelClass, udfMetadata.getStdUDFImplementations(topLevelClass),
+ generateWrapper(
+ topLevelClass,
+ udfMetadata.getStdUDFImplementations(topLevelClass),
+ udfMetadata.getClassToNumberOfTypeParameters(),
context.getSourcesOutputDir());
}
}
- private void generateWrapper(String topLevelClass, Collection implementationClasses, File outputDir) {
+ private void generateWrapper(String topLevelClass, Collection implementationClasses,
+ Map classToNumberOfTypeParameters, File outputDir) {
final String wrapperTemplate;
try (InputStream wrapperTemplateStream = Thread.currentThread()
.getContextClassLoader()
@@ -49,13 +56,15 @@ private void generateWrapper(String topLevelClass, Collection implementa
ClassName wrapperClass =
ClassName.get(topLevelClassName.packageName() + "." + SPARK_PACKAGE_SUFFIX, topLevelClassName.simpleName());
String udfImplementationInstantiations = implementationClasses.stream()
- .map(clazz -> "new " + clazz + "()")
+ .map(clazz -> "new " + clazz + parameters(clazz, classToNumberOfTypeParameters) + "()")
.collect(Collectors.joining(", "));
+ String topLevelClassNameString = topLevelClassName.toString();
ImmutableMap substitutionMap = ImmutableMap.of(
SUBSTITUTOR_KEY_WRAPPER_PACKAGE, wrapperClass.packageName(),
SUBSTITUTOR_KEY_WRAPPER_CLASS, wrapperClass.simpleName(),
- SUBSTITUTOR_KEY_UDF_TOP_LEVEL_CLASS, topLevelClassName.toString(),
+ SUBSTITUTOR_KEY_UDF_TOP_LEVEL_CLASS, topLevelClassNameString
+ + parameters(topLevelClassNameString, classToNumberOfTypeParameters),
SUBSTITUTOR_KEY_UDF_IMPLEMENTATIONS, udfImplementationInstantiations
);
@@ -69,4 +78,11 @@ private void generateWrapper(String topLevelClass, Collection implementa
throw new RuntimeException("Error writing wrapper to file", e);
}
}
+
+ private static String parameters(String clazz, Map classToNumberOfTypeParameters) {
+ int numberOfTypeParameters = classToNumberOfTypeParameters.get(clazz);
+ String[] objectTypes = new String[numberOfTypeParameters];
+ Arrays.fill(objectTypes, "Object");
+ return numberOfTypeParameters > 0 ? "[" + String.join(", ", objectTypes) + "]" : "";
+ }
}
diff --git a/transportable-udfs-codegen/src/test/resources/inputs/sample-udf-metadata.json b/transportable-udfs-codegen/src/test/resources/inputs/sample-udf-metadata.json
index 4d6fb8ae..6da584f2 100644
--- a/transportable-udfs-codegen/src/test/resources/inputs/sample-udf-metadata.json
+++ b/transportable-udfs-codegen/src/test/resources/inputs/sample-udf-metadata.json
@@ -1,17 +1,17 @@
{
- "udfs": [
- {
- "topLevelClass": "udfs.OverloadedUDF",
- "stdUDFImplementations": [
- "udfs.OverloadedUDFInt",
- "udfs.OverloadedUDFString"
- ]
- },
- {
- "topLevelClass": "udfs.SimpleUDF",
- "stdUDFImplementations": [
- "udfs.SimpleUDF"
- ]
- }
- ]
-}
\ No newline at end of file
+ "udfs": {
+ "udfs.OverloadedUDF": [
+ "udfs.OverloadedUDFInt",
+ "udfs.OverloadedUDFString"
+ ],
+ "udfs.SimpleUDF": [
+ "udfs.SimpleUDF"
+ ]
+ },
+ "classToNumberOfTypeParameters": {
+ "udfs.OverloadedUDFString": 0,
+ "udfs.OverloadedUDF": 0,
+ "udfs.OverloadedUDFInt": 0,
+ "udfs.SimpleUDF": 0
+ }
+}
diff --git a/transportable-udfs-compile-utils/src/main/java/com/linkedin/transport/compile/TransportUDFMetadata.java b/transportable-udfs-compile-utils/src/main/java/com/linkedin/transport/compile/TransportUDFMetadata.java
index 48db80f5..c10cc44b 100644
--- a/transportable-udfs-compile-utils/src/main/java/com/linkedin/transport/compile/TransportUDFMetadata.java
+++ b/transportable-udfs-compile-utils/src/main/java/com/linkedin/transport/compile/TransportUDFMetadata.java
@@ -9,14 +9,18 @@
import com.google.common.collect.Multimap;
import com.google.gson.Gson;
import com.google.gson.GsonBuilder;
+import com.google.gson.JsonArray;
+import com.google.gson.JsonElement;
+import com.google.gson.JsonObject;
+import com.google.gson.JsonParser;
import java.io.File;
import java.io.FileReader;
import java.io.IOException;
import java.io.Reader;
import java.io.Writer;
import java.util.Collection;
-import java.util.LinkedList;
-import java.util.List;
+import java.util.HashMap;
+import java.util.Map;
import java.util.Set;
@@ -26,6 +30,7 @@
public class TransportUDFMetadata {
private static final Gson GSON;
private Multimap _udfs;
+ private Map _classToNumberOfTypeParameters;
static {
GSON = new GsonBuilder().setPrettyPrinting().create();
@@ -33,14 +38,15 @@ public class TransportUDFMetadata {
public TransportUDFMetadata() {
_udfs = LinkedHashMultimap.create();
+ _classToNumberOfTypeParameters = new HashMap<>();
}
public void addUDF(String topLevelClass, String stdUDFImplementation) {
_udfs.put(topLevelClass, stdUDFImplementation);
}
- public void addUDF(String topLevelClass, Collection stdUDFImplementations) {
- _udfs.putAll(topLevelClass, stdUDFImplementations);
+ public void setClassNumberOfTypeParameters(String clazz, int numberOfTypeParameters) {
+ _classToNumberOfTypeParameters.put(clazz, numberOfTypeParameters);
}
public Set getTopLevelClasses() {
@@ -51,8 +57,12 @@ public Collection getStdUDFImplementations(String topLevelClass) {
return _udfs.get(topLevelClass);
}
+ public Map getClassToNumberOfTypeParameters() {
+ return _classToNumberOfTypeParameters;
+ }
+
public void toJson(Writer writer) {
- GSON.toJson(TransportUDFMetadataSerDe.fromUDFMetadata(this), writer);
+ GSON.toJson(TransportUDFMetadataSerDe.serialize(this), writer);
}
public static TransportUDFMetadata fromJsonFile(File jsonFile) {
@@ -64,50 +74,49 @@ public static TransportUDFMetadata fromJsonFile(File jsonFile) {
}
public static TransportUDFMetadata fromJson(Reader reader) {
- return TransportUDFMetadataSerDe.toUDFMetadata(GSON.fromJson(reader, TransportUDFMetadataJson.class));
+ return TransportUDFMetadataSerDe.deserialize(new JsonParser().parse(reader));
}
- /**
- * Represents the JSON object structure of the Transport UDF metadata resource file
- */
- private static class TransportUDFMetadataJson {
- private List udfs;
+ private static class TransportUDFMetadataSerDe {
- TransportUDFMetadataJson() {
- this.udfs = new LinkedList<>();
+ public static TransportUDFMetadata deserialize(JsonElement json) {
+ TransportUDFMetadata metadata = new TransportUDFMetadata();
+ JsonObject root = json.getAsJsonObject();
+
+ // Deserialize udfs
+ JsonObject udfs = root.getAsJsonObject("udfs");
+ udfs.keySet().forEach(topLevelClass -> {
+ JsonArray stdUdfImplementations = udfs.getAsJsonArray(topLevelClass);
+ for (int i = 0; i < stdUdfImplementations.size(); i++) {
+ metadata.addUDF(topLevelClass, stdUdfImplementations.get(i).getAsString());
+ }
+ });
+
+ // Deserialize classToNumberOfTypeParameters
+ JsonObject classToNumberOfTypeParameters = root.getAsJsonObject("classToNumberOfTypeParameters");
+ classToNumberOfTypeParameters.entrySet().forEach(
+ e -> metadata.setClassNumberOfTypeParameters(e.getKey(), e.getValue().getAsInt())
+ );
+ return metadata;
}
- static class UDFInfo {
- private String topLevelClass;
- private Collection stdUDFImplementations;
-
- UDFInfo(String topLevelClass, Collection stdUDFImplementations) {
- this.topLevelClass = topLevelClass;
- this.stdUDFImplementations = stdUDFImplementations;
+ public static JsonElement serialize(TransportUDFMetadata metadata) {
+ // Serialzie _udfs
+ JsonObject udfs = new JsonObject();
+ for (Map.Entry> entry : metadata._udfs.asMap().entrySet()) {
+ JsonArray stdUdfImplementations = new JsonArray();
+ entry.getValue().forEach(f -> stdUdfImplementations.add(f));
+ udfs.add(entry.getKey(), stdUdfImplementations);
}
- }
- }
- /**
- * Converts objects between {@link TransportUDFMetadata} and {@link TransportUDFMetadataJson}
- */
- private static class TransportUDFMetadataSerDe {
-
- private static TransportUDFMetadataJson fromUDFMetadata(TransportUDFMetadata metadata) {
- TransportUDFMetadataJson metadataJson = new TransportUDFMetadataJson();
- for (String topLevelClass : metadata.getTopLevelClasses()) {
- metadataJson.udfs.add(
- new TransportUDFMetadataJson.UDFInfo(topLevelClass, metadata.getStdUDFImplementations(topLevelClass)));
- }
- return metadataJson;
- }
+ // Serialize _classToNumberOfTypeParameters
+ JsonObject classToNumberOfTypeParameters = new JsonObject();
+ metadata._classToNumberOfTypeParameters.forEach((clazz, n) -> classToNumberOfTypeParameters.addProperty(clazz, n));
- private static TransportUDFMetadata toUDFMetadata(TransportUDFMetadataJson metadataJson) {
- TransportUDFMetadata metadata = new TransportUDFMetadata();
- for (TransportUDFMetadataJson.UDFInfo udf : metadataJson.udfs) {
- metadata.addUDF(udf.topLevelClass, udf.stdUDFImplementations);
- }
- return metadata;
+ JsonObject root = new JsonObject();
+ root.add("udfs", udfs);
+ root.add("classToNumberOfTypeParameters", classToNumberOfTypeParameters);
+ return root;
}
}
}
diff --git a/transportable-udfs-examples/build.gradle b/transportable-udfs-examples/build.gradle
index 240ba4e5..8714ba39 100644
--- a/transportable-udfs-examples/build.gradle
+++ b/transportable-udfs-examples/build.gradle
@@ -33,8 +33,8 @@ subprojects {
url "https://conjars.org/repo"
}
}
- project.ext.setProperty('presto-version', '319')
- project.ext.setProperty('airlift-slice-version', '0.33')
+ project.ext.setProperty('presto-version', '333')
+ project.ext.setProperty('airlift-slice-version', '0.38')
project.ext.setProperty('spark-group', 'org.apache.spark')
project.ext.setProperty('spark-version', '2.3.0')
}
@@ -62,7 +62,9 @@ subprojects {
}
checkstyle {
- configFile = file("${rootDir}/../gradle/checkstyle/linkedin-checkstyle.xml")
+ configFile = file("${rootDir}/../gradle/checkstyle/checkstyle.xml")
+ configProperties = ['config_loc' : "${rootDir}/../gradle/checkstyle/"]
+ toolVersion '8.23'
}
}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayElementAtFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayElementAtFunction.java
index 2697f8db..e9ee2b22 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayElementAtFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayElementAtFunction.java
@@ -6,15 +6,26 @@
package com.linkedin.transport.examples;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdInteger;
+import com.linkedin.transport.api.data.ArrayData;
import com.linkedin.transport.api.udf.StdUDF2;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class ArrayElementAtFunction extends StdUDF2 implements TopLevelStdUDF {
+/**
+ * Another way to define this class using generics can look like this
+ *
+ * public class ArrayElementAtFunction extends StdUDF2, Integer, K> implements TopLevelStdUDF {
+ *
+ * @Override
+ * public K eval(ArrayData a1, Integer idx) {
+ * return a1.get(idx);
+ * }
+ *
+ * }
+ *
+ */
+public class ArrayElementAtFunction extends StdUDF2 implements TopLevelStdUDF {
@Override
public String getFunctionName() {
@@ -40,7 +51,7 @@ public String getOutputParameterSignature() {
}
@Override
- public StdData eval(StdArray a1, StdInteger idx) {
- return a1.get(idx.get());
+ public Object eval(ArrayData a1, Integer idx) {
+ return a1.get(idx);
}
}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayFillFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayFillFunction.java
index ae5a9ac1..a9cff404 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayFillFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/ArrayFillFunction.java
@@ -7,16 +7,14 @@
import com.google.common.collect.ImmutableList;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdLong;
+import com.linkedin.transport.api.data.ArrayData;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF2;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class ArrayFillFunction extends StdUDF2 implements TopLevelStdUDF {
+public class ArrayFillFunction extends StdUDF2> implements TopLevelStdUDF {
private StdType _arrayType;
@@ -40,9 +38,9 @@ public void init(StdFactory stdFactory) {
}
@Override
- public StdArray eval(StdData a, StdLong length) {
- StdArray array = getStdFactory().createArray(_arrayType);
- for (int i = 0; i < length.get(); i++) {
+ public ArrayData eval(K a, Long length) {
+ ArrayData array = getStdFactory().createArray(_arrayType);
+ for (int i = 0; i < length; i++) {
array.add(a);
}
return array;
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/BinaryDuplicateFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/BinaryDuplicateFunction.java
new file mode 100644
index 00000000..8252b816
--- /dev/null
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/BinaryDuplicateFunction.java
@@ -0,0 +1,46 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.examples;
+
+import com.google.common.collect.ImmutableList;
+import com.linkedin.transport.api.udf.StdUDF1;
+import com.linkedin.transport.api.udf.TopLevelStdUDF;
+import java.nio.ByteBuffer;
+import java.util.List;
+
+
+public class BinaryDuplicateFunction extends StdUDF1 implements TopLevelStdUDF {
+ @Override
+ public ByteBuffer eval(ByteBuffer byteBuffer) {
+ ByteBuffer results = ByteBuffer.allocate(2 * byteBuffer.array().length);
+ for (int i = 0; i < 2; i++) {
+ for (byte b : byteBuffer.array()) {
+ results.put(b);
+ }
+ }
+ return results;
+ }
+
+ @Override
+ public List getInputParameterSignatures() {
+ return ImmutableList.of("varbinary");
+ }
+
+ @Override
+ public String getOutputParameterSignature() {
+ return "varbinary";
+ }
+
+ @Override
+ public String getFunctionName() {
+ return "binary_duplicate";
+ }
+
+ @Override
+ public String getFunctionDescription() {
+ return "Duplicate a binary object";
+ }
+}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/BinaryObjectSizeFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/BinaryObjectSizeFunction.java
new file mode 100644
index 00000000..39b56cd4
--- /dev/null
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/BinaryObjectSizeFunction.java
@@ -0,0 +1,40 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.examples;
+
+import com.google.common.collect.ImmutableList;
+import com.linkedin.transport.api.udf.StdUDF1;
+import com.linkedin.transport.api.udf.TopLevelStdUDF;
+import java.nio.ByteBuffer;
+import java.util.List;
+
+
+public class BinaryObjectSizeFunction extends StdUDF1 implements TopLevelStdUDF {
+ @Override
+ public Integer eval(ByteBuffer byteBuffer) {
+ return byteBuffer.array().length;
+ }
+
+ @Override
+ public List getInputParameterSignatures() {
+ return ImmutableList.of("varbinary");
+ }
+
+ @Override
+ public String getOutputParameterSignature() {
+ return "integer";
+ }
+
+ @Override
+ public String getFunctionName() {
+ return "binary_size";
+ }
+
+ @Override
+ public String getFunctionDescription() {
+ return "Gets the size of a binary object";
+ }
+}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/FileLookupFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/FileLookupFunction.java
index 8112e443..e9ed378f 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/FileLookupFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/FileLookupFunction.java
@@ -7,9 +7,6 @@
import com.google.common.base.Preconditions;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdBoolean;
-import com.linkedin.transport.api.data.StdInteger;
-import com.linkedin.transport.api.data.StdString;
import com.linkedin.transport.api.udf.StdUDF2;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.io.BufferedReader;
@@ -21,14 +18,14 @@
import org.apache.commons.io.IOUtils;
-public class FileLookupFunction extends StdUDF2 implements TopLevelStdUDF {
+public class FileLookupFunction extends StdUDF2 implements TopLevelStdUDF {
private Set ids;
@Override
- public StdBoolean eval(StdString filename, StdInteger intToCheck) {
+ public Boolean eval(String filename, Integer intToCheck) {
Preconditions.checkNotNull(intToCheck, "Integer to check should not be null");
- return getStdFactory().createBoolean(ids.contains(intToCheck.get()));
+ return ids.contains(intToCheck);
}
@Override
@@ -57,8 +54,8 @@ public String getFunctionDescription() {
}
@Override
- public String[] getRequiredFiles(StdString filename, StdInteger intToCheck) {
- return new String[]{filename.get()};
+ public String[] getRequiredFiles(String filename, Integer intToCheck) {
+ return new String[]{filename};
}
public void processRequiredFiles(String[] localPaths) {
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapFromTwoArraysFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapFromTwoArraysFunction.java
index 6fd99981..d6002a50 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapFromTwoArraysFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapFromTwoArraysFunction.java
@@ -7,15 +7,16 @@
import com.google.common.collect.ImmutableList;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdMap;
+import com.linkedin.transport.api.data.ArrayData;
+import com.linkedin.transport.api.data.MapData;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF2;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class MapFromTwoArraysFunction extends StdUDF2 implements TopLevelStdUDF {
+public class MapFromTwoArraysFunction extends StdUDF2, ArrayData, MapData>
+ implements TopLevelStdUDF {
private StdType _mapType;
@@ -35,16 +36,16 @@ public String getOutputParameterSignature() {
@Override
public void init(StdFactory stdFactory) {
super.init(stdFactory);
- // Note: we create the _mapType once in init() and then reuse it to create StdMap objects
+ // Note: we create the _mapType once in init() and then reuse it to create MapData objects
_mapType = getStdFactory().createStdType(getOutputParameterSignature());
}
@Override
- public StdMap eval(StdArray a1, StdArray a2) {
+ public MapData eval(ArrayData a1, ArrayData a2) {
if (a1.size() != a2.size()) {
return null;
}
- StdMap map = getStdFactory().createMap(_mapType);
+ MapData map = getStdFactory().createMap(_mapType);
for (int i = 0; i < a1.size(); i++) {
map.put(a1.get(i), a2.get(i));
}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapKeySetFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapKeySetFunction.java
index a76c4403..af81e024 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapKeySetFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapKeySetFunction.java
@@ -7,16 +7,15 @@
import com.google.common.collect.ImmutableList;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdMap;
+import com.linkedin.transport.api.data.ArrayData;
+import com.linkedin.transport.api.data.MapData;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF1;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class MapKeySetFunction extends StdUDF1 implements TopLevelStdUDF {
+public class MapKeySetFunction extends StdUDF1, ArrayData> implements TopLevelStdUDF {
private StdType _mapType;
@@ -39,9 +38,9 @@ public void init(StdFactory stdFactory) {
}
@Override
- public StdArray eval(StdMap map) {
- StdArray result = getStdFactory().createArray(_mapType);
- for (StdData key : map.keySet()) {
+ public ArrayData eval(MapData map) {
+ ArrayData result = getStdFactory().createArray(_mapType);
+ for (K key : map.keySet()) {
result.add(key);
}
return result;
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapValuesFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapValuesFunction.java
index f22ff7f7..82b34ef1 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapValuesFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/MapValuesFunction.java
@@ -7,16 +7,15 @@
import com.google.common.collect.ImmutableList;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdMap;
+import com.linkedin.transport.api.data.ArrayData;
+import com.linkedin.transport.api.data.MapData;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF1;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class MapValuesFunction extends StdUDF1 implements TopLevelStdUDF {
+public class MapValuesFunction extends StdUDF1, ArrayData> implements TopLevelStdUDF {
private StdType _mapType;
@@ -39,9 +38,9 @@ public void init(StdFactory stdFactory) {
}
@Override
- public StdArray eval(StdMap map) {
- StdArray result = getStdFactory().createArray(_mapType);
- for (StdData value : map.values()) {
+ public ArrayData eval(MapData map) {
+ ArrayData result = getStdFactory().createArray(_mapType);
+ for (V value : map.values()) {
result.add(value);
}
return result;
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddDoubleFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddDoubleFunction.java
new file mode 100644
index 00000000..80b0fb2b
--- /dev/null
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddDoubleFunction.java
@@ -0,0 +1,28 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.examples;
+
+import com.google.common.collect.ImmutableList;
+import com.linkedin.transport.api.udf.StdUDF2;
+import java.util.List;
+
+
+public class NumericAddDoubleFunction extends StdUDF2 implements NumericAddFunction {
+ @Override
+ public Double eval(Double first, Double second) {
+ return first + second;
+ }
+
+ @Override
+ public List getInputParameterSignatures() {
+ return ImmutableList.of("double", "double");
+ }
+
+ @Override
+ public String getOutputParameterSignature() {
+ return "double";
+ }
+}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFloatFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFloatFunction.java
new file mode 100644
index 00000000..a2a0ab47
--- /dev/null
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFloatFunction.java
@@ -0,0 +1,28 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.examples;
+
+import com.google.common.collect.ImmutableList;
+import com.linkedin.transport.api.udf.StdUDF2;
+import java.util.List;
+
+
+public class NumericAddFloatFunction extends StdUDF2 implements NumericAddFunction {
+ @Override
+ public Float eval(Float first, Float second) {
+ return first + second;
+ }
+
+ @Override
+ public List getInputParameterSignatures() {
+ return ImmutableList.of("real", "real");
+ }
+
+ @Override
+ public String getOutputParameterSignature() {
+ return "real";
+ }
+}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFunction.java
index 9c8e26d0..76e805e2 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddFunction.java
@@ -17,6 +17,6 @@ default String getFunctionName() {
@Override
default String getFunctionDescription() {
- return "Adds two integers or longs";
+ return "Adds two integers, longs, reals, or doubles";
}
}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddIntFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddIntFunction.java
index cc5fb900..bcdb3696 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddIntFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddIntFunction.java
@@ -6,16 +6,15 @@
package com.linkedin.transport.examples;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdInteger;
import com.linkedin.transport.api.udf.StdUDF2;
import java.util.List;
-public class NumericAddIntFunction extends StdUDF2
+public class NumericAddIntFunction extends StdUDF2
implements NumericAddFunction {
@Override
- public StdInteger eval(StdInteger first, StdInteger second) {
- return getStdFactory().createInteger(first.get() + second.get());
+ public Integer eval(Integer first, Integer second) {
+ return first + second;
}
@Override
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddLongFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddLongFunction.java
index a530e586..c24d2148 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddLongFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/NumericAddLongFunction.java
@@ -6,15 +6,14 @@
package com.linkedin.transport.examples;
import com.google.common.collect.ImmutableList;
-import com.linkedin.transport.api.data.StdLong;
import com.linkedin.transport.api.udf.StdUDF2;
import java.util.List;
-public class NumericAddLongFunction extends StdUDF2 implements NumericAddFunction {
+public class NumericAddLongFunction extends StdUDF2 implements NumericAddFunction {
@Override
- public StdLong eval(StdLong first, StdLong second) {
- return getStdFactory().createLong(first.get() + second.get());
+ public Long eval(Long first, Long second) {
+ return first + second;
}
@Override
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByIndexFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByIndexFunction.java
index 5a23283d..ffa78ba7 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByIndexFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByIndexFunction.java
@@ -7,15 +7,14 @@
import com.google.common.collect.ImmutableList;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF2;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class StructCreateByIndexFunction extends StdUDF2 implements TopLevelStdUDF {
+public class StructCreateByIndexFunction extends StdUDF2 implements TopLevelStdUDF {
private StdType _field1Type;
private StdType _field2Type;
@@ -41,11 +40,11 @@ public void init(StdFactory stdFactory) {
}
@Override
- public StdStruct eval(StdData field1Value, StdData field2Value) {
- StdStruct struct = getStdFactory().createStruct(ImmutableList.of(_field1Type, _field2Type));
- struct.setField(0, field1Value);
- struct.setField(1, field2Value);
- return struct;
+ public RowData eval(Object field1Value, Object field2Value) {
+ RowData row = getStdFactory().createStruct(ImmutableList.of(_field1Type, _field2Type));
+ row.setField(0, field1Value);
+ row.setField(1, field2Value);
+ return row;
}
@Override
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByNameFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByNameFunction.java
index 36ca3472..b4f2a0c0 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByNameFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/main/java/com/linkedin/transport/examples/StructCreateByNameFunction.java
@@ -7,16 +7,14 @@
import com.google.common.collect.ImmutableList;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdString;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.api.udf.StdUDF4;
import com.linkedin.transport.api.udf.TopLevelStdUDF;
import java.util.List;
-public class StructCreateByNameFunction extends StdUDF4 implements TopLevelStdUDF {
+public class StructCreateByNameFunction extends StdUDF4 implements TopLevelStdUDF {
private StdType _field1Type;
private StdType _field2Type;
@@ -44,13 +42,13 @@ public void init(StdFactory stdFactory) {
}
@Override
- public StdStruct eval(StdString field1Name, StdData field1Value, StdString field2Name, StdData field2Value) {
- StdStruct struct = getStdFactory().createStruct(
- ImmutableList.of(field1Name.get(), field2Name.get()),
+ public RowData eval(String field1Name, Object field1Value, String field2Name, Object field2Value) {
+ RowData struct = getStdFactory().createStruct(
+ ImmutableList.of(field1Name, field2Name),
ImmutableList.of(_field1Type, _field2Type)
);
- struct.setField(field1Name.get(), field1Value);
- struct.setField(field2Name.get(), field2Value);
+ struct.setField(field1Name, field1Value);
+ struct.setField(field2Name, field2Value);
return struct;
}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestBinaryDuplicateFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestBinaryDuplicateFunction.java
new file mode 100644
index 00000000..5e74e47c
--- /dev/null
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestBinaryDuplicateFunction.java
@@ -0,0 +1,59 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.examples;
+
+import com.google.common.collect.ImmutableList;
+import com.google.common.collect.ImmutableMap;
+import com.linkedin.transport.api.udf.StdUDF;
+import com.linkedin.transport.api.udf.TopLevelStdUDF;
+import com.linkedin.transport.test.AbstractStdUDFTest;
+import com.linkedin.transport.test.spi.StdTester;
+import java.nio.ByteBuffer;
+import java.util.List;
+import java.util.Map;
+import org.testng.annotations.Test;
+
+
+public class TestBinaryDuplicateFunction extends AbstractStdUDFTest {
+ @Override
+ protected Map, List>> getTopLevelStdUDFClassesAndImplementations() {
+ return ImmutableMap.of(BinaryDuplicateFunction.class, ImmutableList.of(BinaryDuplicateFunction.class));
+ }
+
+ @Test
+ public void testBinaryDuplicateASCII() {
+ StdTester tester = getTester();
+ testBinaryDuplicateStringHelper(tester, "bar", "barbar");
+ testBinaryDuplicateStringHelper(tester, "", "");
+ testBinaryDuplicateStringHelper(tester, "foobar", "foobarfoobar");
+ }
+
+ @Test
+ public void testBinaryDuplicateUnicode() {
+ StdTester tester = getTester();
+ testBinaryDuplicateStringHelper(tester, "こんにちは世界", "こんにちは世界こんにちは世界");
+ testBinaryDuplicateStringHelper(tester, "\uD83D\uDE02", "\uD83D\uDE02\uD83D\uDE02");
+ }
+
+ private void testBinaryDuplicateStringHelper(StdTester tester, String input, String expectedOutput) {
+ ByteBuffer inputBuffer = ByteBuffer.wrap(input.getBytes());
+ ByteBuffer expected = ByteBuffer.wrap(expectedOutput.getBytes());
+ tester.check(functionCall("binary_duplicate", inputBuffer), expected, "varbinary");
+ }
+
+ @Test
+ public void testBinaryDuplicate() {
+ StdTester tester = getTester();
+ testBinaryDuplicateHelper(tester, new byte[] {1, 2, 3}, new byte[] {1, 2, 3, 1, 2, 3});
+ testBinaryDuplicateHelper(tester, new byte[] {-1, -2, -3}, new byte[] {-1, -2, -3, -1, -2, -3});
+ }
+
+ private void testBinaryDuplicateHelper(StdTester tester, byte[] input, byte[] expectedOutput) {
+ ByteBuffer inputBuffer = ByteBuffer.wrap(input);
+ ByteBuffer expected = ByteBuffer.wrap(expectedOutput);
+ tester.check(functionCall("binary_duplicate", inputBuffer), expected, "varbinary");
+ }
+}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestBinaryObjectSizeFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestBinaryObjectSizeFunction.java
new file mode 100644
index 00000000..b10bf0fc
--- /dev/null
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestBinaryObjectSizeFunction.java
@@ -0,0 +1,36 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.examples;
+
+import com.google.common.collect.ImmutableList;
+import com.google.common.collect.ImmutableMap;
+import com.linkedin.transport.api.udf.StdUDF;
+import com.linkedin.transport.api.udf.TopLevelStdUDF;
+import com.linkedin.transport.test.AbstractStdUDFTest;
+import com.linkedin.transport.test.spi.StdTester;
+import java.nio.ByteBuffer;
+import java.util.List;
+import java.util.Map;
+import org.testng.annotations.Test;
+
+
+public class TestBinaryObjectSizeFunction extends AbstractStdUDFTest {
+ @Override
+ protected Map, List>> getTopLevelStdUDFClassesAndImplementations() {
+ return ImmutableMap.of(BinaryObjectSizeFunction.class, ImmutableList.of(BinaryObjectSizeFunction.class));
+ }
+
+ @Test
+ public void tesBinaryObjectSize() {
+ StdTester tester = getTester();
+ ByteBuffer argTest1 = ByteBuffer.wrap("foo".getBytes());
+ ByteBuffer argTest2 = ByteBuffer.wrap("".getBytes());
+ ByteBuffer argTest3 = ByteBuffer.wrap("fooBar".getBytes());
+ tester.check(functionCall("binary_size", argTest1), 3, "integer");
+ tester.check(functionCall("binary_size", argTest2), 0, "integer");
+ tester.check(functionCall("binary_size", argTest3), 6, "integer");
+ }
+}
diff --git a/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestNumericAddFunction.java b/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestNumericAddFunction.java
index 8114d7e4..12f9791e 100644
--- a/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestNumericAddFunction.java
+++ b/transportable-udfs-examples/transportable-udfs-example-udfs/src/test/java/com/linkedin/transport/examples/TestNumericAddFunction.java
@@ -21,7 +21,11 @@ public class TestNumericAddFunction extends AbstractStdUDFTest {
@Override
protected Map, List>> getTopLevelStdUDFClassesAndImplementations() {
return ImmutableMap.of(NumericAddFunction.class,
- ImmutableList.of(NumericAddIntFunction.class, NumericAddLongFunction.class));
+ ImmutableList.of(
+ NumericAddIntFunction.class,
+ NumericAddLongFunction.class,
+ NumericAddFloatFunction.class,
+ NumericAddDoubleFunction.class));
}
@Test
@@ -29,5 +33,15 @@ public void testNumericAdd() {
StdTester tester = getTester();
tester.check(functionCall("numeric_add", 1, 2), 3, "integer");
tester.check(functionCall("numeric_add", 1L, 2L), 3L, "bigint");
+ tester.check(functionCall("numeric_add", 3.0, 4.0), 7.0, "double");
+
+ Object expectedResult;
+ if (tester.getClass().getCanonicalName().contains("HiveTester")) {
+ // Note that org.apache.hive.service.cli.Column.addValue() converts any elements in RowSet from float to double
+ expectedResult = 5.0;
+ } else {
+ expectedResult = 5.0f;
+ }
+ tester.check(functionCall("numeric_add", 2.0f, 3.0f), expectedResult, "real");
}
}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveFactory.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveFactory.java
index 18b81f03..c058e56b 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveFactory.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveFactory.java
@@ -5,23 +5,14 @@
*/
package com.linkedin.transport.hive;
-import com.google.common.base.Preconditions;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdBoolean;
-import com.linkedin.transport.api.data.StdInteger;
-import com.linkedin.transport.api.data.StdLong;
-import com.linkedin.transport.api.data.StdMap;
-import com.linkedin.transport.api.data.StdString;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.ArrayData;
+import com.linkedin.transport.api.data.MapData;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.api.types.StdType;
-import com.linkedin.transport.hive.data.HiveArray;
-import com.linkedin.transport.hive.data.HiveBoolean;
-import com.linkedin.transport.hive.data.HiveInteger;
-import com.linkedin.transport.hive.data.HiveLong;
-import com.linkedin.transport.hive.data.HiveMap;
-import com.linkedin.transport.hive.data.HiveString;
-import com.linkedin.transport.hive.data.HiveStruct;
+import com.linkedin.transport.hive.data.HiveArrayData;
+import com.linkedin.transport.hive.data.HiveMapData;
+import com.linkedin.transport.hive.data.HiveRowData;
import com.linkedin.transport.hive.types.objectinspector.CacheableObjectInspectorConverters;
import com.linkedin.transport.hive.typesystem.HiveTypeFactory;
import com.linkedin.transport.typesystem.AbstractBoundVariables;
@@ -38,7 +29,6 @@
import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspectorConverters.Converter;
import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspectorFactory;
import org.apache.hadoop.hive.serde2.objectinspector.StructObjectInspector;
-import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory;
public class HiveFactory implements StdFactory {
@@ -54,44 +44,23 @@ public HiveFactory(AbstractBoundVariables boundVariables) {
}
@Override
- public StdInteger createInteger(int value) {
- return new HiveInteger(value, PrimitiveObjectInspectorFactory.javaIntObjectInspector, this);
- }
-
- @Override
- public StdLong createLong(long value) {
- return new HiveLong(value, PrimitiveObjectInspectorFactory.javaLongObjectInspector, this);
- }
-
- @Override
- public StdBoolean createBoolean(boolean value) {
- return new HiveBoolean(value, PrimitiveObjectInspectorFactory.javaBooleanObjectInspector, this);
- }
-
- @Override
- public StdString createString(String value) {
- Preconditions.checkNotNull(value, "Cannot create a null StdString");
- return new HiveString(value, PrimitiveObjectInspectorFactory.javaStringObjectInspector, this);
- }
-
- @Override
- public StdArray createArray(StdType stdType, int expectedSize) {
+ public ArrayData createArray(StdType stdType, int expectedSize) {
ListObjectInspector listObjectInspector = (ListObjectInspector) stdType.underlyingType();
- return new HiveArray(
+ return new HiveArrayData(
new ArrayList(expectedSize),
ObjectInspectorFactory.getStandardListObjectInspector(listObjectInspector.getListElementObjectInspector()),
this);
}
@Override
- public StdArray createArray(StdType stdType) {
+ public ArrayData createArray(StdType stdType) {
return createArray(stdType, 0);
}
@Override
- public StdMap createMap(StdType stdType) {
+ public MapData createMap(StdType stdType) {
MapObjectInspector mapObjectInspector = (MapObjectInspector) stdType.underlyingType();
- return new HiveMap(
+ return new HiveMapData(
new HashMap(),
ObjectInspectorFactory.getStandardMapObjectInspector(
mapObjectInspector.getMapKeyObjectInspector(),
@@ -100,8 +69,8 @@ public StdMap createMap(StdType stdType) {
}
@Override
- public StdStruct createStruct(List fieldNames, List fieldTypes) {
- return new HiveStruct(
+ public RowData createStruct(List fieldNames, List fieldTypes) {
+ return new HiveRowData(
new ArrayList(Arrays.asList(new Object[fieldTypes.size()])),
ObjectInspectorFactory.getStandardStructObjectInspector(
fieldNames,
@@ -111,16 +80,16 @@ public StdStruct createStruct(List fieldNames, List fieldTypes)
}
@Override
- public StdStruct createStruct(List fieldTypes) {
+ public RowData createStruct(List fieldTypes) {
List fieldNames =
IntStream.range(0, fieldTypes.size()).mapToObj(i -> "field" + i).collect(Collectors.toList());
return createStruct(fieldNames, fieldTypes);
}
@Override
- public StdStruct createStruct(StdType stdType) {
+ public RowData createStruct(StdType stdType) {
StructObjectInspector structObjectInspector = (StructObjectInspector) stdType.underlyingType();
- return new HiveStruct(
+ return new HiveRowData(
new ArrayList(Arrays.asList(new Object[structObjectInspector.getAllStructFieldRefs().size()])),
ObjectInspectorFactory.getStandardStructObjectInspector(
structObjectInspector.getAllStructFieldRefs()
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveWrapper.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveWrapper.java
index b2980836..2b06d7a4 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveWrapper.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/HiveWrapper.java
@@ -6,30 +6,42 @@
package com.linkedin.transport.hive;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdData;
import com.linkedin.transport.api.types.StdType;
-import com.linkedin.transport.hive.data.HiveArray;
-import com.linkedin.transport.hive.data.HiveBoolean;
-import com.linkedin.transport.hive.data.HiveInteger;
-import com.linkedin.transport.hive.data.HiveLong;
-import com.linkedin.transport.hive.data.HiveMap;
-import com.linkedin.transport.hive.data.HiveString;
-import com.linkedin.transport.hive.data.HiveStruct;
+import com.linkedin.transport.hive.data.HiveArrayData;
+import com.linkedin.transport.hive.data.HiveData;
+import com.linkedin.transport.hive.data.HiveMapData;
+import com.linkedin.transport.hive.data.HiveRowData;
import com.linkedin.transport.hive.types.HiveArrayType;
import com.linkedin.transport.hive.types.HiveBooleanType;
+import com.linkedin.transport.hive.types.HiveBinaryType;
+import com.linkedin.transport.hive.types.HiveDoubleType;
+import com.linkedin.transport.hive.types.HiveFloatType;
import com.linkedin.transport.hive.types.HiveIntegerType;
import com.linkedin.transport.hive.types.HiveLongType;
import com.linkedin.transport.hive.types.HiveMapType;
import com.linkedin.transport.hive.types.HiveStringType;
-import com.linkedin.transport.hive.types.HiveStructType;
+import com.linkedin.transport.hive.types.HiveRowType;
import com.linkedin.transport.hive.types.HiveUnknownType;
+import java.nio.ByteBuffer;
import org.apache.hadoop.hive.serde2.objectinspector.ListObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.MapObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.PrimitiveObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.StructObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.BinaryObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.BooleanObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.DoubleObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.FloatObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.IntObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.LongObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableBinaryObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableBooleanObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableDoubleObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableFloatObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableIntObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableLongObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.SettableStringObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.StringObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.VoidObjectInspector;
@@ -39,22 +51,22 @@ public final class HiveWrapper {
private HiveWrapper() {
}
- public static StdData createStdData(Object hiveData, ObjectInspector hiveObjectInspector, StdFactory stdFactory) {
- if (hiveObjectInspector instanceof IntObjectInspector) {
- return new HiveInteger(hiveData, (IntObjectInspector) hiveObjectInspector, stdFactory);
- } else if (hiveObjectInspector instanceof LongObjectInspector) {
- return new HiveLong(hiveData, (LongObjectInspector) hiveObjectInspector, stdFactory);
- } else if (hiveObjectInspector instanceof BooleanObjectInspector) {
- return new HiveBoolean(hiveData, (BooleanObjectInspector) hiveObjectInspector, stdFactory);
- } else if (hiveObjectInspector instanceof StringObjectInspector) {
- return new HiveString(hiveData, (StringObjectInspector) hiveObjectInspector, stdFactory);
+ public static Object createStdData(Object hiveData, ObjectInspector hiveObjectInspector, StdFactory stdFactory) {
+ if (hiveObjectInspector instanceof IntObjectInspector || hiveObjectInspector instanceof LongObjectInspector
+ || hiveObjectInspector instanceof FloatObjectInspector || hiveObjectInspector instanceof DoubleObjectInspector
+ || hiveObjectInspector instanceof BooleanObjectInspector
+ || hiveObjectInspector instanceof StringObjectInspector) {
+ return ((PrimitiveObjectInspector) hiveObjectInspector).getPrimitiveJavaObject(hiveData);
+ } else if (hiveObjectInspector instanceof BinaryObjectInspector) {
+ BinaryObjectInspector binaryObjectInspector = (BinaryObjectInspector) hiveObjectInspector;
+ return hiveData == null ? null : ByteBuffer.wrap(binaryObjectInspector.getPrimitiveJavaObject(hiveData));
} else if (hiveObjectInspector instanceof ListObjectInspector) {
ListObjectInspector listObjectInspector = (ListObjectInspector) hiveObjectInspector;
- return new HiveArray(hiveData, listObjectInspector, stdFactory);
+ return new HiveArrayData(hiveData, listObjectInspector, stdFactory);
} else if (hiveObjectInspector instanceof MapObjectInspector) {
- return new HiveMap(hiveData, hiveObjectInspector, stdFactory);
+ return new HiveMapData(hiveData, hiveObjectInspector, stdFactory);
} else if (hiveObjectInspector instanceof StructObjectInspector) {
- return new HiveStruct(((StructObjectInspector) hiveObjectInspector).getStructFieldsDataAsList(hiveData).toArray(),
+ return new HiveRowData(((StructObjectInspector) hiveObjectInspector).getStructFieldsDataAsList(hiveData).toArray(),
hiveObjectInspector, stdFactory);
} else if (hiveObjectInspector instanceof VoidObjectInspector) {
return null;
@@ -72,16 +84,68 @@ public static StdType createStdType(ObjectInspector hiveObjectInspector) {
return new HiveBooleanType((BooleanObjectInspector) hiveObjectInspector);
} else if (hiveObjectInspector instanceof StringObjectInspector) {
return new HiveStringType((StringObjectInspector) hiveObjectInspector);
+ } else if (hiveObjectInspector instanceof FloatObjectInspector) {
+ return new HiveFloatType((FloatObjectInspector) hiveObjectInspector);
+ } else if (hiveObjectInspector instanceof DoubleObjectInspector) {
+ return new HiveDoubleType((DoubleObjectInspector) hiveObjectInspector);
+ } else if (hiveObjectInspector instanceof BinaryObjectInspector) {
+ return new HiveBinaryType((BinaryObjectInspector) hiveObjectInspector);
} else if (hiveObjectInspector instanceof ListObjectInspector) {
return new HiveArrayType((ListObjectInspector) hiveObjectInspector);
} else if (hiveObjectInspector instanceof MapObjectInspector) {
return new HiveMapType((MapObjectInspector) hiveObjectInspector);
} else if (hiveObjectInspector instanceof StructObjectInspector) {
- return new HiveStructType((StructObjectInspector) hiveObjectInspector);
+ return new HiveRowType((StructObjectInspector) hiveObjectInspector);
} else if (hiveObjectInspector instanceof VoidObjectInspector) {
return new HiveUnknownType((VoidObjectInspector) hiveObjectInspector);
}
assert false : "Unrecognized Hive ObjectInspector: " + hiveObjectInspector.getClass();
return null;
}
+
+ public static Object getPlatformDataForObjectInspector(Object transportData, ObjectInspector oi) {
+ if (transportData == null) {
+ return null;
+ } else if (oi instanceof IntObjectInspector) {
+ return ((SettableIntObjectInspector) oi).create((Integer) transportData);
+ } else if (oi instanceof LongObjectInspector) {
+ return ((SettableLongObjectInspector) oi).create((Long) transportData);
+ } else if (oi instanceof FloatObjectInspector) {
+ return ((SettableFloatObjectInspector) oi).create((Float) transportData);
+ } else if (oi instanceof DoubleObjectInspector) {
+ return ((SettableDoubleObjectInspector) oi).create((Double) transportData);
+ } else if (oi instanceof BooleanObjectInspector) {
+ return ((SettableBooleanObjectInspector) oi).create((Boolean) transportData);
+ } else if (oi instanceof StringObjectInspector) {
+ return ((SettableStringObjectInspector) oi).create((String) transportData);
+ } else if (oi instanceof BinaryObjectInspector) {
+ return ((SettableBinaryObjectInspector) oi).create(((ByteBuffer) transportData).array());
+ } else {
+ return ((HiveData) transportData).getUnderlyingDataForObjectInspector(oi);
+ }
+ }
+
+ public static Object getStandardObject(Object transportData) {
+ if (transportData == null) {
+ return null;
+ } else if (transportData instanceof Integer) {
+ return PrimitiveObjectInspectorFactory.writableIntObjectInspector.create((Integer) transportData);
+ } else if (transportData instanceof Long) {
+ return PrimitiveObjectInspectorFactory.writableLongObjectInspector.create((Long) transportData);
+ } else if (transportData instanceof Float) {
+ return PrimitiveObjectInspectorFactory.writableFloatObjectInspector.create((Float) transportData);
+ } else if (transportData instanceof Double) {
+ return PrimitiveObjectInspectorFactory.writableDoubleObjectInspector.create((Double) transportData);
+ } else if (transportData instanceof Boolean) {
+ return PrimitiveObjectInspectorFactory.writableBooleanObjectInspector.create((Boolean) transportData);
+ } else if (transportData instanceof String) {
+ return PrimitiveObjectInspectorFactory.writableStringObjectInspector.create((String) transportData);
+ } else if (transportData instanceof ByteBuffer) {
+ return PrimitiveObjectInspectorFactory.writableBinaryObjectInspector.create(((ByteBuffer) transportData).array());
+ } else {
+ return ((HiveData) transportData).getUnderlyingDataForObjectInspector(
+ ((HiveData) transportData).getUnderlyingObjectInspector()
+ );
+ }
+ }
}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/StdUdfWrapper.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/StdUdfWrapper.java
index 3b3da9ab..5f14689e 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/StdUdfWrapper.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/StdUdfWrapper.java
@@ -7,7 +7,6 @@
import com.linkedin.transport.api.StdFactory;
import com.linkedin.transport.api.data.PlatformData;
-import com.linkedin.transport.api.data.StdData;
import com.linkedin.transport.api.udf.StdUDF;
import com.linkedin.transport.api.udf.StdUDF0;
import com.linkedin.transport.api.udf.StdUDF1;
@@ -23,6 +22,7 @@
import com.linkedin.transport.utils.FileSystemUtils;
import java.io.FileNotFoundException;
import java.io.IOException;
+import java.nio.ByteBuffer;
import java.util.Arrays;
import java.util.List;
import java.util.stream.IntStream;
@@ -35,6 +35,9 @@
import org.apache.hadoop.hive.ql.udf.generic.GenericUDF;
import org.apache.hadoop.hive.serde2.objectinspector.ConstantObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.PrimitiveObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.BinaryObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory;
/**
@@ -49,7 +52,8 @@ public abstract class StdUdfWrapper extends GenericUDF {
protected StdFactory _stdFactory;
private boolean[] _nullableArguments;
private String[] _distributedCacheFiles;
- private StdData[] _args;
+ private Object[] _args;
+ private ObjectInspector _outputObjectInspector;
/**
* Given input object inspectors, this method matches them to the expected type signatures, and finds bindings to the
@@ -70,7 +74,8 @@ public ObjectInspector initialize(ObjectInspector[] arguments) {
_stdUdf.init(_stdFactory);
_requiredFilesProcessed = false;
createStdData();
- return hiveTypeInference.getOutputDataType();
+ _outputObjectInspector= hiveTypeInference.getOutputDataType();
+ return _outputObjectInspector;
}
@Override
@@ -108,14 +113,23 @@ protected boolean containsNullValuedNonNullableConstants() {
return false;
}
- protected StdData wrap(DeferredObject hiveDeferredObject, StdData stdData) {
+ protected Object wrap(DeferredObject hiveDeferredObject, ObjectInspector inputObjectInspector, Object stdData) {
try {
Object hiveObject = hiveDeferredObject.get();
- if (hiveObject != null) {
- ((PlatformData) stdData).setUnderlyingData(hiveObject);
- return stdData;
+ if (inputObjectInspector instanceof BinaryObjectInspector) {
+ return hiveObject == null ? null : ByteBuffer.wrap(
+ ((BinaryObjectInspector) inputObjectInspector).getPrimitiveJavaObject(hiveObject)
+ );
+ }
+ if (inputObjectInspector instanceof PrimitiveObjectInspector) {
+ return ((PrimitiveObjectInspector) inputObjectInspector).getPrimitiveJavaObject(hiveObject);
} else {
- return null;
+ if (hiveObject != null) {
+ ((PlatformData) stdData).setUnderlyingData(hiveObject);
+ return stdData;
+ } else {
+ return null;
+ }
}
} catch (HiveException e) {
throw new RuntimeException("Cannot extract Hive Object from Deferred Object");
@@ -127,21 +141,35 @@ protected StdData wrap(DeferredObject hiveDeferredObject, StdData stdData) {
protected abstract Class extends TopLevelStdUDF> getTopLevelUdfClass();
protected void createStdData() {
- _args = new StdData[_inputObjectInspectors.length];
+ _args = new Object[_inputObjectInspectors.length];
for (int i = 0; i < _inputObjectInspectors.length; i++) {
_args[i] = HiveWrapper.createStdData(null, _inputObjectInspectors[i], _stdFactory);
}
}
- private StdData[] wrapArguments(DeferredObject[] deferredObjects) {
- return IntStream.range(0, _args.length).mapToObj(i -> wrap(deferredObjects[i], _args[i])).toArray(StdData[]::new);
+ private Object getPlatformData(Object transportData) {
+ if (transportData == null) {
+ return null;
+ } else if (transportData instanceof Integer || transportData instanceof Long || transportData instanceof Boolean
+ || transportData instanceof String || transportData instanceof Float || transportData instanceof Double ||
+ transportData instanceof ByteBuffer) {
+ return HiveWrapper.getPlatformDataForObjectInspector(transportData, _outputObjectInspector);
+ } else {
+ return ((PlatformData) transportData).getUnderlyingData();
+ }
+ }
+
+ private Object[] wrapArguments(DeferredObject[] deferredObjects) {
+ return IntStream.range(0, _args.length).mapToObj(
+ i -> wrap(deferredObjects[i], _inputObjectInspectors[i], _args[i])
+ ).toArray(Object[]::new);
}
- private StdData[] wrapConstants() {
+ private Object[] wrapConstants() {
return Arrays.stream(_inputObjectInspectors)
.map(oi -> (oi instanceof ConstantObjectInspector) ? HiveWrapper.createStdData(
((ConstantObjectInspector) oi).getWritableConstantValue(), oi, _stdFactory) : null)
- .toArray(StdData[]::new);
+ .toArray(Object[]::new);
}
@Override
@@ -152,8 +180,8 @@ public Object evaluate(DeferredObject[] arguments) throws HiveException {
if (!_requiredFilesProcessed) {
processRequiredFiles();
}
- StdData[] args = wrapArguments(arguments);
- StdData result;
+ Object[] args = wrapArguments(arguments);
+ Object result;
switch (args.length) {
case 0:
result = ((StdUDF0) _stdUdf).eval();
@@ -185,7 +213,7 @@ public Object evaluate(DeferredObject[] arguments) throws HiveException {
default:
throw new UnsupportedOperationException("eval not yet supported for StdUDF" + args.length);
}
- return result == null ? null : ((PlatformData) result).getUnderlyingData();
+ return getPlatformData(result);
}
@Override
@@ -193,7 +221,7 @@ public String[] getRequiredFiles() {
if (containsNullValuedNonNullableConstants()) {
return new String[]{};
}
- StdData[] args = wrapConstants();
+ Object[] args = wrapConstants();
String[] requiredFiles;
switch (args.length) {
case 0:
@@ -231,7 +259,7 @@ public String[] getRequiredFiles() {
}
_distributedCacheFiles = Arrays.stream(requiredFiles).map(requiredFile -> {
try {
- return FileSystemUtils.resolveLatest(requiredFile, FileSystemUtils.getHDFSFileSystem());
+ return FileSystemUtils.resolveLatest(requiredFile);
} catch (IOException e) {
throw new RuntimeException("Failed to resolve path: [" + requiredFile + "].", e);
}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveArray.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveArrayData.java
similarity index 72%
rename from transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveArray.java
rename to transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveArrayData.java
index 57cb0e8c..d0bf8ec4 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveArray.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveArrayData.java
@@ -6,8 +6,7 @@
package com.linkedin.transport.hive.data;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdData;
+import com.linkedin.transport.api.data.ArrayData;
import com.linkedin.transport.hive.HiveWrapper;
import java.util.Iterator;
import org.apache.hadoop.hive.serde2.objectinspector.ListObjectInspector;
@@ -15,12 +14,12 @@
import org.apache.hadoop.hive.serde2.objectinspector.SettableListObjectInspector;
-public class HiveArray extends HiveData implements StdArray {
+public class HiveArrayData extends HiveData implements ArrayData {
final ListObjectInspector _listObjectInspector;
final ObjectInspector _elementObjectInspector;
- public HiveArray(Object object, ObjectInspector objectInspector, StdFactory stdFactory) {
+ public HiveArrayData(Object object, ObjectInspector objectInspector, StdFactory stdFactory) {
super(stdFactory);
_object = object;
_listObjectInspector = (ListObjectInspector) objectInspector;
@@ -33,19 +32,21 @@ public int size() {
}
@Override
- public StdData get(int idx) {
- return HiveWrapper.createStdData(_listObjectInspector.getListElement(_object, idx), _elementObjectInspector,
+ public E get(int idx) {
+ return (E) HiveWrapper.createStdData(
+ _listObjectInspector.getListElement(_object, idx),
+ _elementObjectInspector,
_stdFactory);
}
@Override
- public void add(StdData e) {
+ public void add(E e) {
if (_listObjectInspector instanceof SettableListObjectInspector) {
SettableListObjectInspector settableListObjectInspector = (SettableListObjectInspector) _listObjectInspector;
int originalSize = size();
settableListObjectInspector.resize(_object, originalSize + 1);
settableListObjectInspector.set(_object, originalSize,
- ((HiveData) e).getUnderlyingDataForObjectInspector(_elementObjectInspector));
+ HiveWrapper.getPlatformDataForObjectInspector(e, _elementObjectInspector));
_isObjectModified = true;
} else {
throw new RuntimeException("Attempt to modify an immutable Hive object of type: "
@@ -59,8 +60,8 @@ public ObjectInspector getUnderlyingObjectInspector() {
}
@Override
- public Iterator iterator() {
- return new Iterator() {
+ public Iterator iterator() {
+ return new Iterator() {
int size = size();
int currentIndex = 0;
@@ -70,8 +71,8 @@ public boolean hasNext() {
}
@Override
- public StdData next() {
- StdData element = HiveWrapper.createStdData(_listObjectInspector.getListElement(_object, currentIndex),
+ public E next() {
+ E element = (E) HiveWrapper.createStdData(_listObjectInspector.getListElement(_object, currentIndex),
_elementObjectInspector, _stdFactory);
currentIndex++;
return element;
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveBoolean.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveBoolean.java
deleted file mode 100644
index b4537170..00000000
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveBoolean.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.hive.data;
-
-import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdBoolean;
-import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
-import org.apache.hadoop.hive.serde2.objectinspector.primitive.BooleanObjectInspector;
-
-
-public class HiveBoolean extends HiveData implements StdBoolean {
-
- final BooleanObjectInspector _booleanObjectInspector;
-
- public HiveBoolean(Object object, BooleanObjectInspector booleanObjectInspector, StdFactory stdFactory) {
- super(stdFactory);
- _object = object;
- _booleanObjectInspector = booleanObjectInspector;
- }
-
- @Override
- public boolean get() {
- return _booleanObjectInspector.get(_object);
- }
-
- @Override
- public ObjectInspector getUnderlyingObjectInspector() {
- return _booleanObjectInspector;
- }
-}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveData.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveData.java
index 51beb456..94337266 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveData.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveData.java
@@ -59,10 +59,6 @@ public ObjectInspector getStandardObjectInspector() {
getUnderlyingObjectInspector(), ObjectInspectorUtils.ObjectInspectorCopyOption.WRITABLE);
}
- public Object getStandardObject() {
- return getUnderlyingDataForObjectInspector(getStandardObjectInspector());
- }
-
private Object getObjectFromCache(ObjectInspector oi) {
if (_isObjectModified) {
_cachedObjectsForObjectInspectors.clear();
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveInteger.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveInteger.java
deleted file mode 100644
index a1d2e38f..00000000
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveInteger.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.hive.data;
-
-import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdInteger;
-import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
-import org.apache.hadoop.hive.serde2.objectinspector.primitive.IntObjectInspector;
-
-
-public class HiveInteger extends HiveData implements StdInteger {
-
- final IntObjectInspector _intObjectInspector;
-
- public HiveInteger(Object object, IntObjectInspector intObjectInspector, StdFactory stdFactory) {
- super(stdFactory);
- _object = object;
- _intObjectInspector = intObjectInspector;
- }
-
- @Override
- public int get() {
- return _intObjectInspector.get(_object);
- }
-
- @Override
- public ObjectInspector getUnderlyingObjectInspector() {
- return _intObjectInspector;
- }
-}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveLong.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveLong.java
deleted file mode 100644
index 0b662b59..00000000
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveLong.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.hive.data;
-
-import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdLong;
-import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
-import org.apache.hadoop.hive.serde2.objectinspector.primitive.LongObjectInspector;
-
-
-public class HiveLong extends HiveData implements StdLong {
-
- final LongObjectInspector _longObjectInspector;
-
- public HiveLong(Object object, LongObjectInspector longObjectInspector, StdFactory stdFactory) {
- super(stdFactory);
- _object = object;
- _longObjectInspector = longObjectInspector;
- }
-
- @Override
- public long get() {
- return _longObjectInspector.get(_object);
- }
-
- @Override
- public ObjectInspector getUnderlyingObjectInspector() {
- return _longObjectInspector;
- }
-}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveMap.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveMapData.java
similarity index 66%
rename from transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveMap.java
rename to transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveMapData.java
index 70f5132b..54da6042 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveMap.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveMapData.java
@@ -6,8 +6,7 @@
package com.linkedin.transport.hive.data;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdMap;
+import com.linkedin.transport.api.data.MapData;
import com.linkedin.transport.hive.HiveWrapper;
import java.util.AbstractCollection;
import java.util.AbstractSet;
@@ -20,13 +19,13 @@
import org.apache.hadoop.hive.serde2.objectinspector.SettableMapObjectInspector;
-public class HiveMap extends HiveData implements StdMap {
+public class HiveMapData extends HiveData implements MapData {
final MapObjectInspector _mapObjectInspector;
final ObjectInspector _keyObjectInspector;
final ObjectInspector _valueObjectInspector;
- public HiveMap(Object object, ObjectInspector objectInspector, StdFactory stdFactory) {
+ public HiveMapData(Object object, ObjectInspector objectInspector, StdFactory stdFactory) {
super(stdFactory);
_object = object;
_mapObjectInspector = (MapObjectInspector) objectInspector;
@@ -40,30 +39,30 @@ public int size() {
}
@Override
- public StdData get(StdData key) {
+ public V get(K key) {
MapObjectInspector mapOI = _mapObjectInspector;
Object mapObj = _object;
Object keyObj;
try {
- keyObj = ((HiveData) key).getUnderlyingDataForObjectInspector(_keyObjectInspector);
+ keyObj = HiveWrapper.getPlatformDataForObjectInspector(key, _keyObjectInspector);
} catch (RuntimeException e) {
// Cannot convert key argument to Map's KeyOI. So convert both the map and the key arg to
// objects having standard OIs
mapOI = (MapObjectInspector) getStandardObjectInspector();
- mapObj = getStandardObject();
- keyObj = ((HiveData) key).getStandardObject();
+ mapObj = HiveWrapper.getStandardObject(this);
+ keyObj = HiveWrapper.getStandardObject(key);
}
- return HiveWrapper.createStdData(
+ return (V) HiveWrapper.createStdData(
mapOI.getMapValueElement(mapObj, keyObj),
mapOI.getMapValueObjectInspector(), _stdFactory);
}
@Override
- public void put(StdData key, StdData value) {
+ public void put(K key, V value) {
if (_mapObjectInspector instanceof SettableMapObjectInspector) {
- Object keyObj = ((HiveData) key).getUnderlyingDataForObjectInspector(_keyObjectInspector);
- Object valueObj = ((HiveData) value).getUnderlyingDataForObjectInspector(_valueObjectInspector);
+ Object keyObj = HiveWrapper.getPlatformDataForObjectInspector(key, _keyObjectInspector);
+ Object valueObj = HiveWrapper.getPlatformDataForObjectInspector(value, _valueObjectInspector);
((SettableMapObjectInspector) _mapObjectInspector).put(
_object,
@@ -79,11 +78,11 @@ public void put(StdData key, StdData value) {
//TODO: Cache the result of .getMap(_object) below for subsequent calls.
@Override
- public Set keySet() {
- return new AbstractSet() {
+ public Set keySet() {
+ return new AbstractSet() {
@Override
- public Iterator iterator() {
- return new Iterator() {
+ public Iterator iterator() {
+ return new Iterator() {
Iterator mapKeyIterator = _mapObjectInspector.getMap(_object).keySet().iterator();
@Override
@@ -92,26 +91,26 @@ public boolean hasNext() {
}
@Override
- public StdData next() {
- return HiveWrapper.createStdData(mapKeyIterator.next(), _keyObjectInspector, _stdFactory);
+ public K next() {
+ return (K) HiveWrapper.createStdData(mapKeyIterator.next(), _keyObjectInspector, _stdFactory);
}
};
}
@Override
public int size() {
- return HiveMap.this.size();
+ return HiveMapData.this.size();
}
};
}
//TODO: Cache the result of .getMap(_object) below for subsequent calls.
@Override
- public Collection values() {
- return new AbstractCollection() {
+ public Collection values() {
+ return new AbstractCollection() {
@Override
- public Iterator iterator() {
- return new Iterator() {
+ public Iterator iterator() {
+ return new Iterator() {
Iterator mapValueIterator = _mapObjectInspector.getMap(_object).values().iterator();
@Override
@@ -120,30 +119,30 @@ public boolean hasNext() {
}
@Override
- public StdData next() {
- return HiveWrapper.createStdData(mapValueIterator.next(), _valueObjectInspector, _stdFactory);
+ public V next() {
+ return (V) HiveWrapper.createStdData(mapValueIterator.next(), _valueObjectInspector, _stdFactory);
}
};
}
@Override
public int size() {
- return HiveMap.this.size();
+ return HiveMapData.this.size();
}
};
}
@Override
- public boolean containsKey(StdData key) {
+ public boolean containsKey(K key) {
Object mapObj = _object;
Object keyObj;
try {
- keyObj = ((HiveData) key).getUnderlyingDataForObjectInspector(_keyObjectInspector);
+ keyObj = HiveWrapper.getPlatformDataForObjectInspector(key, _keyObjectInspector);
} catch (RuntimeException e) {
// Cannot convert key argument to Map's KeyOI. So convertboth the map and the key arg to
// objects having standard OIs
- mapObj = getStandardObject();
- keyObj = ((HiveData) key).getStandardObject();
+ mapObj = HiveWrapper.getStandardObject(this);
+ keyObj = HiveWrapper.getStandardObject(key);
}
return ((Map) mapObj).containsKey(keyObj);
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveStruct.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveRowData.java
similarity index 79%
rename from transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveStruct.java
rename to transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveRowData.java
index 80872eff..5704374e 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveStruct.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveRowData.java
@@ -6,8 +6,7 @@
package com.linkedin.transport.hive.data;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdData;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.hive.HiveWrapper;
import java.util.List;
import java.util.stream.Collectors;
@@ -18,18 +17,18 @@
import org.apache.hadoop.hive.serde2.objectinspector.StructObjectInspector;
-public class HiveStruct extends HiveData implements StdStruct {
+public class HiveRowData extends HiveData implements RowData {
StructObjectInspector _structObjectInspector;
- public HiveStruct(Object object, ObjectInspector objectInspector, StdFactory stdFactory) {
+ public HiveRowData(Object object, ObjectInspector objectInspector, StdFactory stdFactory) {
super(stdFactory);
_object = object;
_structObjectInspector = (StructObjectInspector) objectInspector;
}
@Override
- public StdData getField(int index) {
+ public Object getField(int index) {
StructField structField = _structObjectInspector.getAllStructFieldRefs().get(index);
return HiveWrapper.createStdData(
_structObjectInspector.getStructFieldData(_object, structField),
@@ -38,7 +37,7 @@ public StdData getField(int index) {
}
@Override
- public StdData getField(String name) {
+ public Object getField(String name) {
StructField structField = _structObjectInspector.getStructFieldRef(name);
return HiveWrapper.createStdData(
_structObjectInspector.getStructFieldData(_object, structField),
@@ -47,11 +46,11 @@ public StdData getField(String name) {
}
@Override
- public void setField(int index, StdData value) {
+ public void setField(int index, Object value) {
if (_structObjectInspector instanceof SettableStructObjectInspector) {
StructField field = _structObjectInspector.getAllStructFieldRefs().get(index);
((SettableStructObjectInspector) _structObjectInspector).setStructFieldData(_object,
- field, ((HiveData) value).getUnderlyingDataForObjectInspector(field.getFieldObjectInspector())
+ field, HiveWrapper.getPlatformDataForObjectInspector(value, field.getFieldObjectInspector())
);
_isObjectModified = true;
} else {
@@ -61,11 +60,11 @@ public void setField(int index, StdData value) {
}
@Override
- public void setField(String name, StdData value) {
+ public void setField(String name, Object value) {
if (_structObjectInspector instanceof SettableStructObjectInspector) {
StructField field = _structObjectInspector.getStructFieldRef(name);
((SettableStructObjectInspector) _structObjectInspector).setStructFieldData(_object,
- field, ((HiveData) value).getUnderlyingDataForObjectInspector(field.getFieldObjectInspector()));
+ field, HiveWrapper.getPlatformDataForObjectInspector(value, field.getFieldObjectInspector()));
_isObjectModified = true;
} else {
throw new RuntimeException("Attempt to modify an immutable Hive object of type: "
@@ -74,7 +73,7 @@ public void setField(String name, StdData value) {
}
@Override
- public List fields() {
+ public List fields() {
return IntStream.range(0, _structObjectInspector.getAllStructFieldRefs().size()).mapToObj(i -> getField(i))
.collect(Collectors.toList());
}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveString.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveString.java
deleted file mode 100644
index 83310309..00000000
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/data/HiveString.java
+++ /dev/null
@@ -1,33 +0,0 @@
-/**
- * Copyright 2018 LinkedIn Corporation. All rights reserved.
- * Licensed under the BSD-2 Clause license.
- * See LICENSE in the project root for license information.
- */
-package com.linkedin.transport.hive.data;
-
-import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdString;
-import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
-import org.apache.hadoop.hive.serde2.objectinspector.primitive.StringObjectInspector;
-
-
-public class HiveString extends HiveData implements StdString {
-
- final StringObjectInspector _stringObjectInspector;
-
- public HiveString(Object object, StringObjectInspector stringObjectInspector, StdFactory stdFactory) {
- super(stdFactory);
- _object = object;
- _stringObjectInspector = stringObjectInspector;
- }
-
- @Override
- public String get() {
- return _stringObjectInspector.getPrimitiveJavaObject(_object);
- }
-
- @Override
- public ObjectInspector getUnderlyingObjectInspector() {
- return _stringObjectInspector;
- }
-}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveBinaryType.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveBinaryType.java
new file mode 100644
index 00000000..bc21a5d7
--- /dev/null
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveBinaryType.java
@@ -0,0 +1,24 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.hive.types;
+
+import com.linkedin.transport.api.types.StdBinaryType;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.BinaryObjectInspector;
+
+
+public class HiveBinaryType implements StdBinaryType {
+
+ private final BinaryObjectInspector _binaryObjectInspector;
+
+ public HiveBinaryType(BinaryObjectInspector binaryObjectInspector) {
+ _binaryObjectInspector = binaryObjectInspector;
+ }
+
+ @Override
+ public Object underlyingType() {
+ return _binaryObjectInspector;
+ }
+}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveDoubleType.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveDoubleType.java
new file mode 100644
index 00000000..83659632
--- /dev/null
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveDoubleType.java
@@ -0,0 +1,24 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.hive.types;
+
+import com.linkedin.transport.api.types.StdDoubleType;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.DoubleObjectInspector;
+
+
+public class HiveDoubleType implements StdDoubleType {
+
+ private final DoubleObjectInspector _doubleObjectInspector;
+
+ public HiveDoubleType(DoubleObjectInspector doubleObjectInspector) {
+ _doubleObjectInspector = doubleObjectInspector;
+ }
+
+ @Override
+ public Object underlyingType() {
+ return _doubleObjectInspector;
+ }
+}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveFloatType.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveFloatType.java
new file mode 100644
index 00000000..9f9107f6
--- /dev/null
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveFloatType.java
@@ -0,0 +1,24 @@
+/**
+ * Copyright 2018 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.hive.types;
+
+import com.linkedin.transport.api.types.StdFloatType;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.FloatObjectInspector;
+
+
+public class HiveFloatType implements StdFloatType {
+
+ private final FloatObjectInspector _floatObjectInspector;
+
+ public HiveFloatType(FloatObjectInspector floatObjectInspector) {
+ _floatObjectInspector = floatObjectInspector;
+ }
+
+ @Override
+ public Object underlyingType() {
+ return _floatObjectInspector;
+ }
+}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveStructType.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveRowType.java
similarity index 83%
rename from transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveStructType.java
rename to transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveRowType.java
index f4393776..c9ceef43 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveStructType.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/types/HiveRowType.java
@@ -5,7 +5,7 @@
*/
package com.linkedin.transport.hive.types;
-import com.linkedin.transport.api.types.StdStructType;
+import com.linkedin.transport.api.types.RowType;
import com.linkedin.transport.api.types.StdType;
import com.linkedin.transport.hive.HiveWrapper;
import java.util.List;
@@ -13,11 +13,11 @@
import org.apache.hadoop.hive.serde2.objectinspector.StructObjectInspector;
-public class HiveStructType implements StdStructType {
+public class HiveRowType implements RowType {
final StructObjectInspector _structObjectInspector;
- public HiveStructType(StructObjectInspector structObjectInspector) {
+ public HiveRowType(StructObjectInspector structObjectInspector) {
_structObjectInspector = structObjectInspector;
}
diff --git a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/typesystem/HiveTypeSystem.java b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/typesystem/HiveTypeSystem.java
index 977779f1..4fa3f596 100644
--- a/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/typesystem/HiveTypeSystem.java
+++ b/transportable-udfs-hive/src/main/java/com/linkedin/transport/hive/typesystem/HiveTypeSystem.java
@@ -14,7 +14,10 @@
import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspectorFactory;
import org.apache.hadoop.hive.serde2.objectinspector.StructObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.BinaryObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.BooleanObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.DoubleObjectInspector;
+import org.apache.hadoop.hive.serde2.objectinspector.primitive.FloatObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.IntObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.LongObjectInspector;
import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory;
@@ -69,6 +72,21 @@ protected boolean isStringType(ObjectInspector dataType) {
return dataType instanceof StringObjectInspector;
}
+ @Override
+ protected boolean isFloatType(ObjectInspector dataType) {
+ return dataType instanceof FloatObjectInspector;
+ }
+
+ @Override
+ protected boolean isDoubleType(ObjectInspector dataType) {
+ return dataType instanceof DoubleObjectInspector;
+ }
+
+ @Override
+ protected boolean isBinaryType(ObjectInspector dataType) {
+ return dataType instanceof BinaryObjectInspector;
+ }
+
@Override
protected boolean isArrayType(ObjectInspector dataType) {
return dataType instanceof ListObjectInspector;
@@ -104,6 +122,21 @@ protected ObjectInspector createStringType() {
return PrimitiveObjectInspectorFactory.javaStringObjectInspector;
}
+ @Override
+ protected ObjectInspector createFloatType() {
+ return PrimitiveObjectInspectorFactory.javaFloatObjectInspector;
+ }
+
+ @Override
+ protected ObjectInspector createDoubleType() {
+ return PrimitiveObjectInspectorFactory.javaDoubleObjectInspector;
+ }
+
+ @Override
+ protected ObjectInspector createBinaryType() {
+ return PrimitiveObjectInspectorFactory.javaByteArrayObjectInspector;
+ }
+
@Override
protected ObjectInspector createUnknownType() {
return PrimitiveObjectInspectorFactory.javaVoidObjectInspector;
diff --git a/transportable-udfs-plugin/build.gradle b/transportable-udfs-plugin/build.gradle
index f465e780..6f4ade36 100644
--- a/transportable-udfs-plugin/build.gradle
+++ b/transportable-udfs-plugin/build.gradle
@@ -27,7 +27,7 @@ def writeVersionInfo = { file ->
ant.propertyfile(file: file) {
entry(key: "transport-version", value: version)
entry(key: "hive-version", value: '1.2.2')
- entry(key: "presto-version", value: '319')
+ entry(key: "presto-version", value: '333')
entry(key: "spark-version", value: '2.3.0')
entry(key: "scala-version", value: '2.11.8')
}
diff --git a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Defaults.java b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Defaults.java
index ee96e18d..64a967f4 100644
--- a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Defaults.java
+++ b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Defaults.java
@@ -11,6 +11,7 @@
import com.linkedin.transport.codegen.SparkWrapperGenerator;
import com.linkedin.transport.plugin.packaging.DistributionPackaging;
import com.linkedin.transport.plugin.packaging.ShadedJarPackaging;
+import com.linkedin.transport.plugin.packaging.ThinJarPackaging;
import java.io.IOException;
import java.io.InputStream;
import java.util.List;
@@ -72,7 +73,7 @@ private static Properties loadDefaultVersions() {
// converters drop dependencies with classifiers, so we apply this dependency explicitly
getDependencyConfiguration(RUNTIME_ONLY, "io.prestosql:presto-main", "presto", "tests")
),
- new DistributionPackaging()),
+ ImmutableList.of(new ThinJarPackaging(), new DistributionPackaging())),
new Platform(
"hive",
Language.JAVA,
@@ -85,7 +86,7 @@ private static Properties loadDefaultVersions() {
getDependencyConfiguration(RUNTIME_ONLY, "com.linkedin.transport:transportable-udfs-test-hive",
"transport")
),
- new ShadedJarPackaging(ImmutableList.of("org.apache.hadoop", "org.apache.hive"), null)),
+ ImmutableList.of(new ShadedJarPackaging(ImmutableList.of("org.apache.hadoop", "org.apache.hive"), null))),
new Platform(
"spark",
Language.SCALA,
@@ -99,9 +100,9 @@ private static Properties loadDefaultVersions() {
getDependencyConfiguration(RUNTIME_ONLY, "com.linkedin.transport:transportable-udfs-test-spark",
"transport")
),
- new ShadedJarPackaging(
+ ImmutableList.of(new ShadedJarPackaging(
ImmutableList.of("org.apache.hadoop", "org.apache.spark"),
- ImmutableList.of("com.linkedin.transport.spark.**"))
+ ImmutableList.of("com.linkedin.transport.spark.**")))
)
);
diff --git a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Platform.java b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Platform.java
index fed24365..b3d87679 100644
--- a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Platform.java
+++ b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/Platform.java
@@ -20,11 +20,11 @@ public class Platform {
private final Class extends WrapperGenerator> _wrapperGeneratorClass;
private final List _defaultWrapperDependencyConfigurations;
private final List _defaultTestDependencyConfigurations;
- private final Packaging _packaging;
+ private final List _packaging;
public Platform(String name, Language language, Class extends WrapperGenerator> wrapperGeneratorClass,
List defaultWrapperDependencyConfigurations,
- List defaultTestDependencyConfigurations, Packaging packaging) {
+ List defaultTestDependencyConfigurations, List packaging) {
_name = name;
_language = language;
_wrapperGeneratorClass = wrapperGeneratorClass;
@@ -53,7 +53,7 @@ public List getDefaultTestDependencyConfigurations() {
return _defaultTestDependencyConfigurations;
}
- public Packaging getPackaging() {
+ public List getPackaging() {
return _packaging;
}
}
diff --git a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPlugin.java b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPlugin.java
index 74958b6d..2f8e2984 100644
--- a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPlugin.java
+++ b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPlugin.java
@@ -12,6 +12,7 @@
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.List;
+import java.util.stream.Collectors;
import org.gradle.api.Plugin;
import org.gradle.api.Project;
import org.gradle.api.Task;
@@ -25,6 +26,8 @@
import org.gradle.api.tasks.TaskProvider;
import org.gradle.api.tasks.testing.Test;
import org.gradle.language.base.plugins.LifecycleBasePlugin;
+import org.gradle.testing.jacoco.plugins.JacocoPlugin;
+import org.gradle.testing.jacoco.plugins.JacocoTaskExtension;
import static com.linkedin.transport.plugin.ConfigurationType.*;
import static com.linkedin.transport.plugin.SourceSetUtils.*;
@@ -45,18 +48,32 @@
public class TransportPlugin implements Plugin {
public void apply(Project project) {
+
+ TransportPluginConfig extension = project.getExtensions().create("transport", TransportPluginConfig.class, project);
+
project.getPlugins().withType(JavaPlugin.class, (javaPlugin) -> {
project.getPlugins().apply(ScalaPlugin.class);
project.getPlugins().apply(DistributionPlugin.class);
project.getConfigurations().create(ShadowBasePlugin.getCONFIGURATION_NAME());
JavaPluginConvention javaConvention = project.getConvention().getPlugin(JavaPluginConvention.class);
- SourceSet mainSourceSet = javaConvention.getSourceSets().getByName("main");
- SourceSet testSourceSet = javaConvention.getSourceSets().getByName("test");
+ SourceSet mainSourceSet = javaConvention.getSourceSets().getByName(extension.mainSourceSetName);
+ SourceSet testSourceSet = javaConvention.getSourceSets().getByName(extension.testSourceSetName);
configureBaseSourceSets(project, mainSourceSet, testSourceSet);
Defaults.DEFAULT_PLATFORMS.forEach(
- platform -> configurePlatform(project, platform, mainSourceSet, testSourceSet));
+ platform -> configurePlatform(project, platform, mainSourceSet, testSourceSet, extension.outputDirFile));
+ });
+ // Disable Jacoco for platform test tasks as it is known to cause issues with Presto and Hive tests
+ project.getPlugins().withType(JacocoPlugin.class, (jacocoPlugin) -> {
+ Defaults.DEFAULT_PLATFORMS.forEach(platform -> {
+ project.getTasksByName(testTaskName(platform), true).forEach(task -> {
+ JacocoTaskExtension jacocoExtension = task.getExtensions().findByType(JacocoTaskExtension.class);
+ if (jacocoExtension != null) {
+ jacocoExtension.setEnabled(false);
+ }
+ });
+ });
});
}
@@ -77,8 +94,9 @@ private void configureBaseSourceSets(Project project, SourceSet mainSourceSet, S
/**
* Configures SourceSets, dependencies and tasks related to each Transport UDF platform
*/
- private void configurePlatform(Project project, Platform platform, SourceSet mainSourceSet, SourceSet testSourceSet) {
- SourceSet sourceSet = configureSourceSet(project, platform, mainSourceSet);
+ private void configurePlatform(Project project, Platform platform, SourceSet mainSourceSet, SourceSet testSourceSet,
+ File baseOutputDir) {
+ SourceSet sourceSet = configureSourceSet(project, platform, mainSourceSet, baseOutputDir);
configureGenerateWrappersTask(project, platform, mainSourceSet, sourceSet);
List> packagingTasks =
configurePackagingTasks(project, platform, sourceSet, mainSourceSet);
@@ -94,9 +112,9 @@ private void configurePlatform(Project project, Platform platform, SourceSet mai
* configurations and configures the default dependencies required for compilation and runtime of the wrapper
* SourceSet
*/
- private SourceSet configureSourceSet(Project project, Platform platform, SourceSet mainSourceSet) {
+ private SourceSet configureSourceSet(Project project, Platform platform, SourceSet mainSourceSet, File baseOutputDir) {
JavaPluginConvention javaConvention = project.getConvention().getPlugin(JavaPluginConvention.class);
- Path platformBaseDir = Paths.get(project.getBuildDir().toString(), "generatedWrappers", platform.getName());
+ Path platformBaseDir = Paths.get(baseOutputDir.toString(), "generatedWrappers", platform.getName());
Path wrapperSourceOutputDir = platformBaseDir.resolve("sources");
Path wrapperResourceOutputDir = platformBaseDir.resolve("resources");
@@ -186,7 +204,9 @@ private TaskProvider configureGenerateWrappersTask(Project
*/
private List> configurePackagingTasks(Project project, Platform platform,
SourceSet sourceSet, SourceSet mainSourceSet) {
- return platform.getPackaging().configurePackagingTasks(project, platform, sourceSet, mainSourceSet);
+ return platform.getPackaging().stream()
+ .flatMap(p -> p.configurePackagingTasks(project, platform, sourceSet, mainSourceSet).stream())
+ .collect(Collectors.toList());
}
/**
@@ -230,7 +250,7 @@ task prestoTest(type: Test, dependsOn: test) {
}
*/
- return project.getTasks().register(platform.getName() + "Test", Test.class, task -> {
+ return project.getTasks().register(testTaskName(platform), Test.class, task -> {
task.setGroup(LifecycleBasePlugin.VERIFICATION_GROUP);
task.setDescription("Runs Transport UDF tests on " + platform.getName());
task.setTestClassesDirs(testSourceSet.getOutput().getClassesDirs());
@@ -239,4 +259,8 @@ task prestoTest(type: Test, dependsOn: test) {
task.mustRunAfter(project.getTasks().named("test"));
});
}
+
+ private String testTaskName(Platform platform) {
+ return platform.getName() + "Test";
+ }
}
diff --git a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPluginConfig.java b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPluginConfig.java
new file mode 100644
index 00000000..deff599a
--- /dev/null
+++ b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/TransportPluginConfig.java
@@ -0,0 +1,58 @@
+/**
+ * Copyright 2020 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.plugin;
+
+import java.io.File;
+import org.gradle.api.Project;
+
+
+/**
+ * Custom configuration for the {@link TransportPlugin}.
+ */
+public class TransportPluginConfig {
+
+ /**
+ * Available gradle property names.
+ */
+ private static final String MAIN_SOURCE_SET_NAME_PROP = "transport-plugin-main-source-set";
+ private static final String TEST_SOURCE_SET_NAME_PROP = "transport-plugin-test-source-set";
+ private static final String OUTPUT_DIR_PROP = "transport-plugin-output-dir";
+
+ /**
+ * The main source set used by the plugin.
+ */
+ public String mainSourceSetName;
+ /**
+ * The test source set used by the plugin.
+ */
+ public String testSourceSetName;
+ /**
+ * The output code-gen directory, relative the to the project directory.
+ */
+ public File outputDirFile;
+
+ /**
+ * Create a config object from the gradle {@link Project}.
+ *
+ * On creation, we attempt to populate config values using gradle properties or set to default values.
+ */
+ public TransportPluginConfig(Project project) {
+ mainSourceSetName = getPropertyOrDefault(project, MAIN_SOURCE_SET_NAME_PROP, "main");
+ testSourceSetName = getPropertyOrDefault(project, TEST_SOURCE_SET_NAME_PROP, "test");
+ outputDirFile = project.hasProperty(OUTPUT_DIR_PROP)
+ ? project.file(project.property(OUTPUT_DIR_PROP).toString())
+ : project.getBuildDir(); // Build dir by default
+ }
+
+ /**
+ * Retrieve a string property from a gradle project if it has been set, a default value otherwise.
+ */
+ private String getPropertyOrDefault(Project project, String propertyName, String defaultValue) {
+ return project.hasProperty(propertyName)
+ ? project.property(propertyName).toString()
+ : defaultValue;
+ }
+}
diff --git a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/DistributionPackaging.java b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/DistributionPackaging.java
index 44ef8b55..13de71db 100644
--- a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/DistributionPackaging.java
+++ b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/DistributionPackaging.java
@@ -66,18 +66,18 @@ public List> configurePackagingTasks(Project projec
*/
private TaskProvider createThinJarTask(Project project, SourceSet sourceSet, String platformName) {
/*
- task ThinJar(type: Jar, dependsOn: prestoClasses) {
- classifier 'platformName'
+ task DistThinJar(type: Jar, dependsOn: prestoClasses) {
+ classifier '-dist-thin'
from sourceSets..output
from sourceSets..resources
}
*/
- return project.getTasks().register(sourceSet.getTaskName(null, "thinJar"), Jar.class, task -> {
+ return project.getTasks().register(sourceSet.getTaskName(null, "distThinJar"), Jar.class, task -> {
task.dependsOn(project.getTasks().named(sourceSet.getClassesTaskName()));
task.setDescription("Assembles a thin jar archive containing the " + platformName
+ " classes to be included in the distribution");
- task.setClassifier(platformName + "Thin");
+ task.setClassifier(platformName + "-dist-thin");
task.from(sourceSet.getOutput());
task.from(sourceSet.getResources());
});
diff --git a/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/ThinJarPackaging.java b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/ThinJarPackaging.java
new file mode 100644
index 00000000..7367733c
--- /dev/null
+++ b/transportable-udfs-plugin/src/main/java/com/linkedin/transport/plugin/packaging/ThinJarPackaging.java
@@ -0,0 +1,48 @@
+/**
+ * Copyright 2019 LinkedIn Corporation. All rights reserved.
+ * Licensed under the BSD-2 Clause license.
+ * See LICENSE in the project root for license information.
+ */
+package com.linkedin.transport.plugin.packaging;
+
+import com.google.common.collect.ImmutableList;
+import com.linkedin.transport.plugin.Platform;
+import java.util.List;
+import org.gradle.api.Project;
+import org.gradle.api.Task;
+import org.gradle.api.artifacts.Dependency;
+import org.gradle.api.tasks.SourceSet;
+import org.gradle.api.tasks.TaskProvider;
+import org.gradle.api.tasks.bundling.Jar;
+
+
+/**
+ * A {@link Packaging} class which generates a Jar containing the autogenerated code for the platform
+ */
+public class ThinJarPackaging implements Packaging {
+
+ @Override
+ public List> configurePackagingTasks(Project project, Platform platform,
+ SourceSet platformSourceSet, SourceSet mainSourceSet) {
+ /*
+ task ThinJar(type: Jar, dependsOn: prestoClasses) {
+ classifier '-thin'
+ from sourceSets..output
+ from sourceSets..resources
+ }
+ */
+
+ TaskProvider extends Task> thinJarTask =
+ project.getTasks().register(platformSourceSet.getTaskName(null, "thinJar"), Jar.class, task -> {
+ task.dependsOn(project.getTasks().named(platformSourceSet.getClassesTaskName()));
+ task.setDescription("Assembles a thin jar archive containing the " + platform.getName()
+ + " classes to be included in the distribution");
+ task.setClassifier(platform.getName() + "-thin");
+ task.from(platformSourceSet.getOutput());
+ task.from(platformSourceSet.getResources());
+ });
+
+ project.getArtifacts().add(Dependency.ARCHIVES_CONFIGURATION, thinJarTask);
+ return ImmutableList.of(thinJarTask);
+ }
+}
diff --git a/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/FileSystemClient.java b/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/FileSystemClient.java
index 850ef304..f964433d 100644
--- a/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/FileSystemClient.java
+++ b/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/FileSystemClient.java
@@ -53,7 +53,9 @@ public String copyToLocalFile(String remoteFilename) {
Path remotePath = new Path(remoteFilename);
Path localPath = new Path(Paths.get(getAndCreateLocalDir(), new File(remoteFilename).getName()).toString());
FileSystem fs = remotePath.getFileSystem(conf);
- String resolvedRemoteFilename = FileSystemUtils.resolveLatest(remoteFilename, fs);
+ // It is important to pass the custom configuration object to FileSystemUtils since we load some extra
+ // properties from etc/**.xml in getConfiguration() for Presto
+ String resolvedRemoteFilename = FileSystemUtils.resolveLatest(remoteFilename, conf);
Path resolvedRemotePath = new Path(resolvedRemoteFilename);
fs.copyToLocalFile(resolvedRemotePath, localPath);
return localPath.toString();
diff --git a/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/PrestoFactory.java b/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/PrestoFactory.java
index c5e0c53e..94539835 100644
--- a/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/PrestoFactory.java
+++ b/transportable-udfs-presto/src/main/java/com/linkedin/transport/presto/PrestoFactory.java
@@ -5,38 +5,31 @@
*/
package com.linkedin.transport.presto;
-import com.google.common.base.Preconditions;
+import com.google.common.collect.ImmutableSet;
import com.linkedin.transport.api.StdFactory;
-import com.linkedin.transport.api.data.StdArray;
-import com.linkedin.transport.api.data.StdBoolean;
-import com.linkedin.transport.api.data.StdInteger;
-import com.linkedin.transport.api.data.StdLong;
-import com.linkedin.transport.api.data.StdMap;
-import com.linkedin.transport.api.data.StdString;
-import com.linkedin.transport.api.data.StdStruct;
+import com.linkedin.transport.api.data.ArrayData;
+import com.linkedin.transport.api.data.MapData;
+import com.linkedin.transport.api.data.RowData;
import com.linkedin.transport.api.types.StdType;
-import com.linkedin.transport.presto.data.PrestoArray;
-import com.linkedin.transport.presto.data.PrestoBoolean;
-import com.linkedin.transport.presto.data.PrestoInteger;
-import com.linkedin.transport.presto.data.PrestoLong;
-import com.linkedin.transport.presto.data.PrestoMap;
-import com.linkedin.transport.presto.data.PrestoString;
-import com.linkedin.transport.presto.data.PrestoStruct;
-import io.airlift.slice.Slices;
+import com.linkedin.transport.presto.data.PrestoArrayData;
+import com.linkedin.transport.presto.data.PrestoMapData;
+import com.linkedin.transport.presto.data.PrestoRowData;
import io.prestosql.metadata.BoundVariables;
import io.prestosql.metadata.Metadata;
-import io.prestosql.metadata.Signature;
+import io.prestosql.metadata.OperatorNotFoundException;
+import io.prestosql.metadata.ResolvedFunction;
import io.prestosql.operator.scalar.ScalarFunctionImplementation;
+import io.prestosql.spi.function.OperatorType;
import io.prestosql.spi.type.ArrayType;
import io.prestosql.spi.type.MapType;
import io.prestosql.spi.type.RowType;
import io.prestosql.spi.type.Type;
-import io.prestosql.spi.type.TypeSignature;
+import java.nio.ByteBuffer;
import java.util.List;
import java.util.stream.Collectors;
import static io.prestosql.metadata.SignatureBinder.*;
-
+import static io.prestosql.operator.TypeSignatureParser.*;
public class PrestoFactory implements StdFactory {
@@ -49,65 +42,48 @@ public PrestoFactory(BoundVariables boundVariables, Metadata metadata) {
}
@Override
- public StdInteger createInteger(int value) {
- return new PrestoInteger(value);
- }
-
- @Override
- public StdLong createLong(long value) {
- return new PrestoLong(value);
+ public ArrayData createArray(StdType stdType, int expectedSize) {
+ return new PrestoArrayData((ArrayType) stdType.underlyingType(), expectedSize, this);
}
@Override
- public StdBoolean createBoolean(boolean value) {
- return new PrestoBoolean(value);
- }
-
- @Override
- public StdString createString(String value) {
- Preconditions.checkNotNull(value, "Cannot create a null StdString");
- return new PrestoString(Slices.utf8Slice(value));
- }
-
- @Override
- public StdArray createArray(StdType stdType, int expectedSize) {
- return new PrestoArray((ArrayType) stdType.underlyingType(), expectedSize, this);
- }
-
- @Override
- public StdArray createArray(StdType stdType) {
+ public ArrayData createArray(StdType stdType) {
return createArray(stdType, 0);
}
@Override
- public StdMap createMap(StdType stdType) {
- return new PrestoMap((MapType) stdType.underlyingType(), this);
+ public MapData createMap(StdType stdType) {
+ return new PrestoMapData((MapType) stdType.underlyingType(), this);
}
@Override
- public PrestoStruct createStruct(List fieldNames, List fieldTypes) {
- return new PrestoStruct(fieldNames,
+ public PrestoRowData createStruct(List fieldNames, List fieldTypes) {
+ return new PrestoRowData(fieldNames,
fieldTypes.stream().map(stdType -> (Type) stdType.underlyingType()).collect(Collectors.toList()), this);
}
@Override
- public PrestoStruct createStruct(List fieldTypes) {
- return new PrestoStruct(
+ public PrestoRowData createStruct(List